Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
d44945b475 | ||
|
|
3b88c386a1 | ||
|
|
b95b74f0a3 | ||
|
|
5d9ff5df03 | ||
|
|
2cecb4c11c | ||
|
|
316d9b0e7f | ||
|
|
17ae8e050a | ||
|
|
f0410221d8 | ||
|
|
1c217b546c | ||
|
|
1bcdf29206 | ||
|
|
5c31deb630 | ||
|
|
c2def00bcf | ||
|
|
784dc1f0da | ||
|
|
7d93bee4bd | ||
|
|
2aecd1312e | ||
|
|
60c5cc40b2 | ||
|
|
5edb004799 | ||
|
|
764d00c249 | ||
|
|
47b77763cb | ||
|
|
7805d9b6f0 | ||
|
|
40a0e6a0aa | ||
|
|
99d63aa5f4 | ||
|
|
ee94ddc133 | ||
|
|
651c7aa3f4 | ||
|
|
1cd9cd8803 | ||
|
|
ab735d1f3a |
@@ -151,8 +151,26 @@ pkg/merge/ Schema merging
|
|||||||
pkg/models/ Internal data models
|
pkg/models/ Internal data models
|
||||||
pkg/transform/ Transformation logic
|
pkg/transform/ Transformation logic
|
||||||
pkg/pgsql/ PostgreSQL utilities
|
pkg/pgsql/ PostgreSQL utilities
|
||||||
|
pkg/sqltypes/ Nullable SQL types for generated/hand-written models (see below)
|
||||||
```
|
```
|
||||||
|
|
||||||
|
## Nullable Types (`pkg/sqltypes`)
|
||||||
|
|
||||||
|
The `bun` and `gorm` writers can generate model structs using
|
||||||
|
[`pkg/sqltypes`](./pkg/sqltypes/README.md) — nullable types (`SqlString`,
|
||||||
|
`SqlInt32`, `SqlTimeStamp`, `SqlStringArray`, …) that implement
|
||||||
|
`database/sql.Scanner`, `driver.Valuer`, and JSON/YAML/XML marshalling in one
|
||||||
|
type, selected via `--types sqltypes`. See the
|
||||||
|
[`pkg/sqltypes` README](./pkg/sqltypes/README.md) for the full type
|
||||||
|
reference, or the [`bun`](./pkg/writers/bun/README.md) /
|
||||||
|
[`gorm`](./pkg/writers/gorm/README.md) writer docs for the `--types` flag
|
||||||
|
(`sqltypes`, `stdlib`, or `baselib`). PostgreSQL array columns are the one
|
||||||
|
exception: the `bun` writer always generates native Go slices (`[]string`,
|
||||||
|
`[]int32`, …) with an explicit `array` bun tag, regardless of `--types` —
|
||||||
|
see [`bun`'s `--array-nullable`](./pkg/writers/bun/README.md#nullablearrays)
|
||||||
|
flag for nullable-array handling. The `SqlXxxArray` wrapper types remain
|
||||||
|
available in `pkg/sqltypes` and are still used by the `gorm` writer.
|
||||||
|
|
||||||
## Contributing
|
## Contributing
|
||||||
|
|
||||||
1. Register or sign in with GitHub at [git.warky.dev](https://git.warky.dev)
|
1. Register or sign in with GitHub at [git.warky.dev](https://git.warky.dev)
|
||||||
|
|||||||
@@ -0,0 +1,214 @@
|
|||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
|
||||||
|
"github.com/spf13/cobra"
|
||||||
|
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/assetloader"
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/pgsql"
|
||||||
|
)
|
||||||
|
|
||||||
|
var (
|
||||||
|
assetsDir string
|
||||||
|
assetsConn string
|
||||||
|
assetsIgnoreErrors bool
|
||||||
|
)
|
||||||
|
|
||||||
|
var assetsCmd = &cobra.Command{
|
||||||
|
Use: "assets",
|
||||||
|
Short: "Load and execute asset manifests against a database",
|
||||||
|
Long: `Load local binary and text asset files into a PostgreSQL database.
|
||||||
|
|
||||||
|
Assets are described by YAML manifests (assets.yaml) colocated with the files.
|
||||||
|
Each manifest entry specifies the file to load and the SQL call to execute.
|
||||||
|
File bytes are bound as native pgx parameters — never as SQL text literals —
|
||||||
|
so binary files stay byte-exact with no size or encoding limitations.
|
||||||
|
|
||||||
|
Manifests must live in directories that follow the naming pattern used by
|
||||||
|
relspec scripts:
|
||||||
|
{priority}_{sequence}_{name}/ or {priority}-{sequence}-{name}/
|
||||||
|
|
||||||
|
This allows asset-loading steps to be ordered correctly alongside SQL scripts
|
||||||
|
in a migrate-apply pipeline.
|
||||||
|
|
||||||
|
Manifest format (assets.yaml):
|
||||||
|
- file: invoice.md
|
||||||
|
call: |
|
||||||
|
INSERT INTO org.filepointer (rid_owner, filename, contenttype, jsonstore)
|
||||||
|
VALUES (1, :filename, 'text/markdown', jsonb_build_object('content', :bytes::text))
|
||||||
|
- file: logo.png
|
||||||
|
call: UPDATE branding SET logo = :bytes WHERE id = 1
|
||||||
|
params:
|
||||||
|
owner_id: "42"
|
||||||
|
|
||||||
|
Built-in placeholders:
|
||||||
|
:bytes — the file's raw content as bytea
|
||||||
|
:filename — the base name of the file (string)
|
||||||
|
:any_key — a static value declared in the entry's params map`,
|
||||||
|
}
|
||||||
|
|
||||||
|
var assetsListCmd = &cobra.Command{
|
||||||
|
Use: "list",
|
||||||
|
Short: "List asset manifests from a directory",
|
||||||
|
Long: `List all asset manifest entries from a directory in execution order.
|
||||||
|
|
||||||
|
The directory is scanned recursively for assets.yaml files located in
|
||||||
|
directories that follow the {priority}_{sequence}_{name} naming convention.
|
||||||
|
|
||||||
|
Example:
|
||||||
|
relspec assets list --dir ./sql`,
|
||||||
|
RunE: runAssetsList,
|
||||||
|
}
|
||||||
|
|
||||||
|
var assetsExecuteCmd = &cobra.Command{
|
||||||
|
Use: "execute",
|
||||||
|
Short: "Execute asset manifests against a database",
|
||||||
|
Long: `Execute asset manifest entries from a directory against a PostgreSQL database.
|
||||||
|
|
||||||
|
Asset manifests are executed in order: Priority (ascending), Sequence (ascending),
|
||||||
|
Directory name (alphabetical). By default, execution stops on the first error.
|
||||||
|
Use --ignore-errors to continue even when individual entries fail.
|
||||||
|
|
||||||
|
PostgreSQL Connection String Examples:
|
||||||
|
postgres://username:password@localhost:5432/database_name
|
||||||
|
postgresql://user:pass@host/dbname?sslmode=disable
|
||||||
|
|
||||||
|
Examples:
|
||||||
|
relspec assets execute --dir ./sql \
|
||||||
|
--conn "postgres://user:pass@localhost:5432/mydb"
|
||||||
|
|
||||||
|
relspec assets execute --dir ./sql \
|
||||||
|
--conn "postgres://localhost/mydb" \
|
||||||
|
--ignore-errors`,
|
||||||
|
RunE: runAssetsExecute,
|
||||||
|
}
|
||||||
|
|
||||||
|
func init() {
|
||||||
|
assetsListCmd.Flags().StringVar(&assetsDir, "dir", "", "Directory to scan for asset manifests (required)")
|
||||||
|
if err := assetsListCmd.MarkFlagRequired("dir"); err != nil {
|
||||||
|
fmt.Fprintf(os.Stderr, "Error marking dir flag as required: %v\n", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
assetsExecuteCmd.Flags().StringVar(&assetsDir, "dir", "", "Directory to scan for asset manifests (required)")
|
||||||
|
assetsExecuteCmd.Flags().StringVar(&assetsConn, "conn", "", "PostgreSQL connection string (required)")
|
||||||
|
assetsExecuteCmd.Flags().BoolVar(&assetsIgnoreErrors, "ignore-errors", false, "Continue executing even if entries fail")
|
||||||
|
if err := assetsExecuteCmd.MarkFlagRequired("dir"); err != nil {
|
||||||
|
fmt.Fprintf(os.Stderr, "Error marking dir flag as required: %v\n", err)
|
||||||
|
}
|
||||||
|
if err := assetsExecuteCmd.MarkFlagRequired("conn"); err != nil {
|
||||||
|
fmt.Fprintf(os.Stderr, "Error marking conn flag as required: %v\n", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
assetsCmd.AddCommand(assetsListCmd)
|
||||||
|
assetsCmd.AddCommand(assetsExecuteCmd)
|
||||||
|
}
|
||||||
|
|
||||||
|
func runAssetsList(cmd *cobra.Command, args []string) error {
|
||||||
|
fmt.Fprintf(os.Stderr, "\n=== Asset Manifests List ===\n")
|
||||||
|
fmt.Fprintf(os.Stderr, "Directory: %s\n\n", assetsDir)
|
||||||
|
|
||||||
|
items, err := assetloader.ScanDir(assetsDir)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("scanning directory: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(items) == 0 {
|
||||||
|
fmt.Fprintf(os.Stderr, "No asset manifests found.\n\n")
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Fprintf(os.Stderr, "Found %d asset entry(ies) in execution order:\n\n", len(items))
|
||||||
|
fmt.Fprintf(os.Stderr, "%-4s %-10s %-8s %-20s %s\n", "No.", "Priority", "Sequence", "Dir", "File")
|
||||||
|
fmt.Fprintf(os.Stderr, "%-4s %-10s %-8s %-20s %s\n", "----", "--------", "--------", "--------------------", "----")
|
||||||
|
|
||||||
|
for i, item := range items {
|
||||||
|
fmt.Fprintf(os.Stderr, "%-4d %-10d %-8d %-20s %s\n",
|
||||||
|
i+1,
|
||||||
|
item.Priority,
|
||||||
|
item.Sequence,
|
||||||
|
item.DirName,
|
||||||
|
filepath.Base(item.Entry.File),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Fprintf(os.Stderr, "\n")
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func runAssetsExecute(cmd *cobra.Command, args []string) error {
|
||||||
|
fmt.Fprintf(os.Stderr, "\n=== Asset Manifests Execution ===\n")
|
||||||
|
fmt.Fprintf(os.Stderr, "Started at: %s\n", getCurrentTimestamp())
|
||||||
|
fmt.Fprintf(os.Stderr, "Directory: %s\n", assetsDir)
|
||||||
|
fmt.Fprintf(os.Stderr, "Database: %s\n\n", maskPassword(assetsConn))
|
||||||
|
|
||||||
|
fmt.Fprintf(os.Stderr, "[1/2] Scanning asset manifests...\n")
|
||||||
|
|
||||||
|
items, err := assetloader.ScanDir(assetsDir)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("scanning directory: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(items) == 0 {
|
||||||
|
fmt.Fprintf(os.Stderr, " No asset manifests found. Nothing to execute.\n\n")
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Fprintf(os.Stderr, " ✓ Found %d asset entry(ies)\n\n", len(items))
|
||||||
|
|
||||||
|
fmt.Fprintf(os.Stderr, "[2/2] Executing assets in order (Priority → Sequence → Dir)...\n\n")
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
conn, err := pgsql.Connect(ctx, assetsConn, "assets-execute")
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("connecting to database: %w", err)
|
||||||
|
}
|
||||||
|
defer conn.Close(ctx)
|
||||||
|
|
||||||
|
successCount := 0
|
||||||
|
var failures []struct {
|
||||||
|
item assetloader.Item
|
||||||
|
err error
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, item := range items {
|
||||||
|
name := filepath.Base(item.Entry.File)
|
||||||
|
fmt.Printf("Executing asset: %s (Priority=%d, Sequence=%d, Dir=%s)\n",
|
||||||
|
name, item.Priority, item.Sequence, item.DirName)
|
||||||
|
|
||||||
|
if err := assetloader.ExecuteItem(ctx, conn, item); err != nil {
|
||||||
|
if assetsIgnoreErrors {
|
||||||
|
fmt.Printf("⚠ Error loading %s: %v (continuing due to --ignore-errors)\n", name, err)
|
||||||
|
failures = append(failures, struct {
|
||||||
|
item assetloader.Item
|
||||||
|
err error
|
||||||
|
}{item, err})
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
return fmt.Errorf("asset %s (Priority=%d, Sequence=%d): %w",
|
||||||
|
name, item.Priority, item.Sequence, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
successCount++
|
||||||
|
fmt.Printf("✓ Successfully loaded: %s\n", name)
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Fprintf(os.Stderr, "\n=== Execution Complete ===\n")
|
||||||
|
fmt.Fprintf(os.Stderr, "Completed at: %s\n", getCurrentTimestamp())
|
||||||
|
fmt.Fprintf(os.Stderr, "Total entries: %d\n", len(items))
|
||||||
|
fmt.Fprintf(os.Stderr, "Successful: %d\n", successCount)
|
||||||
|
if len(failures) > 0 {
|
||||||
|
fmt.Fprintf(os.Stderr, "Failed: %d\n", len(failures))
|
||||||
|
fmt.Fprintf(os.Stderr, "\n⚠ Failed Entries Summary (%d failed):\n", len(failures))
|
||||||
|
for i, f := range failures {
|
||||||
|
fmt.Fprintf(os.Stderr, " %d. %s (Priority=%d, Sequence=%d)\n Error: %v\n",
|
||||||
|
i+1, filepath.Base(f.item.Entry.File), f.item.Priority, f.item.Sequence, f.err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
fmt.Fprintf(os.Stderr, "\n")
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
+29
-4
@@ -1,6 +1,7 @@
|
|||||||
package main
|
package main
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
stdjson "encoding/json"
|
||||||
"fmt"
|
"fmt"
|
||||||
"os"
|
"os"
|
||||||
"strings"
|
"strings"
|
||||||
@@ -53,7 +54,9 @@ var (
|
|||||||
convertSchemaFilter string
|
convertSchemaFilter string
|
||||||
convertFlattenSchema bool
|
convertFlattenSchema bool
|
||||||
convertNullableTypes string
|
convertNullableTypes string
|
||||||
|
convertNullableArrays string
|
||||||
convertContinueOnError bool
|
convertContinueOnError bool
|
||||||
|
convertExtraFields string
|
||||||
)
|
)
|
||||||
|
|
||||||
var convertCmd = &cobra.Command{
|
var convertCmd = &cobra.Command{
|
||||||
@@ -177,8 +180,10 @@ func init() {
|
|||||||
convertCmd.Flags().StringVar(&convertPackageName, "package", "", "Package name (for code generation formats like gorm/bun)")
|
convertCmd.Flags().StringVar(&convertPackageName, "package", "", "Package name (for code generation formats like gorm/bun)")
|
||||||
convertCmd.Flags().StringVar(&convertSchemaFilter, "schema", "", "Filter to a specific schema by name (required for formats like dctx that only support single schemas)")
|
convertCmd.Flags().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().BoolVar(&convertFlattenSchema, "flatten-schema", false, "Flatten schema.table names to schema_table (useful for databases like SQLite that do not support schemas)")
|
||||||
convertCmd.Flags().StringVar(&convertNullableTypes, "types", "", "Nullable type package for code-gen writers (bun/gorm): 'resolvespec' (default) or 'stdlib' (database/sql)")
|
convertCmd.Flags().StringVar(&convertNullableTypes, "types", "", "Nullable type package for code-gen writers (bun/gorm): 'baselib' (default, Go pointer types), 'stdlib' (database/sql), or 'sqltypes'")
|
||||||
|
convertCmd.Flags().StringVar(&convertNullableArrays, "array-nullable", "", "Nullable PostgreSQL array representation for the Bun writer in stdlib/baselib --types mode: 'slice' (default, plain slice) or 'pointer_slice' (*[]T, distinguishes NULL from '{}')")
|
||||||
convertCmd.Flags().BoolVar(&convertContinueOnError, "continue-on-error", false, "Prepend \\set ON_ERROR_STOP off to generated SQL so psql continues past errors (pgsql output only)")
|
convertCmd.Flags().BoolVar(&convertContinueOnError, "continue-on-error", false, "Prepend \\set ON_ERROR_STOP off to generated SQL so psql continues past errors (pgsql output only)")
|
||||||
|
convertCmd.Flags().StringVar(&convertExtraFields, "extra-fields", "", "Path to JSON file containing extra Bun model fields to inject (bun output only); fields support target_table, name, type, bun_tag, json_tag, comment")
|
||||||
|
|
||||||
err := convertCmd.MarkFlagRequired("from")
|
err := convertCmd.MarkFlagRequired("from")
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -245,7 +250,7 @@ func runConvert(cmd *cobra.Command, args []string) error {
|
|||||||
fmt.Fprintf(os.Stderr, " Schema: %s\n", convertSchemaFilter)
|
fmt.Fprintf(os.Stderr, " Schema: %s\n", convertSchemaFilter)
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := writeDatabase(db, convertTargetType, convertTargetPath, convertPackageName, convertSchemaFilter, convertFlattenSchema, convertNullableTypes, convertContinueOnError); 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)
|
return fmt.Errorf("failed to write target: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -385,10 +390,30 @@ func readDatabaseForConvert(dbType, filePath, connString string) (*models.Databa
|
|||||||
return db, nil
|
return db, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func writeDatabase(db *models.Database, dbType, outputPath, packageName, schemaFilter string, flattenSchema bool, nullableTypes string, continueOnError bool) 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
|
var writer writers.Writer
|
||||||
|
|
||||||
writerOpts := newWriterOptions(outputPath, packageName, flattenSchema, nullableTypes, continueOnError)
|
writerOpts := newWriterOptions(outputPath, packageName, flattenSchema, nullableTypes, nullableArrays, continueOnError)
|
||||||
|
if extraFields != "" {
|
||||||
|
if !strings.EqualFold(dbType, "bun") {
|
||||||
|
return fmt.Errorf("--extra-fields is only supported for Bun output")
|
||||||
|
}
|
||||||
|
extraFieldsJSON, err := os.ReadFile(extraFields)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("failed to read --extra-fields file %q: %w", extraFields, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var parsed []wbun.ExtraFieldConfig
|
||||||
|
if err := stdjson.Unmarshal(extraFieldsJSON, &parsed); err != nil {
|
||||||
|
return fmt.Errorf("invalid --extra-fields JSON in %q: %w", extraFields, err)
|
||||||
|
}
|
||||||
|
if len(parsed) == 0 {
|
||||||
|
return fmt.Errorf("--extra-fields must contain at least one field")
|
||||||
|
}
|
||||||
|
writerOpts.Metadata = map[string]interface{}{
|
||||||
|
"extra_fields": string(extraFieldsJSON),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
switch strings.ToLower(dbType) {
|
switch strings.ToLower(dbType) {
|
||||||
case "dbml":
|
case "dbml":
|
||||||
|
|||||||
+13
-13
@@ -323,31 +323,31 @@ func writeDatabaseForEdit(dbType, filePath, connString string, db *models.Databa
|
|||||||
|
|
||||||
switch strings.ToLower(dbType) {
|
switch strings.ToLower(dbType) {
|
||||||
case "dbml":
|
case "dbml":
|
||||||
writer = wdbml.NewWriter(newWriterOptions(filePath, "", false, "", false))
|
writer = wdbml.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
|
||||||
case "dctx":
|
case "dctx":
|
||||||
writer = wdctx.NewWriter(newWriterOptions(filePath, "", false, "", false))
|
writer = wdctx.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
|
||||||
case "drawdb":
|
case "drawdb":
|
||||||
writer = wdrawdb.NewWriter(newWriterOptions(filePath, "", false, "", false))
|
writer = wdrawdb.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
|
||||||
case "graphql":
|
case "graphql":
|
||||||
writer = wgraphql.NewWriter(newWriterOptions(filePath, "", false, "", false))
|
writer = wgraphql.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
|
||||||
case "json":
|
case "json":
|
||||||
writer = wjson.NewWriter(newWriterOptions(filePath, "", false, "", false))
|
writer = wjson.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
|
||||||
case "yaml":
|
case "yaml":
|
||||||
writer = wyaml.NewWriter(newWriterOptions(filePath, "", false, "", false))
|
writer = wyaml.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
|
||||||
case "gorm":
|
case "gorm":
|
||||||
writer = wgorm.NewWriter(newWriterOptions(filePath, "", false, "", false))
|
writer = wgorm.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
|
||||||
case "bun":
|
case "bun":
|
||||||
writer = wbun.NewWriter(newWriterOptions(filePath, "", false, "", false))
|
writer = wbun.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
|
||||||
case "drizzle":
|
case "drizzle":
|
||||||
writer = wdrizzle.NewWriter(newWriterOptions(filePath, "", false, "", false))
|
writer = wdrizzle.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
|
||||||
case "prisma":
|
case "prisma":
|
||||||
writer = wprisma.NewWriter(newWriterOptions(filePath, "", false, "", false))
|
writer = wprisma.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
|
||||||
case "typeorm":
|
case "typeorm":
|
||||||
writer = wtypeorm.NewWriter(newWriterOptions(filePath, "", false, "", false))
|
writer = wtypeorm.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
|
||||||
case "sqlite", "sqlite3":
|
case "sqlite", "sqlite3":
|
||||||
writer = wsqlite.NewWriter(newWriterOptions(filePath, "", false, "", false))
|
writer = wsqlite.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
|
||||||
case "pgsql":
|
case "pgsql":
|
||||||
writer = wpgsql.NewWriter(newWriterOptions(filePath, "", false, "", false))
|
writer = wpgsql.NewWriter(newWriterOptions(filePath, "", false, "", "", false))
|
||||||
default:
|
default:
|
||||||
return fmt.Errorf("%s: unsupported format: %s", label, dbType)
|
return fmt.Errorf("%s: unsupported format: %s", label, dbType)
|
||||||
}
|
}
|
||||||
|
|||||||
+13
-13
@@ -375,61 +375,61 @@ func writeDatabaseForMerge(dbType, filePath, connString string, db *models.Datab
|
|||||||
if filePath == "" {
|
if filePath == "" {
|
||||||
return fmt.Errorf("%s: file path is required for DBML format", label)
|
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":
|
case "dctx":
|
||||||
if filePath == "" {
|
if filePath == "" {
|
||||||
return fmt.Errorf("%s: file path is required for DCTX format", label)
|
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":
|
case "drawdb":
|
||||||
if filePath == "" {
|
if filePath == "" {
|
||||||
return fmt.Errorf("%s: file path is required for DrawDB format", label)
|
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":
|
case "graphql":
|
||||||
if filePath == "" {
|
if filePath == "" {
|
||||||
return fmt.Errorf("%s: file path is required for GraphQL format", label)
|
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":
|
case "json":
|
||||||
if filePath == "" {
|
if filePath == "" {
|
||||||
return fmt.Errorf("%s: file path is required for JSON format", label)
|
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":
|
case "yaml":
|
||||||
if filePath == "" {
|
if filePath == "" {
|
||||||
return fmt.Errorf("%s: file path is required for YAML format", label)
|
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":
|
case "gorm":
|
||||||
if filePath == "" {
|
if filePath == "" {
|
||||||
return fmt.Errorf("%s: file path is required for GORM format", label)
|
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":
|
case "bun":
|
||||||
if filePath == "" {
|
if filePath == "" {
|
||||||
return fmt.Errorf("%s: file path is required for Bun format", label)
|
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":
|
case "drizzle":
|
||||||
if filePath == "" {
|
if filePath == "" {
|
||||||
return fmt.Errorf("%s: file path is required for Drizzle format", label)
|
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":
|
case "prisma":
|
||||||
if filePath == "" {
|
if filePath == "" {
|
||||||
return fmt.Errorf("%s: file path is required for Prisma format", label)
|
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":
|
case "typeorm":
|
||||||
if filePath == "" {
|
if filePath == "" {
|
||||||
return fmt.Errorf("%s: file path is required for TypeORM format", label)
|
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":
|
case "sqlite", "sqlite3":
|
||||||
writer = wsqlite.NewWriter(newWriterOptions(filePath, "", flattenSchema, "", false))
|
writer = wsqlite.NewWriter(newWriterOptions(filePath, "", flattenSchema, "", "", false))
|
||||||
case "pgsql":
|
case "pgsql":
|
||||||
writerOpts := newWriterOptions(filePath, "", flattenSchema, "", false)
|
writerOpts := newWriterOptions(filePath, "", flattenSchema, "", "", false)
|
||||||
if connString != "" {
|
if connString != "" {
|
||||||
writerOpts.Metadata = map[string]interface{}{
|
writerOpts.Metadata = map[string]interface{}{
|
||||||
"connection_string": connString,
|
"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{
|
return &writers.WriterOptions{
|
||||||
OutputPath: outputPath,
|
OutputPath: outputPath,
|
||||||
PackageName: packageName,
|
PackageName: packageName,
|
||||||
FlattenSchema: flattenSchema,
|
FlattenSchema: flattenSchema,
|
||||||
NullableTypes: nullableTypes,
|
NullableTypes: nullableTypes,
|
||||||
|
NullableArrays: nullableArrays,
|
||||||
Prisma7: prisma7,
|
Prisma7: prisma7,
|
||||||
ContinueOnError: continueOnError,
|
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
|
||||||
|
}
|
||||||
@@ -64,10 +64,12 @@ func init() {
|
|||||||
rootCmd.AddCommand(diffCmd)
|
rootCmd.AddCommand(diffCmd)
|
||||||
rootCmd.AddCommand(inspectCmd)
|
rootCmd.AddCommand(inspectCmd)
|
||||||
rootCmd.AddCommand(scriptsCmd)
|
rootCmd.AddCommand(scriptsCmd)
|
||||||
|
rootCmd.AddCommand(assetsCmd)
|
||||||
rootCmd.AddCommand(templCmd)
|
rootCmd.AddCommand(templCmd)
|
||||||
rootCmd.AddCommand(editCmd)
|
rootCmd.AddCommand(editCmd)
|
||||||
rootCmd.AddCommand(mergeCmd)
|
rootCmd.AddCommand(mergeCmd)
|
||||||
rootCmd.AddCommand(splitCmd)
|
rootCmd.AddCommand(splitCmd)
|
||||||
rootCmd.AddCommand(versionCmd)
|
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(&prisma7, "prisma7", false, "Use Prisma 7 generator conventions when reading/writing Prisma schemas")
|
||||||
}
|
}
|
||||||
|
|||||||
+17
-13
@@ -11,18 +11,19 @@ import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
var (
|
var (
|
||||||
splitSourceType string
|
splitSourceType string
|
||||||
splitSourcePath string
|
splitSourcePath string
|
||||||
splitSourceConn string
|
splitSourceConn string
|
||||||
splitTargetType string
|
splitTargetType string
|
||||||
splitTargetPath string
|
splitTargetPath string
|
||||||
splitSchemas string
|
splitSchemas string
|
||||||
splitTables string
|
splitTables string
|
||||||
splitPackageName string
|
splitPackageName string
|
||||||
splitDatabaseName string
|
splitDatabaseName string
|
||||||
splitExcludeSchema string
|
splitExcludeSchema string
|
||||||
splitExcludeTables string
|
splitExcludeTables string
|
||||||
splitNullableTypes string
|
splitNullableTypes string
|
||||||
|
splitNullableArrays string
|
||||||
)
|
)
|
||||||
|
|
||||||
var splitCmd = &cobra.Command{
|
var splitCmd = &cobra.Command{
|
||||||
@@ -111,7 +112,8 @@ func init() {
|
|||||||
splitCmd.Flags().StringVar(&splitTables, "tables", "", "Comma-separated list of table names to include (case-insensitive)")
|
splitCmd.Flags().StringVar(&splitTables, "tables", "", "Comma-separated list of table names to include (case-insensitive)")
|
||||||
splitCmd.Flags().StringVar(&splitExcludeSchema, "exclude-schema", "", "Comma-separated list of schema names to exclude")
|
splitCmd.Flags().StringVar(&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(&splitExcludeTables, "exclude-tables", "", "Comma-separated list of table names to exclude (case-insensitive)")
|
||||||
splitCmd.Flags().StringVar(&splitNullableTypes, "types", "", "Nullable type package for code-gen writers (bun/gorm): 'resolvespec' (default) or 'stdlib' (database/sql)")
|
splitCmd.Flags().StringVar(&splitNullableTypes, "types", "", "Nullable type package for code-gen writers (bun/gorm): 'baselib' (default, Go pointer types), 'stdlib' (database/sql), or 'sqltypes'")
|
||||||
|
splitCmd.Flags().StringVar(&splitNullableArrays, "array-nullable", "", "Nullable PostgreSQL array representation for the Bun writer in stdlib/baselib --types mode: 'slice' (default, plain slice) or 'pointer_slice' (*[]T, distinguishes NULL from '{}')")
|
||||||
|
|
||||||
err := splitCmd.MarkFlagRequired("from")
|
err := splitCmd.MarkFlagRequired("from")
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -188,7 +190,9 @@ func runSplit(cmd *cobra.Command, args []string) error {
|
|||||||
"", // no schema filter for split
|
"", // no schema filter for split
|
||||||
false, // no flatten-schema for split
|
false, // no flatten-schema for split
|
||||||
splitNullableTypes,
|
splitNullableTypes,
|
||||||
|
splitNullableArrays,
|
||||||
false, // no continue-on-error for split
|
false, // no continue-on-error for split
|
||||||
|
"", // no extra fields for split
|
||||||
)
|
)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("failed to write output: %w", err)
|
return fmt.Errorf("failed to write output: %w", err)
|
||||||
|
|||||||
@@ -85,6 +85,23 @@ migrations/
|
|||||||
|
|
||||||
All files will be found and executed in Priority→Sequence order regardless of directory structure.
|
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
|
## Commands
|
||||||
|
|
||||||
### relspec scripts list
|
### relspec scripts list
|
||||||
|
|||||||
@@ -3,17 +3,17 @@ package models_bun
|
|||||||
// //ModelCoreMasterprocess - Generated Table for Schema core
|
// //ModelCoreMasterprocess - Generated Table for Schema core
|
||||||
// type ModelCoreMasterprocess struct {
|
// type ModelCoreMasterprocess struct {
|
||||||
// bun.BaseModel `bun:"table:core.masterprocess,alias:masterprocess"`
|
// bun.BaseModel `bun:"table:core.masterprocess,alias:masterprocess"`
|
||||||
// Description resolvespec_common.SqlString `json:"description" bun:"description,type:citext,"`
|
// Description sql_types.SqlString `json:"description" bun:"description,type:citext,"`
|
||||||
// GUID resolvespec_common.SqlUUID `json:"guid" bun:"guid,type:uuid,default:newid(),"`
|
// GUID sql_types.SqlUUID `json:"guid" bun:"guid,type:uuid,default:newid(),"`
|
||||||
// Inactive resolvespec_common.SqlInt16 `json:"inactive" bun:"inactive,type:smallint,"`
|
// Inactive sql_types.SqlInt16 `json:"inactive" bun:"inactive,type:smallint,"`
|
||||||
// Jsonvalue resolvespec_common.SqlJSONB `json:"jsonvalue" bun:"jsonvalue,type:jsonb,"`
|
// Jsonvalue sql_types.SqlJSONB `json:"jsonvalue" bun:"jsonvalue,type:jsonb,"`
|
||||||
// Ridjsonschema resolvespec_common.SqlInt32 `json:"rid_jsonschema" bun:"rid_jsonschema,type:integer,"`
|
// Ridjsonschema sql_types.SqlInt32 `json:"rid_jsonschema" bun:"rid_jsonschema,type:integer,"`
|
||||||
// Ridmasterprocess resolvespec_common.SqlInt32 `json:"rid_masterprocess" bun:"rid_masterprocess,type:integer,pk,default:nextval('core.identity_masterprocess_rid_masterprocess'::regclass),"`
|
// Ridmasterprocess sql_types.SqlInt32 `json:"rid_masterprocess" bun:"rid_masterprocess,type:integer,pk,default:nextval('core.identity_masterprocess_rid_masterprocess'::regclass),"`
|
||||||
// Ridmastertypehubtype resolvespec_common.SqlInt32 `json:"rid_mastertype_hubtype" bun:"rid_mastertype_hubtype,type:integer,"`
|
// Ridmastertypehubtype sql_types.SqlInt32 `json:"rid_mastertype_hubtype" bun:"rid_mastertype_hubtype,type:integer,"`
|
||||||
// Ridmastertypeprocesstype resolvespec_common.SqlInt32 `json:"rid_mastertype_processtype" bun:"rid_mastertype_processtype,type:integer,"`
|
// Ridmastertypeprocesstype sql_types.SqlInt32 `json:"rid_mastertype_processtype" bun:"rid_mastertype_processtype,type:integer,"`
|
||||||
// Ridprogrammodule resolvespec_common.SqlInt32 `json:"rid_programmodule" bun:"rid_programmodule,type:integer,"`
|
// Ridprogrammodule sql_types.SqlInt32 `json:"rid_programmodule" bun:"rid_programmodule,type:integer,"`
|
||||||
// Sequenceno resolvespec_common.SqlInt32 `json:"sequenceno" bun:"sequenceno,type:integer,"`
|
// Sequenceno sql_types.SqlInt32 `json:"sequenceno" bun:"sequenceno,type:integer,"`
|
||||||
// Singleprocess resolvespec_common.SqlInt16 `json:"singleprocess" bun:"singleprocess,type:smallint,"`
|
// Singleprocess sql_types.SqlInt16 `json:"singleprocess" bun:"singleprocess,type:smallint,"`
|
||||||
// Updatecnt int64 `json:"updatecnt" bun:"updatecnt,type:integer,default:0,"`
|
// Updatecnt int64 `json:"updatecnt" bun:"updatecnt,type:integer,default:0,"`
|
||||||
// JSON *ModelCoreJsonschema `json:"JSON,omitempty" bun:"rel:has-one,join:rid_jsonschema=rid_jsonschema"`
|
// JSON *ModelCoreJsonschema `json:"JSON,omitempty" bun:"rel:has-one,join:rid_jsonschema=rid_jsonschema"`
|
||||||
// MTT_RID_MASTERTYPE_HUBTYPE *ModelCoreMastertype `json:"MTT_RID_MASTERTYPE_HUBTYPE,omitempty" bun:"rel:has-one,join:rid_mastertype_hubtype=rid_mastertype"`
|
// MTT_RID_MASTERTYPE_HUBTYPE *ModelCoreMastertype `json:"MTT_RID_MASTERTYPE_HUBTYPE,omitempty" bun:"rel:has-one,join:rid_mastertype_hubtype=rid_mastertype"`
|
||||||
|
|||||||
@@ -3,26 +3,26 @@ package models_bun
|
|||||||
// //ModelCoreMastertask - Generated Table for Schema core
|
// //ModelCoreMastertask - Generated Table for Schema core
|
||||||
// type ModelCoreMastertask struct {
|
// type ModelCoreMastertask struct {
|
||||||
// bun.BaseModel `bun:"table:core.mastertask,alias:mastertask"`
|
// bun.BaseModel `bun:"table:core.mastertask,alias:mastertask"`
|
||||||
// Allactionsmustcomplete resolvespec_common.SqlInt16 `json:"allactionsmustcomplete" bun:"allactionsmustcomplete,type:smallint,"`
|
// Allactionsmustcomplete sql_types.SqlInt16 `json:"allactionsmustcomplete" bun:"allactionsmustcomplete,type:smallint,"`
|
||||||
// Condition resolvespec_common.SqlString `json:"condition" bun:"condition,type:citext,"`
|
// Condition sql_types.SqlString `json:"condition" bun:"condition,type:citext,"`
|
||||||
// Description resolvespec_common.SqlString `json:"description" bun:"description,type:citext,"`
|
// Description sql_types.SqlString `json:"description" bun:"description,type:citext,"`
|
||||||
// Dueday resolvespec_common.SqlInt16 `json:"dueday" bun:"dueday,type:smallint,"`
|
// Dueday sql_types.SqlInt16 `json:"dueday" bun:"dueday,type:smallint,"`
|
||||||
// Dueoption resolvespec_common.SqlString `json:"dueoption" bun:"dueoption,type:citext,"`
|
// Dueoption sql_types.SqlString `json:"dueoption" bun:"dueoption,type:citext,"`
|
||||||
// Escalation resolvespec_common.SqlInt32 `json:"escalation" bun:"escalation,type:integer,"`
|
// Escalation sql_types.SqlInt32 `json:"escalation" bun:"escalation,type:integer,"`
|
||||||
// Escalationoption resolvespec_common.SqlString `json:"escalationoption" bun:"escalationoption,type:citext,"`
|
// Escalationoption sql_types.SqlString `json:"escalationoption" bun:"escalationoption,type:citext,"`
|
||||||
// GUID resolvespec_common.SqlUUID `json:"guid" bun:"guid,type:uuid,default:newid(),"`
|
// GUID sql_types.SqlUUID `json:"guid" bun:"guid,type:uuid,default:newid(),"`
|
||||||
// Inactive resolvespec_common.SqlInt16 `json:"inactive" bun:"inactive,type:smallint,"`
|
// Inactive sql_types.SqlInt16 `json:"inactive" bun:"inactive,type:smallint,"`
|
||||||
// Jsonvalue resolvespec_common.SqlJSONB `json:"jsonvalue" bun:"jsonvalue,type:jsonb,"`
|
// Jsonvalue sql_types.SqlJSONB `json:"jsonvalue" bun:"jsonvalue,type:jsonb,"`
|
||||||
// Mastertasknote resolvespec_common.SqlString `json:"mastertasknote" bun:"mastertasknote,type:citext,"`
|
// Mastertasknote sql_types.SqlString `json:"mastertasknote" bun:"mastertasknote,type:citext,"`
|
||||||
// Repeatinterval resolvespec_common.SqlInt16 `json:"repeatinterval" bun:"repeatinterval,type:smallint,"`
|
// Repeatinterval sql_types.SqlInt16 `json:"repeatinterval" bun:"repeatinterval,type:smallint,"`
|
||||||
// Repeattype resolvespec_common.SqlString `json:"repeattype" bun:"repeattype,type:citext,"`
|
// Repeattype sql_types.SqlString `json:"repeattype" bun:"repeattype,type:citext,"`
|
||||||
// Ridjsonschema resolvespec_common.SqlInt32 `json:"rid_jsonschema" bun:"rid_jsonschema,type:integer,"`
|
// Ridjsonschema sql_types.SqlInt32 `json:"rid_jsonschema" bun:"rid_jsonschema,type:integer,"`
|
||||||
// Ridmasterprocess resolvespec_common.SqlInt32 `json:"rid_masterprocess" bun:"rid_masterprocess,type:integer,"`
|
// Ridmasterprocess sql_types.SqlInt32 `json:"rid_masterprocess" bun:"rid_masterprocess,type:integer,"`
|
||||||
// Ridmastertask resolvespec_common.SqlInt32 `json:"rid_mastertask" bun:"rid_mastertask,type:integer,pk,default:nextval('core.identity_mastertask_rid_mastertask'::regclass),"`
|
// Ridmastertask sql_types.SqlInt32 `json:"rid_mastertask" bun:"rid_mastertask,type:integer,pk,default:nextval('core.identity_mastertask_rid_mastertask'::regclass),"`
|
||||||
// Ridmastertypetasktype resolvespec_common.SqlInt32 `json:"rid_mastertype_tasktype" bun:"rid_mastertype_tasktype,type:integer,"`
|
// Ridmastertypetasktype sql_types.SqlInt32 `json:"rid_mastertype_tasktype" bun:"rid_mastertype_tasktype,type:integer,"`
|
||||||
// Sequenceno resolvespec_common.SqlInt32 `json:"sequenceno" bun:"sequenceno,type:integer,"`
|
// Sequenceno sql_types.SqlInt32 `json:"sequenceno" bun:"sequenceno,type:integer,"`
|
||||||
// Singletask resolvespec_common.SqlInt16 `json:"singletask" bun:"singletask,type:smallint,"`
|
// Singletask sql_types.SqlInt16 `json:"singletask" bun:"singletask,type:smallint,"`
|
||||||
// Startday resolvespec_common.SqlInt16 `json:"startday" bun:"startday,type:smallint,"`
|
// Startday sql_types.SqlInt16 `json:"startday" bun:"startday,type:smallint,"`
|
||||||
// Updatecnt int64 `json:"updatecnt" bun:"updatecnt,type:integer,default:0,"`
|
// Updatecnt int64 `json:"updatecnt" bun:"updatecnt,type:integer,default:0,"`
|
||||||
// JSON *ModelCoreJsonschema `json:"JSON,omitempty" bun:"rel:has-one,join:rid_jsonschema=rid_jsonschema"`
|
// JSON *ModelCoreJsonschema `json:"JSON,omitempty" bun:"rel:has-one,join:rid_jsonschema=rid_jsonschema"`
|
||||||
// MPR *ModelCoreMasterprocess `json:"MPR,omitempty" bun:"rel:has-one,join:rid_masterprocess=rid_masterprocess"`
|
// MPR *ModelCoreMasterprocess `json:"MPR,omitempty" bun:"rel:has-one,join:rid_masterprocess=rid_masterprocess"`
|
||||||
|
|||||||
@@ -3,18 +3,18 @@ package models_bun
|
|||||||
// //ModelCoreMastertype - Generated Table for Schema core
|
// //ModelCoreMastertype - Generated Table for Schema core
|
||||||
// type ModelCoreMastertype struct {
|
// type ModelCoreMastertype struct {
|
||||||
// bun.BaseModel `bun:"table:core.mastertype,alias:mastertype"`
|
// bun.BaseModel `bun:"table:core.mastertype,alias:mastertype"`
|
||||||
// Category resolvespec_common.SqlString `json:"category" bun:"category,type:citext,"`
|
// Category sql_types.SqlString `json:"category" bun:"category,type:citext,"`
|
||||||
// Description resolvespec_common.SqlString `json:"description" bun:"description,type:citext,"`
|
// Description sql_types.SqlString `json:"description" bun:"description,type:citext,"`
|
||||||
// Disableedit resolvespec_common.SqlInt16 `json:"disableedit" bun:"disableedit,type:smallint,"`
|
// Disableedit sql_types.SqlInt16 `json:"disableedit" bun:"disableedit,type:smallint,"`
|
||||||
// Forprefix resolvespec_common.SqlString `json:"forprefix" bun:"forprefix,type:citext,"`
|
// Forprefix sql_types.SqlString `json:"forprefix" bun:"forprefix,type:citext,"`
|
||||||
// GUID resolvespec_common.SqlUUID `json:"guid" bun:"guid,type:uuid,default:newid(),"`
|
// GUID sql_types.SqlUUID `json:"guid" bun:"guid,type:uuid,default:newid(),"`
|
||||||
// Hidden resolvespec_common.SqlInt16 `json:"hidden" bun:"hidden,type:smallint,"`
|
// Hidden sql_types.SqlInt16 `json:"hidden" bun:"hidden,type:smallint,"`
|
||||||
// Inactive resolvespec_common.SqlInt16 `json:"inactive" bun:"inactive,type:smallint,"`
|
// Inactive sql_types.SqlInt16 `json:"inactive" bun:"inactive,type:smallint,"`
|
||||||
// Jsonvalue resolvespec_common.SqlJSONB `json:"jsonvalue" bun:"jsonvalue,type:jsonb,"`
|
// Jsonvalue sql_types.SqlJSONB `json:"jsonvalue" bun:"jsonvalue,type:jsonb,"`
|
||||||
// Mastertype resolvespec_common.SqlString `json:"mastertype" bun:"mastertype,type:citext,"`
|
// Mastertype sql_types.SqlString `json:"mastertype" bun:"mastertype,type:citext,"`
|
||||||
// Note resolvespec_common.SqlString `json:"note" bun:"note,type:citext,"`
|
// Note sql_types.SqlString `json:"note" bun:"note,type:citext,"`
|
||||||
// Ridmastertype resolvespec_common.SqlInt32 `json:"rid_mastertype" bun:"rid_mastertype,type:integer,pk,default:nextval('core.identity_mastertype_rid_mastertype'::regclass),"`
|
// Ridmastertype sql_types.SqlInt32 `json:"rid_mastertype" bun:"rid_mastertype,type:integer,pk,default:nextval('core.identity_mastertype_rid_mastertype'::regclass),"`
|
||||||
// Ridparent resolvespec_common.SqlInt32 `json:"rid_parent" bun:"rid_parent,type:integer,"`
|
// Ridparent sql_types.SqlInt32 `json:"rid_parent" bun:"rid_parent,type:integer,"`
|
||||||
// Updatecnt int64 `json:"updatecnt" bun:"updatecnt,type:integer,default:0,"`
|
// Updatecnt int64 `json:"updatecnt" bun:"updatecnt,type:integer,default:0,"`
|
||||||
// MTT *ModelCoreMastertype `json:"MTT,omitempty" bun:"rel:has-one,join:rid_mastertype=rid_parent"`
|
// MTT *ModelCoreMastertype `json:"MTT,omitempty" bun:"rel:has-one,join:rid_mastertype=rid_parent"`
|
||||||
|
|
||||||
|
|||||||
@@ -3,15 +3,15 @@ package models_bun
|
|||||||
// //ModelCoreProcess - Generated Table for Schema core
|
// //ModelCoreProcess - Generated Table for Schema core
|
||||||
// type ModelCoreProcess struct {
|
// type ModelCoreProcess struct {
|
||||||
// bun.BaseModel `bun:"table:core.process,alias:process"`
|
// bun.BaseModel `bun:"table:core.process,alias:process"`
|
||||||
// Completedate resolvespec_common.SqlDate `json:"completedate" bun:"completedate,type:date,"`
|
// Completedate sql_types.SqlDate `json:"completedate" bun:"completedate,type:date,"`
|
||||||
// Completetime types.CustomIntTime `json:"completetime" bun:"completetime,type:integer,"`
|
// Completetime types.CustomIntTime `json:"completetime" bun:"completetime,type:integer,"`
|
||||||
// Description resolvespec_common.SqlString `json:"description" bun:"description,type:citext,"`
|
// Description sql_types.SqlString `json:"description" bun:"description,type:citext,"`
|
||||||
// GUID resolvespec_common.SqlUUID `json:"guid" bun:"guid,type:uuid,default:newid(),"`
|
// GUID sql_types.SqlUUID `json:"guid" bun:"guid,type:uuid,default:newid(),"`
|
||||||
// Ridcompleteuser resolvespec_common.SqlInt32 `json:"rid_completeuser" bun:"rid_completeuser,type:integer,"`
|
// Ridcompleteuser sql_types.SqlInt32 `json:"rid_completeuser" bun:"rid_completeuser,type:integer,"`
|
||||||
// Ridhub resolvespec_common.SqlInt32 `json:"rid_hub" bun:"rid_hub,type:integer,"`
|
// Ridhub sql_types.SqlInt32 `json:"rid_hub" bun:"rid_hub,type:integer,"`
|
||||||
// Ridmasterprocess resolvespec_common.SqlInt32 `json:"rid_masterprocess" bun:"rid_masterprocess,type:integer,"`
|
// Ridmasterprocess sql_types.SqlInt32 `json:"rid_masterprocess" bun:"rid_masterprocess,type:integer,"`
|
||||||
// Ridprocess resolvespec_common.SqlInt32 `json:"rid_process" bun:"rid_process,type:integer,pk,default:nextval('core.identity_process_rid_process'::regclass),"`
|
// Ridprocess sql_types.SqlInt32 `json:"rid_process" bun:"rid_process,type:integer,pk,default:nextval('core.identity_process_rid_process'::regclass),"`
|
||||||
// Status resolvespec_common.SqlString `json:"status" bun:"status,type:citext,"`
|
// Status sql_types.SqlString `json:"status" bun:"status,type:citext,"`
|
||||||
// Updatecnt int64 `json:"updatecnt" bun:"updatecnt,type:integer,default:0,"`
|
// Updatecnt int64 `json:"updatecnt" bun:"updatecnt,type:integer,default:0,"`
|
||||||
// HUB *ModelCoreHub `json:"HUB,omitempty" bun:"rel:has-one,join:rid_hub=rid_hub"`
|
// HUB *ModelCoreHub `json:"HUB,omitempty" bun:"rel:has-one,join:rid_hub=rid_hub"`
|
||||||
// MPR *ModelCoreMasterprocess `json:"MPR,omitempty" bun:"rel:has-one,join:rid_masterprocess=rid_masterprocess"`
|
// MPR *ModelCoreMasterprocess `json:"MPR,omitempty" bun:"rel:has-one,join:rid_masterprocess=rid_masterprocess"`
|
||||||
|
|||||||
@@ -11,6 +11,7 @@ require (
|
|||||||
github.com/spf13/cobra v1.10.2
|
github.com/spf13/cobra v1.10.2
|
||||||
github.com/stretchr/testify v1.11.1
|
github.com/stretchr/testify v1.11.1
|
||||||
github.com/uptrace/bun v1.2.18
|
github.com/uptrace/bun v1.2.18
|
||||||
|
github.com/uptrace/bun/dialect/pgdialect v1.2.18
|
||||||
golang.org/x/text v0.37.0
|
golang.org/x/text v0.37.0
|
||||||
gopkg.in/yaml.v3 v3.0.1
|
gopkg.in/yaml.v3 v3.0.1
|
||||||
modernc.org/sqlite v1.50.1
|
modernc.org/sqlite v1.50.1
|
||||||
@@ -25,6 +26,7 @@ require (
|
|||||||
github.com/inconshreveable/mousetrap v1.1.0 // indirect
|
github.com/inconshreveable/mousetrap v1.1.0 // indirect
|
||||||
github.com/jackc/pgpassfile v1.0.0 // indirect
|
github.com/jackc/pgpassfile v1.0.0 // indirect
|
||||||
github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 // 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/jinzhu/inflection v1.0.0 // indirect
|
||||||
github.com/kr/pretty v0.3.1 // indirect
|
github.com/kr/pretty v0.3.1 // indirect
|
||||||
github.com/lucasb-eyer/go-colorful v1.4.0 // 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/msgpack/v5 v5.4.1 // indirect
|
||||||
github.com/vmihailenco/tagparser/v2 v2.0.0 // indirect
|
github.com/vmihailenco/tagparser/v2 v2.0.0 // indirect
|
||||||
golang.org/x/crypto v0.51.0 // indirect
|
golang.org/x/crypto v0.51.0 // indirect
|
||||||
|
golang.org/x/sync v0.20.0 // indirect
|
||||||
golang.org/x/sys v0.44.0 // indirect
|
golang.org/x/sys v0.44.0 // indirect
|
||||||
golang.org/x/term v0.43.0 // indirect
|
golang.org/x/term v0.43.0 // indirect
|
||||||
golang.org/x/tools v0.45.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/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 h1:3HnRcMfS6OBPMG1eSOzlbFJ/X/AyMEJb7rMxE6VQvDU=
|
||||||
github.com/uptrace/bun v1.2.18/go.mod h1:wNltaKJk4JtOt4SG5I5zmA7v0/Mzjh1+/S906Rayd3Y=
|
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 h1:cQriyiUvjTwOHg8QZaPihLWeRAAVoCpE00IUPn0Bjt8=
|
||||||
github.com/vmihailenco/msgpack/v5 v5.4.1/go.mod h1:GaZTsDaehaPpQVyxrf5mtQlH+pc21PIudVV/E3rRQok=
|
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=
|
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>
|
# Maintainer: Hein (Warky Devs) <hein@warky.dev>
|
||||||
pkgname=relspec
|
pkgname=relspec
|
||||||
pkgver=1.0.59
|
pkgver=1.0.66
|
||||||
pkgrel=1
|
pkgrel=1
|
||||||
pkgdesc="RelSpec is a comprehensive database relations management tool that reads, transforms, and writes database table specifications across multiple formats and ORMs."
|
pkgdesc="RelSpec is a comprehensive database relations management tool that reads, transforms, and writes database table specifications across multiple formats and ORMs."
|
||||||
arch=('x86_64' 'aarch64')
|
arch=('x86_64' 'aarch64')
|
||||||
|
|||||||
@@ -1,5 +1,5 @@
|
|||||||
Name: relspec
|
Name: relspec
|
||||||
Version: 1.0.59
|
Version: 1.0.66
|
||||||
Release: 1%{?dist}
|
Release: 1%{?dist}
|
||||||
Summary: RelSpec is a comprehensive database relations management tool that reads, transforms, and writes database table specifications across multiple formats and ORMs.
|
Summary: RelSpec is a comprehensive database relations management tool that reads, transforms, and writes database table specifications across multiple formats and ORMs.
|
||||||
|
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
+44
-16
@@ -2,10 +2,22 @@ package diff
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"reflect"
|
"reflect"
|
||||||
|
"sort"
|
||||||
|
|
||||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
"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
|
// CompareDatabases compares two database models and returns the differences
|
||||||
func CompareDatabases(source, target *models.Database) *DiffResult {
|
func CompareDatabases(source, target *models.Database) *DiffResult {
|
||||||
result := &DiffResult{
|
result := &DiffResult{
|
||||||
@@ -34,7 +46,8 @@ func compareSchemas(source, target []*models.Schema) *SchemaDiff {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Find missing and modified schemas
|
// Find missing and modified schemas
|
||||||
for name, srcSchema := range sourceMap {
|
for _, name := range sortedKeys(sourceMap) {
|
||||||
|
srcSchema := sourceMap[name]
|
||||||
if tgtSchema, exists := targetMap[name]; !exists {
|
if tgtSchema, exists := targetMap[name]; !exists {
|
||||||
diff.Missing = append(diff.Missing, srcSchema)
|
diff.Missing = append(diff.Missing, srcSchema)
|
||||||
} else {
|
} else {
|
||||||
@@ -45,7 +58,8 @@ func compareSchemas(source, target []*models.Schema) *SchemaDiff {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Find extra schemas
|
// Find extra schemas
|
||||||
for name, tgtSchema := range targetMap {
|
for _, name := range sortedKeys(targetMap) {
|
||||||
|
tgtSchema := targetMap[name]
|
||||||
if _, exists := sourceMap[name]; !exists {
|
if _, exists := sourceMap[name]; !exists {
|
||||||
diff.Extra = append(diff.Extra, tgtSchema)
|
diff.Extra = append(diff.Extra, tgtSchema)
|
||||||
}
|
}
|
||||||
@@ -106,7 +120,8 @@ func compareTables(source, target []*models.Table) *TableDiff {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Find missing and modified tables
|
// Find missing and modified tables
|
||||||
for name, srcTable := range sourceMap {
|
for _, name := range sortedKeys(sourceMap) {
|
||||||
|
srcTable := sourceMap[name]
|
||||||
if tgtTable, exists := targetMap[name]; !exists {
|
if tgtTable, exists := targetMap[name]; !exists {
|
||||||
diff.Missing = append(diff.Missing, srcTable)
|
diff.Missing = append(diff.Missing, srcTable)
|
||||||
} else {
|
} else {
|
||||||
@@ -117,7 +132,8 @@ func compareTables(source, target []*models.Table) *TableDiff {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Find extra tables
|
// Find extra tables
|
||||||
for name, tgtTable := range targetMap {
|
for _, name := range sortedKeys(targetMap) {
|
||||||
|
tgtTable := targetMap[name]
|
||||||
if _, exists := sourceMap[name]; !exists {
|
if _, exists := sourceMap[name]; !exists {
|
||||||
diff.Extra = append(diff.Extra, tgtTable)
|
diff.Extra = append(diff.Extra, tgtTable)
|
||||||
}
|
}
|
||||||
@@ -176,7 +192,8 @@ func compareColumns(source, target map[string]*models.Column) *ColumnDiff {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Find missing and modified columns
|
// Find missing and modified columns
|
||||||
for name, srcCol := range source {
|
for _, name := range sortedKeys(source) {
|
||||||
|
srcCol := source[name]
|
||||||
if tgtCol, exists := target[name]; !exists {
|
if tgtCol, exists := target[name]; !exists {
|
||||||
diff.Missing = append(diff.Missing, srcCol)
|
diff.Missing = append(diff.Missing, srcCol)
|
||||||
} else {
|
} else {
|
||||||
@@ -192,7 +209,8 @@ func compareColumns(source, target map[string]*models.Column) *ColumnDiff {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Find extra columns
|
// Find extra columns
|
||||||
for name, tgtCol := range target {
|
for _, name := range sortedKeys(target) {
|
||||||
|
tgtCol := target[name]
|
||||||
if _, exists := source[name]; !exists {
|
if _, exists := source[name]; !exists {
|
||||||
diff.Extra = append(diff.Extra, tgtCol)
|
diff.Extra = append(diff.Extra, tgtCol)
|
||||||
}
|
}
|
||||||
@@ -240,7 +258,8 @@ func compareIndexes(source, target map[string]*models.Index) *IndexDiff {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Find missing and modified indexes
|
// Find missing and modified indexes
|
||||||
for name, srcIdx := range source {
|
for _, name := range sortedKeys(source) {
|
||||||
|
srcIdx := source[name]
|
||||||
if tgtIdx, exists := target[name]; !exists {
|
if tgtIdx, exists := target[name]; !exists {
|
||||||
diff.Missing = append(diff.Missing, srcIdx)
|
diff.Missing = append(diff.Missing, srcIdx)
|
||||||
} else {
|
} else {
|
||||||
@@ -256,7 +275,8 @@ func compareIndexes(source, target map[string]*models.Index) *IndexDiff {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Find extra indexes
|
// Find extra indexes
|
||||||
for name, tgtIdx := range target {
|
for _, name := range sortedKeys(target) {
|
||||||
|
tgtIdx := target[name]
|
||||||
if _, exists := source[name]; !exists {
|
if _, exists := source[name]; !exists {
|
||||||
diff.Extra = append(diff.Extra, tgtIdx)
|
diff.Extra = append(diff.Extra, tgtIdx)
|
||||||
}
|
}
|
||||||
@@ -292,7 +312,8 @@ func compareConstraints(source, target map[string]*models.Constraint) *Constrain
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Find missing and modified constraints
|
// Find missing and modified constraints
|
||||||
for name, srcCon := range source {
|
for _, name := range sortedKeys(source) {
|
||||||
|
srcCon := source[name]
|
||||||
if tgtCon, exists := target[name]; !exists {
|
if tgtCon, exists := target[name]; !exists {
|
||||||
diff.Missing = append(diff.Missing, srcCon)
|
diff.Missing = append(diff.Missing, srcCon)
|
||||||
} else {
|
} else {
|
||||||
@@ -308,7 +329,8 @@ func compareConstraints(source, target map[string]*models.Constraint) *Constrain
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Find extra constraints
|
// Find extra constraints
|
||||||
for name, tgtCon := range target {
|
for _, name := range sortedKeys(target) {
|
||||||
|
tgtCon := target[name]
|
||||||
if _, exists := source[name]; !exists {
|
if _, exists := source[name]; !exists {
|
||||||
diff.Extra = append(diff.Extra, tgtCon)
|
diff.Extra = append(diff.Extra, tgtCon)
|
||||||
}
|
}
|
||||||
@@ -350,7 +372,8 @@ func compareRelationships(source, target map[string]*models.Relationship) *Relat
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Find missing and modified relationships
|
// Find missing and modified relationships
|
||||||
for name, srcRel := range source {
|
for _, name := range sortedKeys(source) {
|
||||||
|
srcRel := source[name]
|
||||||
if tgtRel, exists := target[name]; !exists {
|
if tgtRel, exists := target[name]; !exists {
|
||||||
diff.Missing = append(diff.Missing, srcRel)
|
diff.Missing = append(diff.Missing, srcRel)
|
||||||
} else {
|
} else {
|
||||||
@@ -366,7 +389,8 @@ func compareRelationships(source, target map[string]*models.Relationship) *Relat
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Find extra relationships
|
// Find extra relationships
|
||||||
for name, tgtRel := range target {
|
for _, name := range sortedKeys(target) {
|
||||||
|
tgtRel := target[name]
|
||||||
if _, exists := source[name]; !exists {
|
if _, exists := source[name]; !exists {
|
||||||
diff.Extra = append(diff.Extra, tgtRel)
|
diff.Extra = append(diff.Extra, tgtRel)
|
||||||
}
|
}
|
||||||
@@ -415,7 +439,8 @@ func compareViews(source, target []*models.View) *ViewDiff {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Find missing and modified views
|
// Find missing and modified views
|
||||||
for name, srcView := range sourceMap {
|
for _, name := range sortedKeys(sourceMap) {
|
||||||
|
srcView := sourceMap[name]
|
||||||
if tgtView, exists := targetMap[name]; !exists {
|
if tgtView, exists := targetMap[name]; !exists {
|
||||||
diff.Missing = append(diff.Missing, srcView)
|
diff.Missing = append(diff.Missing, srcView)
|
||||||
} else {
|
} else {
|
||||||
@@ -431,7 +456,8 @@ func compareViews(source, target []*models.View) *ViewDiff {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Find extra views
|
// Find extra views
|
||||||
for name, tgtView := range targetMap {
|
for _, name := range sortedKeys(targetMap) {
|
||||||
|
tgtView := targetMap[name]
|
||||||
if _, exists := sourceMap[name]; !exists {
|
if _, exists := sourceMap[name]; !exists {
|
||||||
diff.Extra = append(diff.Extra, tgtView)
|
diff.Extra = append(diff.Extra, tgtView)
|
||||||
}
|
}
|
||||||
@@ -468,7 +494,8 @@ func compareSequences(source, target []*models.Sequence) *SequenceDiff {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Find missing and modified sequences
|
// Find missing and modified sequences
|
||||||
for name, srcSeq := range sourceMap {
|
for _, name := range sortedKeys(sourceMap) {
|
||||||
|
srcSeq := sourceMap[name]
|
||||||
if tgtSeq, exists := targetMap[name]; !exists {
|
if tgtSeq, exists := targetMap[name]; !exists {
|
||||||
diff.Missing = append(diff.Missing, srcSeq)
|
diff.Missing = append(diff.Missing, srcSeq)
|
||||||
} else {
|
} else {
|
||||||
@@ -484,7 +511,8 @@ func compareSequences(source, target []*models.Sequence) *SequenceDiff {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Find extra sequences
|
// Find extra sequences
|
||||||
for name, tgtSeq := range targetMap {
|
for _, name := range sortedKeys(targetMap) {
|
||||||
|
tgtSeq := targetMap[name]
|
||||||
if _, exists := sourceMap[name]; !exists {
|
if _, exists := sourceMap[name]; !exists {
|
||||||
diff.Extra = append(diff.Extra, tgtSeq)
|
diff.Extra = append(diff.Extra, tgtSeq)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,6 +1,7 @@
|
|||||||
package diff
|
package diff
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"reflect"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
"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) {
|
func TestCompareColumnDetails(t *testing.T) {
|
||||||
tests := []struct {
|
tests := []struct {
|
||||||
name string
|
name string
|
||||||
|
|||||||
@@ -2,6 +2,7 @@ package inspector
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"sort"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
"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) {
|
func (i *Inspector) Inspect() (*InspectorReport, error) {
|
||||||
results := []ValidationResult{}
|
results := []ValidationResult{}
|
||||||
|
|
||||||
// Run all enabled validators
|
// Run all enabled validators in deterministic (alphabetical) rule-name order
|
||||||
for ruleName, rule := range i.config.Rules {
|
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() {
|
if !rule.IsEnabled() {
|
||||||
continue
|
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) {
|
func TestInspectWithDisabledRules(t *testing.T) {
|
||||||
db := createTestDatabase()
|
db := createTestDatabase()
|
||||||
config := GetDefaultConfig()
|
config := GetDefaultConfig()
|
||||||
|
|||||||
@@ -5,6 +5,7 @@ import (
|
|||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
"io"
|
||||||
"os"
|
"os"
|
||||||
|
"sort"
|
||||||
"strings"
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
)
|
)
|
||||||
@@ -199,12 +200,18 @@ func (f *MarkdownFormatter) formatContext(context map[string]interface{}) string
|
|||||||
"column": true,
|
"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] {
|
if skipKeys[key] {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
parts = append(parts, fmt.Sprintf("%s=%v", key, value))
|
parts = append(parts, fmt.Sprintf("%s=%v", key, context[key]))
|
||||||
}
|
}
|
||||||
|
|
||||||
return strings.Join(parts, ", ")
|
return strings.Join(parts, ", ")
|
||||||
|
|||||||
+54
-12
@@ -2,12 +2,54 @@ package inspector
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"regexp"
|
"regexp"
|
||||||
|
"sort"
|
||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||||
"git.warky.dev/wdevs/relspecgo/pkg/pgsql"
|
"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
|
// validatePrimaryKeyNaming checks that primary key column names match a pattern
|
||||||
func validatePrimaryKeyNaming(db *models.Database, rule Rule, ruleName string) []ValidationResult {
|
func validatePrimaryKeyNaming(db *models.Database, rule Rule, ruleName string) []ValidationResult {
|
||||||
results := []ValidationResult{}
|
results := []ValidationResult{}
|
||||||
@@ -18,7 +60,7 @@ func validatePrimaryKeyNaming(db *models.Database, rule Rule, ruleName string) [
|
|||||||
|
|
||||||
for _, schema := range db.Schemas {
|
for _, schema := range db.Schemas {
|
||||||
for _, table := range schema.Tables {
|
for _, table := range schema.Tables {
|
||||||
for _, col := range table.Columns {
|
for _, col := range sortColumns(table.Columns) {
|
||||||
if col.IsPrimaryKey {
|
if col.IsPrimaryKey {
|
||||||
location := formatLocation(schema.Name, table.Name, col.Name)
|
location := formatLocation(schema.Name, table.Name, col.Name)
|
||||||
passed := pattern.MatchString(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 _, schema := range db.Schemas {
|
||||||
for _, table := range schema.Tables {
|
for _, table := range schema.Tables {
|
||||||
for _, col := range table.Columns {
|
for _, col := range sortColumns(table.Columns) {
|
||||||
if col.IsPrimaryKey {
|
if col.IsPrimaryKey {
|
||||||
location := formatLocation(schema.Name, table.Name, col.Name)
|
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 _, schema := range db.Schemas {
|
||||||
for _, table := range schema.Tables {
|
for _, table := range schema.Tables {
|
||||||
for _, col := range table.Columns {
|
for _, col := range sortColumns(table.Columns) {
|
||||||
if col.IsPrimaryKey {
|
if col.IsPrimaryKey {
|
||||||
location := formatLocation(schema.Name, table.Name, col.Name)
|
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 _, schema := range db.Schemas {
|
||||||
for _, table := range schema.Tables {
|
for _, table := range schema.Tables {
|
||||||
// Check foreign key constraints
|
// Check foreign key constraints
|
||||||
for _, constraint := range table.Constraints {
|
for _, constraint := range sortConstraints(table.Constraints) {
|
||||||
if constraint.Type == models.ForeignKeyConstraint {
|
if constraint.Type == models.ForeignKeyConstraint {
|
||||||
for _, colName := range constraint.Columns {
|
for _, colName := range constraint.Columns {
|
||||||
location := formatLocation(schema.Name, table.Name, colName)
|
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 _, schema := range db.Schemas {
|
||||||
for _, table := range schema.Tables {
|
for _, table := range schema.Tables {
|
||||||
for _, constraint := range table.Constraints {
|
for _, constraint := range sortConstraints(table.Constraints) {
|
||||||
if constraint.Type == models.ForeignKeyConstraint {
|
if constraint.Type == models.ForeignKeyConstraint {
|
||||||
location := formatLocation(schema.Name, table.Name, "")
|
location := formatLocation(schema.Name, table.Name, "")
|
||||||
passed := pattern.MatchString(constraint.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
|
// Check if each FK column has an index
|
||||||
for fkCol := range fkColumns {
|
for _, fkCol := range sortedKeys(fkColumns) {
|
||||||
hasIndex := false
|
hasIndex := false
|
||||||
|
|
||||||
// Check table indexes
|
// Check table indexes
|
||||||
@@ -282,7 +324,7 @@ func validateColumnNamingCase(db *models.Database, rule Rule, ruleName string) [
|
|||||||
|
|
||||||
for _, schema := range db.Schemas {
|
for _, schema := range db.Schemas {
|
||||||
for _, table := range schema.Tables {
|
for _, table := range schema.Tables {
|
||||||
for _, col := range table.Columns {
|
for _, col := range sortColumns(table.Columns) {
|
||||||
location := formatLocation(schema.Name, table.Name, col.Name)
|
location := formatLocation(schema.Name, table.Name, col.Name)
|
||||||
passed := pattern.MatchString(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 _, schema := range db.Schemas {
|
||||||
for _, table := range schema.Tables {
|
for _, table := range schema.Tables {
|
||||||
for _, col := range table.Columns {
|
for _, col := range sortColumns(table.Columns) {
|
||||||
location := formatLocation(schema.Name, table.Name, col.Name)
|
location := formatLocation(schema.Name, table.Name, col.Name)
|
||||||
passed := len(col.Name) <= rule.MaxLength
|
passed := len(col.Name) <= rule.MaxLength
|
||||||
|
|
||||||
@@ -396,7 +438,7 @@ func validateReservedKeywords(db *models.Database, rule Rule, ruleName string) [
|
|||||||
|
|
||||||
// Check column names
|
// Check column names
|
||||||
if rule.CheckColumns {
|
if rule.CheckColumns {
|
||||||
for _, col := range table.Columns {
|
for _, col := range sortColumns(table.Columns) {
|
||||||
location := formatLocation(schema.Name, table.Name, col.Name)
|
location := formatLocation(schema.Name, table.Name, col.Name)
|
||||||
passed := !keywords[strings.ToUpper(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
|
// Check all foreign key constraints
|
||||||
for _, schema := range db.Schemas {
|
for _, schema := range db.Schemas {
|
||||||
for _, table := range schema.Tables {
|
for _, table := range schema.Tables {
|
||||||
for _, constraint := range table.Constraints {
|
for _, constraint := range sortConstraints(table.Constraints) {
|
||||||
if constraint.Type == models.ForeignKeyConstraint {
|
if constraint.Type == models.ForeignKeyConstraint {
|
||||||
// Build referenced table key
|
// Build referenced table key
|
||||||
refSchema := constraint.ReferencedSchema
|
refSchema := constraint.ReferencedSchema
|
||||||
@@ -522,7 +564,7 @@ func validateCircularDependency(db *models.Database, rule Rule, ruleName string)
|
|||||||
for _, table := range schema.Tables {
|
for _, table := range schema.Tables {
|
||||||
tableKey := schema.Name + "." + table.Name
|
tableKey := schema.Name + "." + table.Name
|
||||||
|
|
||||||
for _, constraint := range table.Constraints {
|
for _, constraint := range sortConstraints(table.Constraints) {
|
||||||
if constraint.Type == models.ForeignKeyConstraint {
|
if constraint.Type == models.ForeignKeyConstraint {
|
||||||
refSchema := constraint.ReferencedSchema
|
refSchema := constraint.ReferencedSchema
|
||||||
if refSchema == "" {
|
if refSchema == "" {
|
||||||
@@ -537,7 +579,7 @@ func validateCircularDependency(db *models.Database, rule Rule, ruleName string)
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Check for cycles using DFS
|
// Check for cycles using DFS
|
||||||
for tableKey := range dependencies {
|
for _, tableKey := range sortedKeys(dependencies) {
|
||||||
visited := make(map[string]bool)
|
visited := make(map[string]bool)
|
||||||
recStack := make(map[string]bool)
|
recStack := make(map[string]bool)
|
||||||
|
|
||||||
|
|||||||
+12
-2
@@ -5,6 +5,7 @@ package merge
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"sort"
|
||||||
"strconv"
|
"strconv"
|
||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
@@ -156,8 +157,17 @@ func (r *MergeResult) mergeColumns(table *models.Table, srcTable *models.Table)
|
|||||||
existingColumns[colName] = table.Columns[colName]
|
existingColumns[colName] = table.Columns[colName]
|
||||||
}
|
}
|
||||||
|
|
||||||
// Merge columns
|
// Merge columns in deterministic (alphabetical) order so that, when a
|
||||||
for colName, srcCol := range srcTable.Columns {
|
// 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 {
|
if tgtCol, exists := existingColumns[colName]; !exists {
|
||||||
// Column doesn't exist, add it
|
// Column doesn't exist, add it
|
||||||
newCol := cloneColumn(srcCol)
|
newCol := cloneColumn(srcCol)
|
||||||
|
|||||||
+23
-1
@@ -1,6 +1,9 @@
|
|||||||
package models
|
package models
|
||||||
|
|
||||||
import "fmt"
|
import (
|
||||||
|
"fmt"
|
||||||
|
"sort"
|
||||||
|
)
|
||||||
|
|
||||||
// Flat/Denormalized Views
|
// 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
|
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
|
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
|
return flatRelationships
|
||||||
}
|
}
|
||||||
|
|||||||
+36
-14
@@ -5,6 +5,7 @@
|
|||||||
package models
|
package models
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"sort"
|
||||||
"strings"
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
@@ -141,15 +142,28 @@ func (d *Table) SQLName() string {
|
|||||||
|
|
||||||
// GetPrimaryKey returns the primary key column for the table, or nil if none exists.
|
// GetPrimaryKey returns the primary key column for the table, or nil if none exists.
|
||||||
func (m Table) GetPrimaryKey() *Column {
|
func (m Table) GetPrimaryKey() *Column {
|
||||||
|
var pk *Column
|
||||||
for _, column := range m.Columns {
|
for _, column := range m.Columns {
|
||||||
if column.IsPrimaryKey {
|
if !column.IsPrimaryKey {
|
||||||
return column
|
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 {
|
func (m Table) GetForeignKeys() []*Constraint {
|
||||||
keys := make([]*Constraint, 0)
|
keys := make([]*Constraint, 0)
|
||||||
|
|
||||||
@@ -158,6 +172,12 @@ func (m Table) GetForeignKeys() []*Constraint {
|
|||||||
keys = append(keys, c)
|
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
|
return keys
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -350,16 +370,17 @@ const (
|
|||||||
// Script represents a database migration or initialization script.
|
// Script represents a database migration or initialization script.
|
||||||
// Scripts can have dependencies and rollback capabilities.
|
// Scripts can have dependencies and rollback capabilities.
|
||||||
type Script struct {
|
type Script struct {
|
||||||
Name string `json:"name" yaml:"name" xml:"name"`
|
Name string `json:"name" yaml:"name" xml:"name"`
|
||||||
Description string `json:"description" yaml:"description" xml:"description"`
|
Description string `json:"description" yaml:"description" xml:"description"`
|
||||||
SQL string `json:"sql" yaml:"sql" xml:"sql"`
|
SQL string `json:"sql" yaml:"sql" xml:"sql"`
|
||||||
Rollback string `json:"rollback,omitempty" yaml:"rollback,omitempty" xml:"rollback,omitempty"`
|
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"`
|
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"`
|
Schema string `json:"schema,omitempty" yaml:"schema,omitempty" xml:"schema,omitempty"`
|
||||||
Version string `json:"version,omitempty" yaml:"version,omitempty" xml:"version,omitempty"`
|
Version string `json:"version,omitempty" yaml:"version,omitempty" xml:"version,omitempty"`
|
||||||
Priority int `json:"priority,omitempty" yaml:"priority,omitempty" xml:"priority,omitempty"`
|
Priority int `json:"priority,omitempty" yaml:"priority,omitempty" xml:"priority,omitempty"`
|
||||||
Sequence uint `json:"sequence,omitempty" yaml:"sequence,omitempty" xml:"sequence,omitempty"`
|
Sequence uint `json:"sequence,omitempty" yaml:"sequence,omitempty" xml:"sequence,omitempty"`
|
||||||
GUID string `json:"guid" yaml:"guid" xml:"guid"`
|
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.
|
// SQLName returns the script name in lowercase for SQL compatibility.
|
||||||
@@ -468,6 +489,7 @@ func InitScript(name string) *Script {
|
|||||||
return &Script{
|
return &Script{
|
||||||
Name: name,
|
Name: name,
|
||||||
RunAfter: make([]string, 0),
|
RunAfter: make([]string, 0),
|
||||||
|
Metadata: make(map[string]any),
|
||||||
GUID: uuid.New().String(),
|
GUID: uuid.New().String(),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -746,7 +746,7 @@ func (r *Reader) goTypeToSQL(expr ast.Expr) string {
|
|||||||
if t.Sel.Name == "Time" {
|
if t.Sel.Name == "Time" {
|
||||||
return "timestamp"
|
return "timestamp"
|
||||||
}
|
}
|
||||||
case "resolvespec_common", "sql_types":
|
case "sql_types":
|
||||||
return r.sqlTypeToSQL(t.Sel.Name)
|
return r.sqlTypeToSQL(t.Sel.Name)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -787,7 +787,7 @@ func (r *Reader) isNullableGoType(expr ast.Expr) bool {
|
|||||||
case *ast.SelectorExpr:
|
case *ast.SelectorExpr:
|
||||||
// Check for sql_types nullable types
|
// Check for sql_types nullable types
|
||||||
if ident, ok := t.X.(*ast.Ident); ok {
|
if ident, ok := t.X.(*ast.Ident); ok {
|
||||||
if ident.Name == "resolvespec_common" || ident.Name == "sql_types" {
|
if ident.Name == "sql_types" {
|
||||||
return strings.HasPrefix(t.Sel.Name, "Sql")
|
return strings.HasPrefix(t.Sel.Name, "Sql")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -4,6 +4,7 @@ import (
|
|||||||
"encoding/xml"
|
"encoding/xml"
|
||||||
"fmt"
|
"fmt"
|
||||||
"os"
|
"os"
|
||||||
|
"sort"
|
||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
"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 len(columns) == 0 {
|
||||||
if dctxKey.Primary {
|
if dctxKey.Primary {
|
||||||
// Look for common primary key column patterns
|
// Look for common primary key column patterns
|
||||||
|
colNames := make([]string, 0, len(table.Columns))
|
||||||
for colName := range table.Columns {
|
for colName := range table.Columns {
|
||||||
|
colNames = append(colNames, colName)
|
||||||
|
}
|
||||||
|
sort.Strings(colNames)
|
||||||
|
|
||||||
|
for _, colName := range colNames {
|
||||||
colNameLower := strings.ToLower(colName)
|
colNameLower := strings.ToLower(colName)
|
||||||
if strings.HasPrefix(colNameLower, "rid_") || strings.HasSuffix(colNameLower, "id") {
|
if strings.HasPrefix(colNameLower, "rid_") || strings.HasSuffix(colNameLower, "id") {
|
||||||
columns = append(columns, colName)
|
columns = append(columns, colName)
|
||||||
|
|||||||
@@ -820,17 +820,31 @@ func (r *Reader) createImplicitJoinTable(model1, model2 string, tableMap map[str
|
|||||||
tableMap[joinTableName] = joinTable
|
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 {
|
func (r *Reader) getPrimaryKeyColumn(table *models.Table) *models.Column {
|
||||||
if table == nil {
|
if table == nil {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
var pk *models.Column
|
||||||
for _, col := range table.Columns {
|
for _, col := range table.Columns {
|
||||||
if col.IsPrimaryKey {
|
if !col.IsPrimaryKey {
|
||||||
return col
|
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
|
- `1_001_test.txt` - Wrong extension
|
||||||
- `readme.md` - Not a SQL file
|
- `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
|
## Usage
|
||||||
|
|
||||||
### Basic Usage
|
### Basic Usage
|
||||||
|
|||||||
@@ -7,6 +7,7 @@ import (
|
|||||||
"regexp"
|
"regexp"
|
||||||
"strconv"
|
"strconv"
|
||||||
|
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/assetloader"
|
||||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||||
"git.warky.dev/wdevs/relspecgo/pkg/readers"
|
"git.warky.dev/wdevs/relspecgo/pkg/readers"
|
||||||
)
|
)
|
||||||
@@ -151,6 +152,10 @@ func (r *Reader) readScripts() ([]*models.Script, error) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("failed to read file %s: %w", path, err)
|
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
|
// Get relative path from base directory
|
||||||
relPath, err := filepath.Rel(r.options.FilePath, path)
|
relPath, err := filepath.Rel(r.options.FilePath, path)
|
||||||
@@ -161,9 +166,10 @@ func (r *Reader) readScripts() ([]*models.Script, error) {
|
|||||||
// Create Script model
|
// Create Script model
|
||||||
script := models.InitScript(name)
|
script := models.InitScript(name)
|
||||||
script.Description = fmt.Sprintf("SQL script from %s", relPath)
|
script.Description = fmt.Sprintf("SQL script from %s", relPath)
|
||||||
script.SQL = string(content)
|
script.SQL = sql
|
||||||
script.Priority = priority
|
script.Priority = priority
|
||||||
script.Sequence = uint(sequence)
|
script.Sequence = uint(sequence)
|
||||||
|
script.Metadata[assetloader.ScriptSourcePathMetadataKey] = path
|
||||||
|
|
||||||
scripts = append(scripts, script)
|
scripts = append(scripts, script)
|
||||||
|
|
||||||
|
|||||||
@@ -1,8 +1,10 @@
|
|||||||
package sqldir
|
package sqldir
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"encoding/base64"
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"git.warky.dev/wdevs/relspecgo/pkg/readers"
|
"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
|
// Create test SQL files with both underscore and hyphen separators
|
||||||
testFiles := map[string]string{
|
testFiles := map[string]string{
|
||||||
"1_001_create_users.sql": "CREATE TABLE users (id SERIAL PRIMARY KEY, name TEXT);",
|
"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);",
|
"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);",
|
"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_003_seed_data.pgsql": "INSERT INTO users (name) VALUES ('Alice'), ('Bob');",
|
||||||
"10-10-create-newid.pgsql": "CREATE TABLE newid (id SERIAL PRIMARY KEY);",
|
"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 {
|
for filename, content := range testFiles {
|
||||||
@@ -267,10 +269,10 @@ func TestReader_HyphenFormat(t *testing.T) {
|
|||||||
|
|
||||||
// Create test files with hyphen separators
|
// Create test files with hyphen separators
|
||||||
testFiles := map[string]string{
|
testFiles := map[string]string{
|
||||||
"1-001-create-table.sql": "CREATE TABLE test (id INT);",
|
"1-001-create-table.sql": "CREATE TABLE test (id INT);",
|
||||||
"1-002-insert-data.pgsql": "INSERT INTO test VALUES (1);",
|
"1-002-insert-data.pgsql": "INSERT INTO test VALUES (1);",
|
||||||
"10-10-create-newid.pgsql": "CREATE TABLE newid (id SERIAL);",
|
"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 {
|
for filename, content := range testFiles {
|
||||||
@@ -301,10 +303,10 @@ func TestReader_HyphenFormat(t *testing.T) {
|
|||||||
priority int
|
priority int
|
||||||
sequence uint
|
sequence uint
|
||||||
}{
|
}{
|
||||||
"create-table": {1, 1},
|
"create-table": {1, 1},
|
||||||
"insert-data": {1, 2},
|
"insert-data": {1, 2},
|
||||||
"add-index": {2, 5},
|
"add-index": {2, 5},
|
||||||
"create-newid": {10, 10},
|
"create-newid": {10, 10},
|
||||||
}
|
}
|
||||||
|
|
||||||
for _, script := range schema.Scripts {
|
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")
|
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
|
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 {
|
func (r *Reader) getPrimaryKeyColumn(table *models.Table) *models.Column {
|
||||||
if table == nil {
|
if table == nil {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
var pk *models.Column
|
||||||
for _, col := range table.Columns {
|
for _, col := range table.Columns {
|
||||||
if col.IsPrimaryKey {
|
if !col.IsPrimaryKey {
|
||||||
return col
|
continue
|
||||||
|
}
|
||||||
|
if pk == nil {
|
||||||
|
pk = col
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if col.Sequence > 0 && pk.Sequence > 0 {
|
||||||
|
if col.Sequence < pk.Sequence {
|
||||||
|
pk = col
|
||||||
|
}
|
||||||
|
} else if col.Name < pk.Name {
|
||||||
|
pk = col
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil
|
return pk
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,14 +1,16 @@
|
|||||||
package reflectutil
|
package reflectutil
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"fmt"
|
||||||
"reflect"
|
"reflect"
|
||||||
|
"sort"
|
||||||
"strings"
|
"strings"
|
||||||
)
|
)
|
||||||
|
|
||||||
// Deref dereferences pointers until it reaches a non-pointer value
|
// Deref dereferences pointers until it reaches a non-pointer value
|
||||||
// Returns the dereferenced value and true if successful, or the original value and false if nil
|
// Returns the dereferenced value and true if successful, or the original value and false if nil
|
||||||
func Deref(v reflect.Value) (reflect.Value, bool) {
|
func Deref(v reflect.Value) (reflect.Value, bool) {
|
||||||
for v.Kind() == reflect.Ptr {
|
for v.Kind() == reflect.Pointer {
|
||||||
if v.IsNil() {
|
if v.IsNil() {
|
||||||
return v, false
|
return v, false
|
||||||
}
|
}
|
||||||
@@ -134,7 +136,7 @@ func MapKeys(i interface{}) []interface{} {
|
|||||||
return []interface{}{}
|
return []interface{}{}
|
||||||
}
|
}
|
||||||
|
|
||||||
keys := v.MapKeys()
|
keys := sortedMapKeys(v)
|
||||||
result := make([]interface{}, len(keys))
|
result := make([]interface{}, len(keys))
|
||||||
for i, key := range keys {
|
for i, key := range keys {
|
||||||
result[i] = key.Interface()
|
result[i] = key.Interface()
|
||||||
@@ -155,14 +157,39 @@ func MapValues(i interface{}) []interface{} {
|
|||||||
return []interface{}{}
|
return []interface{}{}
|
||||||
}
|
}
|
||||||
|
|
||||||
result := make([]interface{}, 0, v.Len())
|
keys := sortedMapKeys(v)
|
||||||
iter := v.MapRange()
|
result := make([]interface{}, 0, len(keys))
|
||||||
for iter.Next() {
|
for _, key := range keys {
|
||||||
result = append(result, iter.Value().Interface())
|
result = append(result, v.MapIndex(key).Interface())
|
||||||
}
|
}
|
||||||
return result
|
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
|
// MapGet safely gets a value from a map by key
|
||||||
// Returns nil if key doesn't exist or not a map
|
// Returns nil if key doesn't exist or not a map
|
||||||
func MapGet(m interface{}, key interface{}) interface{} {
|
func MapGet(m interface{}, key interface{}) interface{} {
|
||||||
|
|||||||
@@ -0,0 +1,136 @@
|
|||||||
|
# sqltypes
|
||||||
|
|
||||||
|
Nullable SQL types for hand-written or generated Go models. Each type wraps a
|
||||||
|
value with a `Valid` flag and implements `database/sql.Scanner`,
|
||||||
|
`driver.Valuer`, `encoding/json`, `gopkg.in/yaml.v3`, and `encoding/xml`
|
||||||
|
marshalling — so a single struct field can be scanned from a database row,
|
||||||
|
round-tripped through JSON/YAML/XML, and written back to the database without
|
||||||
|
any per-format glue code.
|
||||||
|
|
||||||
|
This package is what the `bun` and `gorm` writers emit when generating models
|
||||||
|
with `--types sqltypes` (see [`pkg/writers/bun`](../writers/bun/README.md) and
|
||||||
|
[`pkg/writers/gorm`](../writers/gorm/README.md)). It can also be imported
|
||||||
|
directly in hand-written models.
|
||||||
|
|
||||||
|
## Import
|
||||||
|
|
||||||
|
```go
|
||||||
|
import sql_types "git.warky.dev/wdevs/relspecgo/pkg/sqltypes"
|
||||||
|
```
|
||||||
|
|
||||||
|
## Scalar types
|
||||||
|
|
||||||
|
All scalar types are instantiations of the generic `SqlNull[T]`:
|
||||||
|
|
||||||
|
| Type | Underlying | Typical SQL type |
|
||||||
|
|---|---|---|
|
||||||
|
| `SqlInt16` | `int16` | `smallint` |
|
||||||
|
| `SqlInt32` | `int32` | `integer` |
|
||||||
|
| `SqlInt64` | `int64` | `bigint` |
|
||||||
|
| `SqlFloat32` | `float32` | `real`, `float4` |
|
||||||
|
| `SqlFloat64` | `float64` | `double precision`, `numeric`, `decimal`, `money` |
|
||||||
|
| `SqlBool` | `bool` | `boolean` |
|
||||||
|
| `SqlString` | `string` | `text`, `varchar`, `char`, `citext`, `inet`, `cidr`, `macaddr` |
|
||||||
|
| `SqlByteArray` | `[]byte` | `bytea` (base64-encoded in JSON/YAML/XML) |
|
||||||
|
| `SqlUUID` | `uuid.UUID` (`github.com/google/uuid`) | `uuid` |
|
||||||
|
|
||||||
|
You can also instantiate `SqlNull[T]` directly for any type not covered
|
||||||
|
above, e.g. `SqlNull[MyEnum]`.
|
||||||
|
|
||||||
|
### Date/time types
|
||||||
|
|
||||||
|
Plain `time.Time` doesn't distinguish date-only, time-only, and timestamp
|
||||||
|
semantics, and its zero value marshals to a confusing `0001-01-01T00:00:00Z`.
|
||||||
|
These wrapper types fix both problems:
|
||||||
|
|
||||||
|
| Type | Format | Notes |
|
||||||
|
|---|---|---|
|
||||||
|
| `SqlTimeStamp` | `2006-01-02T15:04:05` | Full timestamp |
|
||||||
|
| `SqlDate` | `2006-01-02` | Date only |
|
||||||
|
| `SqlTime` | `15:04:05` | Time only |
|
||||||
|
|
||||||
|
Zero/pre-epoch values (`time.Time{}` or anything before `0002-01-01`) marshal
|
||||||
|
to `null` and `Value()` returns `nil`, instead of leaking Go's zero-time
|
||||||
|
sentinel into the database or API responses.
|
||||||
|
|
||||||
|
### JSON types
|
||||||
|
|
||||||
|
| Type | Underlying | Notes |
|
||||||
|
|---|---|---|
|
||||||
|
| `SqlJSONB` | `[]byte` | Raw JSON bytes; `MarshalYAML` decodes to native YAML mappings/sequences instead of an embedded JSON string |
|
||||||
|
| `SqlJSON` | `= SqlJSONB` | Alias — PostgreSQL's `json` and `jsonb` share the same Go representation |
|
||||||
|
|
||||||
|
`SqlJSONB` has `AsMap()` / `AsSlice()` helpers for pulling out
|
||||||
|
`map[string]any` / `[]any` without a separate `json.Unmarshal` call.
|
||||||
|
|
||||||
|
### Vector type (pgvector)
|
||||||
|
|
||||||
|
`SqlVector` wraps `[]float32` for the `vector` column type ([pgvector](https://github.com/pgvector/pgvector)),
|
||||||
|
scanning/writing the `[1,2,3]` literal format pgvector uses over the wire.
|
||||||
|
|
||||||
|
## Array types
|
||||||
|
|
||||||
|
PostgreSQL array columns (`text[]`, `integer[]`, …) map to `SqlXxxArray`
|
||||||
|
types, each wrapping `Val []T` + `Valid bool` and handling PostgreSQL's
|
||||||
|
`{a,b,c}` array literal format on `Scan`/`Value`:
|
||||||
|
|
||||||
|
`SqlStringArray`, `SqlInt16Array`, `SqlInt32Array`, `SqlInt64Array`,
|
||||||
|
`SqlFloat32Array`, `SqlFloat64Array`, `SqlBoolArray`, `SqlUUIDArray`.
|
||||||
|
|
||||||
|
## Constructing values
|
||||||
|
|
||||||
|
Every type has a `NewSqlXxx(v)` constructor that sets `Valid: true`:
|
||||||
|
|
||||||
|
```go
|
||||||
|
name := sql_types.NewSqlString("Ada Lovelace")
|
||||||
|
age := sql_types.NewSqlInt32(36)
|
||||||
|
tags := sql_types.NewSqlStringArray([]string{"engineer", "mathematician"})
|
||||||
|
```
|
||||||
|
|
||||||
|
The zero value of any type (`sql_types.SqlString{}`) is null/invalid — use it
|
||||||
|
directly for a `NULL` field instead of a separate constructor.
|
||||||
|
|
||||||
|
Generic helpers:
|
||||||
|
|
||||||
|
```go
|
||||||
|
sql_types.Null(v, valid) // SqlNull[T]{Val: v, Valid: valid}
|
||||||
|
sql_types.NewSql[T](anyValue) // best-effort conversion from any Go value
|
||||||
|
```
|
||||||
|
|
||||||
|
## Reading values back
|
||||||
|
|
||||||
|
Each scalar type has typed accessors that return the zero value instead of
|
||||||
|
panicking when `Valid` is false:
|
||||||
|
|
||||||
|
```go
|
||||||
|
n.Int64() // SqlInt16/32/64, SqlFloat32/64, SqlBool, SqlString → int64
|
||||||
|
n.Float64() // → float64
|
||||||
|
n.Bool() // → bool
|
||||||
|
n.Time() // SqlNull[time.Time]-based types → time.Time
|
||||||
|
n.UUID() // SqlUUID → uuid.UUID
|
||||||
|
n.String() // fmt.Stringer — empty string when invalid
|
||||||
|
```
|
||||||
|
|
||||||
|
## Example
|
||||||
|
|
||||||
|
```go
|
||||||
|
type User struct {
|
||||||
|
ID sql_types.SqlUUID `json:"id"`
|
||||||
|
Name sql_types.SqlString `json:"name"`
|
||||||
|
Tags sql_types.SqlStringArray `json:"tags"`
|
||||||
|
Metadata sql_types.SqlJSONB `json:"metadata"`
|
||||||
|
CreatedAt sql_types.SqlTimeStamp `json:"created_at"`
|
||||||
|
}
|
||||||
|
|
||||||
|
u := User{
|
||||||
|
ID: sql_types.NewSqlUUID(uuid.New()),
|
||||||
|
Name: sql_types.NewSqlString("Ada Lovelace"),
|
||||||
|
Tags: sql_types.NewSqlStringArray([]string{"engineer"}),
|
||||||
|
CreatedAt: sql_types.SqlTimeStampNow(),
|
||||||
|
}
|
||||||
|
// Metadata left as the zero value → serializes as null, scans as NULL.
|
||||||
|
```
|
||||||
|
|
||||||
|
Every type implements `sql.Scanner` and `driver.Valuer`, so these fields can
|
||||||
|
be used directly as struct fields with `database/sql`, `bun`, or `gorm`
|
||||||
|
without additional tags or hooks.
|
||||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,485 @@
|
|||||||
|
package sqltypes
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/google/uuid"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestParsePostgresArrayElements(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
input string
|
||||||
|
want []string
|
||||||
|
wantErr bool
|
||||||
|
}{
|
||||||
|
{"simple", "{a,b,c}", []string{"a", "b", "c"}, false},
|
||||||
|
{"empty array", "{}", []string{}, false},
|
||||||
|
{"null", "NULL", nil, false},
|
||||||
|
{"lowercase null", "null", nil, false},
|
||||||
|
{"empty string", "", nil, false},
|
||||||
|
{"quoted with comma", `{a,"b,c",d}`, []string{"a", "b,c", "d"}, false},
|
||||||
|
{"escaped quote", `{"a""b"}`, []string{`a"b`}, false},
|
||||||
|
{"escaped backslash", `{"a\\b"}`, []string{`a\b`}, false},
|
||||||
|
{"not an array", "abc", nil, true},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
got, err := parsePostgresArrayElements(tt.input)
|
||||||
|
if tt.wantErr {
|
||||||
|
if err == nil {
|
||||||
|
t.Fatalf("expected error, got nil")
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("unexpected error: %v", err)
|
||||||
|
}
|
||||||
|
if len(got) != len(tt.want) {
|
||||||
|
t.Fatalf("expected %v, got %v", tt.want, got)
|
||||||
|
}
|
||||||
|
for i := range got {
|
||||||
|
if got[i] != tt.want[i] {
|
||||||
|
t.Errorf("index %d: expected %q, got %q", i, tt.want[i], got[i])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFormatPostgresStringArray(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
input []string
|
||||||
|
want string
|
||||||
|
}{
|
||||||
|
{"nil", nil, "NULL"},
|
||||||
|
{"empty", []string{}, "{}"},
|
||||||
|
{"simple", []string{"a", "b"}, "{a,b}"},
|
||||||
|
{"needs quoting comma", []string{"a,b"}, `{"a,b"}`},
|
||||||
|
{"needs quoting empty elem", []string{""}, `{""}`},
|
||||||
|
{"needs quoting quote", []string{`a"b`}, `{"a""b"}`},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
got := formatPostgresStringArray(tt.input)
|
||||||
|
if got != tt.want {
|
||||||
|
t.Errorf("expected %q, got %q", tt.want, got)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSqlStringArray(t *testing.T) {
|
||||||
|
t.Run("scan and value round-trip", func(t *testing.T) {
|
||||||
|
var a SqlStringArray
|
||||||
|
if err := a.Scan(`{a,"b,c",d}`); err != nil {
|
||||||
|
t.Fatalf("Scan failed: %v", err)
|
||||||
|
}
|
||||||
|
want := []string{"a", "b,c", "d"}
|
||||||
|
if len(a.Val) != len(want) {
|
||||||
|
t.Fatalf("expected %v, got %v", want, a.Val)
|
||||||
|
}
|
||||||
|
for i := range want {
|
||||||
|
if a.Val[i] != want[i] {
|
||||||
|
t.Errorf("index %d: expected %q, got %q", i, want[i], a.Val[i])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
val, err := a.Value()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Value failed: %v", err)
|
||||||
|
}
|
||||||
|
var b SqlStringArray
|
||||||
|
if err := b.Scan(val); err != nil {
|
||||||
|
t.Fatalf("re-scan failed: %v", err)
|
||||||
|
}
|
||||||
|
for i := range want {
|
||||||
|
if b.Val[i] != want[i] {
|
||||||
|
t.Errorf("round-trip index %d: expected %q, got %q", i, want[i], b.Val[i])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("scan nil", func(t *testing.T) {
|
||||||
|
var a SqlStringArray
|
||||||
|
if err := a.Scan(nil); err != nil {
|
||||||
|
t.Fatalf("Scan failed: %v", err)
|
||||||
|
}
|
||||||
|
if a.Valid {
|
||||||
|
t.Error("expected invalid")
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("value invalid", func(t *testing.T) {
|
||||||
|
a := SqlStringArray{Valid: false}
|
||||||
|
val, err := a.Value()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Value failed: %v", err)
|
||||||
|
}
|
||||||
|
if val != nil {
|
||||||
|
t.Errorf("expected nil, got %v", val)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("scan wrong type", func(t *testing.T) {
|
||||||
|
var a SqlStringArray
|
||||||
|
if err := a.Scan(42); err == nil {
|
||||||
|
t.Error("expected error for unsupported scan type")
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("json round-trip", func(t *testing.T) {
|
||||||
|
a := NewSqlStringArray([]string{"x", "y", "z"})
|
||||||
|
data, err := json.Marshal(a)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Marshal failed: %v", err)
|
||||||
|
}
|
||||||
|
if string(data) != `["x","y","z"]` {
|
||||||
|
t.Errorf("unexpected JSON: %s", data)
|
||||||
|
}
|
||||||
|
var a2 SqlStringArray
|
||||||
|
if err := json.Unmarshal(data, &a2); err != nil {
|
||||||
|
t.Fatalf("Unmarshal failed: %v", err)
|
||||||
|
}
|
||||||
|
if !a2.Valid || len(a2.Val) != 3 {
|
||||||
|
t.Fatalf("expected 3 valid elements, got %v", a2)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("json null", func(t *testing.T) {
|
||||||
|
var a SqlStringArray
|
||||||
|
data, err := json.Marshal(a)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Marshal failed: %v", err)
|
||||||
|
}
|
||||||
|
if string(data) != "null" {
|
||||||
|
t.Errorf("expected null, got %s", data)
|
||||||
|
}
|
||||||
|
var a2 SqlStringArray
|
||||||
|
if err := json.Unmarshal([]byte("null"), &a2); err != nil {
|
||||||
|
t.Fatalf("Unmarshal failed: %v", err)
|
||||||
|
}
|
||||||
|
if a2.Valid {
|
||||||
|
t.Error("expected invalid after unmarshaling null")
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("json unmarshal invalid type errors", func(t *testing.T) {
|
||||||
|
var a SqlStringArray
|
||||||
|
if err := json.Unmarshal([]byte(`42`), &a); err == nil {
|
||||||
|
t.Error("expected error unmarshaling non-array JSON")
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSqlInt16Array(t *testing.T) {
|
||||||
|
t.Run("scan and value", func(t *testing.T) {
|
||||||
|
var a SqlInt16Array
|
||||||
|
if err := a.Scan("{1,2,-3}"); err != nil {
|
||||||
|
t.Fatalf("Scan failed: %v", err)
|
||||||
|
}
|
||||||
|
want := []int16{1, 2, -3}
|
||||||
|
for i := range want {
|
||||||
|
if a.Val[i] != want[i] {
|
||||||
|
t.Errorf("index %d: expected %d, got %d", i, want[i], a.Val[i])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
val, err := a.Value()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Value failed: %v", err)
|
||||||
|
}
|
||||||
|
if val != "{1,2,-3}" {
|
||||||
|
t.Errorf("expected {1,2,-3}, got %v", val)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("scan invalid element", func(t *testing.T) {
|
||||||
|
var a SqlInt16Array
|
||||||
|
if err := a.Scan("{1,abc}"); err == nil {
|
||||||
|
t.Error("expected error for non-numeric element")
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("json round-trip", func(t *testing.T) {
|
||||||
|
a := NewSqlInt16Array([]int16{5, 10, 15})
|
||||||
|
data, err := json.Marshal(a)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Marshal failed: %v", err)
|
||||||
|
}
|
||||||
|
var a2 SqlInt16Array
|
||||||
|
if err := json.Unmarshal(data, &a2); err != nil {
|
||||||
|
t.Fatalf("Unmarshal failed: %v", err)
|
||||||
|
}
|
||||||
|
for i, v := range a.Val {
|
||||||
|
if a2.Val[i] != v {
|
||||||
|
t.Errorf("index %d: expected %d, got %d", i, v, a2.Val[i])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSqlInt32Array(t *testing.T) {
|
||||||
|
a := NewSqlInt32Array([]int32{100000, -200000})
|
||||||
|
val, err := a.Value()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Value failed: %v", err)
|
||||||
|
}
|
||||||
|
var b SqlInt32Array
|
||||||
|
if err := b.Scan(val); err != nil {
|
||||||
|
t.Fatalf("Scan failed: %v", err)
|
||||||
|
}
|
||||||
|
for i, v := range a.Val {
|
||||||
|
if b.Val[i] != v {
|
||||||
|
t.Errorf("index %d: expected %d, got %d", i, v, b.Val[i])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
data, err := json.Marshal(a)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Marshal failed: %v", err)
|
||||||
|
}
|
||||||
|
var c SqlInt32Array
|
||||||
|
if err := json.Unmarshal(data, &c); err != nil {
|
||||||
|
t.Fatalf("Unmarshal failed: %v", err)
|
||||||
|
}
|
||||||
|
for i, v := range a.Val {
|
||||||
|
if c.Val[i] != v {
|
||||||
|
t.Errorf("index %d: expected %d, got %d", i, v, c.Val[i])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSqlInt64Array(t *testing.T) {
|
||||||
|
a := NewSqlInt64Array([]int64{9223372036854775807, -9223372036854775808})
|
||||||
|
val, err := a.Value()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Value failed: %v", err)
|
||||||
|
}
|
||||||
|
var b SqlInt64Array
|
||||||
|
if err := b.Scan(val); err != nil {
|
||||||
|
t.Fatalf("Scan failed: %v", err)
|
||||||
|
}
|
||||||
|
for i, v := range a.Val {
|
||||||
|
if b.Val[i] != v {
|
||||||
|
t.Errorf("index %d: expected %d, got %d", i, v, b.Val[i])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
t.Run("null json", func(t *testing.T) {
|
||||||
|
var n SqlInt64Array
|
||||||
|
data, _ := json.Marshal(n)
|
||||||
|
if string(data) != "null" {
|
||||||
|
t.Errorf("expected null, got %s", data)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSqlFloat32Array(t *testing.T) {
|
||||||
|
a := NewSqlFloat32Array([]float32{1.5, -2.25, 0})
|
||||||
|
val, err := a.Value()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Value failed: %v", err)
|
||||||
|
}
|
||||||
|
var b SqlFloat32Array
|
||||||
|
if err := b.Scan(val); err != nil {
|
||||||
|
t.Fatalf("Scan failed: %v", err)
|
||||||
|
}
|
||||||
|
for i, v := range a.Val {
|
||||||
|
if b.Val[i] != v {
|
||||||
|
t.Errorf("index %d: expected %v, got %v", i, v, b.Val[i])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
data, err := json.Marshal(a)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Marshal failed: %v", err)
|
||||||
|
}
|
||||||
|
var c SqlFloat32Array
|
||||||
|
if err := json.Unmarshal(data, &c); err != nil {
|
||||||
|
t.Fatalf("Unmarshal failed: %v", err)
|
||||||
|
}
|
||||||
|
for i, v := range a.Val {
|
||||||
|
if c.Val[i] != v {
|
||||||
|
t.Errorf("index %d: expected %v, got %v", i, v, c.Val[i])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSqlFloat64Array(t *testing.T) {
|
||||||
|
a := NewSqlFloat64Array([]float64{3.14159, -2.71828})
|
||||||
|
val, err := a.Value()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Value failed: %v", err)
|
||||||
|
}
|
||||||
|
var b SqlFloat64Array
|
||||||
|
if err := b.Scan(val); err != nil {
|
||||||
|
t.Fatalf("Scan failed: %v", err)
|
||||||
|
}
|
||||||
|
for i, v := range a.Val {
|
||||||
|
if b.Val[i] != v {
|
||||||
|
t.Errorf("index %d: expected %v, got %v", i, v, b.Val[i])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSqlBoolArray(t *testing.T) {
|
||||||
|
t.Run("scan various truthy forms", func(t *testing.T) {
|
||||||
|
var a SqlBoolArray
|
||||||
|
if err := a.Scan("{t,f,true,false,1,0,yes}"); err != nil {
|
||||||
|
t.Fatalf("Scan failed: %v", err)
|
||||||
|
}
|
||||||
|
want := []bool{true, false, true, false, true, false, true}
|
||||||
|
for i := range want {
|
||||||
|
if a.Val[i] != want[i] {
|
||||||
|
t.Errorf("index %d: expected %v, got %v", i, want[i], a.Val[i])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("value formatting", func(t *testing.T) {
|
||||||
|
a := NewSqlBoolArray([]bool{true, false})
|
||||||
|
val, err := a.Value()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Value failed: %v", err)
|
||||||
|
}
|
||||||
|
if val != "{t,f}" {
|
||||||
|
t.Errorf("expected {t,f}, got %v", val)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("json round-trip", func(t *testing.T) {
|
||||||
|
a := NewSqlBoolArray([]bool{true, false, true})
|
||||||
|
data, err := json.Marshal(a)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Marshal failed: %v", err)
|
||||||
|
}
|
||||||
|
var a2 SqlBoolArray
|
||||||
|
if err := json.Unmarshal(data, &a2); err != nil {
|
||||||
|
t.Fatalf("Unmarshal failed: %v", err)
|
||||||
|
}
|
||||||
|
for i, v := range a.Val {
|
||||||
|
if a2.Val[i] != v {
|
||||||
|
t.Errorf("index %d: expected %v, got %v", i, v, a2.Val[i])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSqlUUIDArray(t *testing.T) {
|
||||||
|
u1, u2 := uuid.New(), uuid.New()
|
||||||
|
|
||||||
|
t.Run("scan and value round-trip", func(t *testing.T) {
|
||||||
|
a := NewSqlUUIDArray([]uuid.UUID{u1, u2})
|
||||||
|
val, err := a.Value()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Value failed: %v", err)
|
||||||
|
}
|
||||||
|
var b SqlUUIDArray
|
||||||
|
if err := b.Scan(val); err != nil {
|
||||||
|
t.Fatalf("Scan failed: %v", err)
|
||||||
|
}
|
||||||
|
if b.Val[0] != u1 || b.Val[1] != u2 {
|
||||||
|
t.Errorf("expected [%v %v], got %v", u1, u2, b.Val)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("scan invalid uuid element", func(t *testing.T) {
|
||||||
|
var a SqlUUIDArray
|
||||||
|
if err := a.Scan("{not-a-uuid}"); err == nil {
|
||||||
|
t.Error("expected error for invalid uuid element")
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("json round-trip", func(t *testing.T) {
|
||||||
|
a := NewSqlUUIDArray([]uuid.UUID{u1, u2})
|
||||||
|
data, err := json.Marshal(a)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Marshal failed: %v", err)
|
||||||
|
}
|
||||||
|
var a2 SqlUUIDArray
|
||||||
|
if err := json.Unmarshal(data, &a2); err != nil {
|
||||||
|
t.Fatalf("Unmarshal failed: %v", err)
|
||||||
|
}
|
||||||
|
if a2.Val[0] != u1 || a2.Val[1] != u2 {
|
||||||
|
t.Errorf("expected [%v %v], got %v", u1, u2, a2.Val)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSqlVector(t *testing.T) {
|
||||||
|
t.Run("scan and value round-trip", func(t *testing.T) {
|
||||||
|
var v SqlVector
|
||||||
|
if err := v.Scan("[1,2.5,-3]"); err != nil {
|
||||||
|
t.Fatalf("Scan failed: %v", err)
|
||||||
|
}
|
||||||
|
want := []float32{1, 2.5, -3}
|
||||||
|
for i := range want {
|
||||||
|
if v.Val[i] != want[i] {
|
||||||
|
t.Errorf("index %d: expected %v, got %v", i, want[i], v.Val[i])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
val, err := v.Value()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Value failed: %v", err)
|
||||||
|
}
|
||||||
|
if val != "[1,2.5,-3]" {
|
||||||
|
t.Errorf("expected [1,2.5,-3], got %v", val)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("scan empty vector", func(t *testing.T) {
|
||||||
|
var v SqlVector
|
||||||
|
if err := v.Scan("[]"); err != nil {
|
||||||
|
t.Fatalf("Scan failed: %v", err)
|
||||||
|
}
|
||||||
|
if !v.Valid || len(v.Val) != 0 {
|
||||||
|
t.Errorf("expected valid empty vector, got %v", v)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("scan invalid literal", func(t *testing.T) {
|
||||||
|
var v SqlVector
|
||||||
|
if err := v.Scan("not-a-vector"); err == nil {
|
||||||
|
t.Error("expected error for invalid vector literal")
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("json round-trip", func(t *testing.T) {
|
||||||
|
v := NewSqlVector([]float32{0.1, 0.2, 0.3})
|
||||||
|
data, err := json.Marshal(v)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Marshal failed: %v", err)
|
||||||
|
}
|
||||||
|
var v2 SqlVector
|
||||||
|
if err := json.Unmarshal(data, &v2); err != nil {
|
||||||
|
t.Fatalf("Unmarshal failed: %v", err)
|
||||||
|
}
|
||||||
|
for i, f := range v.Val {
|
||||||
|
if v2.Val[i] != f {
|
||||||
|
t.Errorf("index %d: expected %v, got %v", i, f, v2.Val[i])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("json null", func(t *testing.T) {
|
||||||
|
var v SqlVector
|
||||||
|
data, err := json.Marshal(v)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Marshal failed: %v", err)
|
||||||
|
}
|
||||||
|
if string(data) != "null" {
|
||||||
|
t.Errorf("expected null, got %s", data)
|
||||||
|
}
|
||||||
|
var v2 SqlVector
|
||||||
|
if err := json.Unmarshal([]byte("null"), &v2); err != nil {
|
||||||
|
t.Fatalf("Unmarshal failed: %v", err)
|
||||||
|
}
|
||||||
|
if v2.Valid {
|
||||||
|
t.Error("expected invalid after unmarshaling null")
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
@@ -0,0 +1,981 @@
|
|||||||
|
// Package sqltypes provides nullable SQL types with automatic casting and conversion methods.
|
||||||
|
package sqltypes
|
||||||
|
|
||||||
|
import (
|
||||||
|
"database/sql"
|
||||||
|
"database/sql/driver"
|
||||||
|
"encoding/base64"
|
||||||
|
"encoding/json"
|
||||||
|
"encoding/xml"
|
||||||
|
"fmt"
|
||||||
|
"reflect"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/google/uuid"
|
||||||
|
"gopkg.in/yaml.v3"
|
||||||
|
)
|
||||||
|
|
||||||
|
// tryParseDT attempts to parse a string into a time.Time using various formats.
|
||||||
|
func tryParseDT(str string) (time.Time, error) {
|
||||||
|
var lasterror error
|
||||||
|
tryFormats := []string{
|
||||||
|
time.RFC3339,
|
||||||
|
"2006-01-02T15:04:05.000-0700",
|
||||||
|
"2006-01-02T15:04:05.000",
|
||||||
|
"06-01-02T15:04:05.000",
|
||||||
|
"2006-01-02T15:04:05",
|
||||||
|
"2006-01-02 15:04:05",
|
||||||
|
"02/01/2006",
|
||||||
|
"02-01-2006",
|
||||||
|
"2006-01-02",
|
||||||
|
"15:04:05.000",
|
||||||
|
"15:04:05",
|
||||||
|
"15:04",
|
||||||
|
}
|
||||||
|
for _, f := range tryFormats {
|
||||||
|
tx, err := time.Parse(f, str)
|
||||||
|
if err == nil {
|
||||||
|
return tx, nil
|
||||||
|
}
|
||||||
|
lasterror = err
|
||||||
|
}
|
||||||
|
return time.Time{}, lasterror // Return zero time on failure
|
||||||
|
}
|
||||||
|
|
||||||
|
// ToJSONDT formats a time.Time to RFC3339 string.
|
||||||
|
func ToJSONDT(dt time.Time) string {
|
||||||
|
return dt.Format(time.RFC3339)
|
||||||
|
}
|
||||||
|
|
||||||
|
// SqlNull is a generic nullable type that behaves like sql.NullXXX with auto-casting.
|
||||||
|
type SqlNull[T any] struct {
|
||||||
|
Val T
|
||||||
|
Valid bool
|
||||||
|
}
|
||||||
|
|
||||||
|
// Scan implements sql.Scanner.
|
||||||
|
func (n *SqlNull[T]) Scan(value any) error {
|
||||||
|
if value == nil {
|
||||||
|
n.Valid = false
|
||||||
|
n.Val = *new(T)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check if T is []byte, and decode base64 if applicable
|
||||||
|
// Do this BEFORE trying sql.Null to ensure base64 is handled
|
||||||
|
var zero T
|
||||||
|
if _, ok := any(zero).([]byte); ok {
|
||||||
|
// For []byte types, try to decode from base64
|
||||||
|
var strVal string
|
||||||
|
switch v := value.(type) {
|
||||||
|
case string:
|
||||||
|
strVal = v
|
||||||
|
case []byte:
|
||||||
|
strVal = string(v)
|
||||||
|
default:
|
||||||
|
strVal = fmt.Sprintf("%v", value)
|
||||||
|
}
|
||||||
|
// Try base64 decode
|
||||||
|
if decoded, err := base64.StdEncoding.DecodeString(strVal); err == nil {
|
||||||
|
n.Val = any(decoded).(T)
|
||||||
|
n.Valid = true
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
// Fallback to raw bytes
|
||||||
|
n.Val = any([]byte(strVal)).(T)
|
||||||
|
n.Valid = true
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Try standard sql.Null[T] for other types.
|
||||||
|
var sqlNull sql.Null[T]
|
||||||
|
if err := sqlNull.Scan(value); err == nil {
|
||||||
|
n.Val = sqlNull.V
|
||||||
|
n.Valid = sqlNull.Valid
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Fallback: parse from string/bytes.
|
||||||
|
switch v := value.(type) {
|
||||||
|
case string:
|
||||||
|
return n.FromString(v)
|
||||||
|
case []byte:
|
||||||
|
return n.FromString(string(v))
|
||||||
|
case float32, float64:
|
||||||
|
return n.FromString(fmt.Sprintf("%f", value))
|
||||||
|
case int, int8, int16, int32, int64, uint, uint8, uint16, uint32, uint64:
|
||||||
|
return n.FromString(fmt.Sprintf("%d", value))
|
||||||
|
default:
|
||||||
|
return n.FromString(fmt.Sprintf("%v", value))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
func (n *SqlNull[T]) FromString(s string) error {
|
||||||
|
s = strings.TrimSpace(s)
|
||||||
|
n.Valid = false
|
||||||
|
n.Val = *new(T)
|
||||||
|
|
||||||
|
if s == "" || strings.EqualFold(s, "null") {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
var zero T
|
||||||
|
switch any(zero).(type) {
|
||||||
|
case int, int8, int16, int32, int64:
|
||||||
|
if i, err := strconv.ParseInt(s, 10, 64); err == nil {
|
||||||
|
reflect.ValueOf(&n.Val).Elem().SetInt(i)
|
||||||
|
n.Valid = true
|
||||||
|
} else if f, err := strconv.ParseFloat(s, 64); err == nil {
|
||||||
|
reflect.ValueOf(&n.Val).Elem().SetInt(int64(f))
|
||||||
|
n.Valid = true
|
||||||
|
}
|
||||||
|
case uint, uint8, uint16, uint32, uint64:
|
||||||
|
if u, err := strconv.ParseUint(s, 10, 64); err == nil {
|
||||||
|
reflect.ValueOf(&n.Val).Elem().SetUint(u)
|
||||||
|
n.Valid = true
|
||||||
|
} else if f, err := strconv.ParseFloat(s, 64); err == nil && f >= 0 {
|
||||||
|
reflect.ValueOf(&n.Val).Elem().SetUint(uint64(f))
|
||||||
|
n.Valid = true
|
||||||
|
}
|
||||||
|
case float32, float64:
|
||||||
|
if f, err := strconv.ParseFloat(s, 64); err == nil {
|
||||||
|
reflect.ValueOf(&n.Val).Elem().SetFloat(f)
|
||||||
|
n.Valid = true
|
||||||
|
}
|
||||||
|
case bool:
|
||||||
|
if b, err := strconv.ParseBool(s); err == nil {
|
||||||
|
n.Val = any(b).(T)
|
||||||
|
n.Valid = true
|
||||||
|
}
|
||||||
|
case time.Time:
|
||||||
|
if t, err := tryParseDT(s); err == nil && !t.IsZero() {
|
||||||
|
n.Val = any(t).(T)
|
||||||
|
n.Valid = true
|
||||||
|
}
|
||||||
|
case uuid.UUID:
|
||||||
|
if u, err := uuid.Parse(s); err == nil {
|
||||||
|
n.Val = any(u).(T)
|
||||||
|
n.Valid = true
|
||||||
|
}
|
||||||
|
case []byte:
|
||||||
|
n.Val = any([]byte(s)).(T)
|
||||||
|
n.Valid = true
|
||||||
|
case string:
|
||||||
|
n.Val = any(s).(T)
|
||||||
|
n.Valid = true
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Value implements driver.Valuer.
|
||||||
|
func (n SqlNull[T]) Value() (driver.Value, error) {
|
||||||
|
if !n.Valid {
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check if the type implements fmt.Stringer (e.g., uuid.UUID, custom types)
|
||||||
|
// Convert to string for driver compatibility
|
||||||
|
if stringer, ok := any(n.Val).(fmt.Stringer); ok {
|
||||||
|
return stringer.String(), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
return any(n.Val), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// MarshalJSON implements json.Marshaler.
|
||||||
|
func (n SqlNull[T]) MarshalJSON() ([]byte, error) {
|
||||||
|
if !n.Valid {
|
||||||
|
return []byte("null"), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check if T is []byte, and encode to base64
|
||||||
|
if _, ok := any(n.Val).([]byte); ok {
|
||||||
|
// Encode []byte as base64
|
||||||
|
encoded := base64.StdEncoding.EncodeToString(any(n.Val).([]byte))
|
||||||
|
return json.Marshal(encoded)
|
||||||
|
}
|
||||||
|
|
||||||
|
return json.Marshal(n.Val)
|
||||||
|
}
|
||||||
|
|
||||||
|
// UnmarshalJSON implements json.Unmarshaler.
|
||||||
|
func (n *SqlNull[T]) UnmarshalJSON(b []byte) error {
|
||||||
|
if len(b) == 0 || string(b) == "null" || strings.TrimSpace(string(b)) == "" {
|
||||||
|
n.Valid = false
|
||||||
|
n.Val = *new(T)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check if T is []byte, and decode from base64
|
||||||
|
var val T
|
||||||
|
if _, ok := any(val).([]byte); ok {
|
||||||
|
// Unmarshal as string first (JSON representation)
|
||||||
|
var s string
|
||||||
|
if err := json.Unmarshal(b, &s); err == nil {
|
||||||
|
// Decode from base64
|
||||||
|
if decoded, err := base64.StdEncoding.DecodeString(s); err == nil {
|
||||||
|
n.Val = any(decoded).(T)
|
||||||
|
n.Valid = true
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
// Fallback to raw string as bytes
|
||||||
|
n.Val = any([]byte(s)).(T)
|
||||||
|
n.Valid = true
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := json.Unmarshal(b, &val); err == nil {
|
||||||
|
n.Val = val
|
||||||
|
n.Valid = true
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Fallback: unmarshal as string and parse.
|
||||||
|
var s string
|
||||||
|
if err := json.Unmarshal(b, &s); err == nil {
|
||||||
|
return n.FromString(s)
|
||||||
|
}
|
||||||
|
|
||||||
|
return fmt.Errorf("cannot unmarshal %s into SqlNull[%T]", b, n.Val)
|
||||||
|
}
|
||||||
|
|
||||||
|
// MarshalYAML implements yaml.Marshaler.
|
||||||
|
func (n SqlNull[T]) MarshalYAML() (any, error) {
|
||||||
|
if !n.Valid {
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check if T is []byte, and encode to base64 (mirrors MarshalJSON).
|
||||||
|
if b, ok := any(n.Val).([]byte); ok {
|
||||||
|
return base64.StdEncoding.EncodeToString(b), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
return n.Val, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// UnmarshalYAML implements yaml.Unmarshaler.
|
||||||
|
func (n *SqlNull[T]) UnmarshalYAML(value *yaml.Node) error {
|
||||||
|
if value == nil || value.Tag == "!!null" {
|
||||||
|
n.Valid = false
|
||||||
|
n.Val = *new(T)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check if T is []byte, and decode from base64.
|
||||||
|
var zero T
|
||||||
|
if _, ok := any(zero).([]byte); ok {
|
||||||
|
var s string
|
||||||
|
if err := value.Decode(&s); err == nil {
|
||||||
|
if decoded, err := base64.StdEncoding.DecodeString(s); err == nil {
|
||||||
|
n.Val = any(decoded).(T)
|
||||||
|
n.Valid = true
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
n.Val = any([]byte(s)).(T)
|
||||||
|
n.Valid = true
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
var val T
|
||||||
|
if err := value.Decode(&val); err == nil {
|
||||||
|
n.Val = val
|
||||||
|
n.Valid = true
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Fallback: decode as string and parse.
|
||||||
|
var s string
|
||||||
|
if err := value.Decode(&s); err == nil {
|
||||||
|
return n.FromString(s)
|
||||||
|
}
|
||||||
|
|
||||||
|
return fmt.Errorf("cannot unmarshal %q into SqlNull[%T]", value.Value, n.Val)
|
||||||
|
}
|
||||||
|
|
||||||
|
// MarshalXML implements xml.Marshaler.
|
||||||
|
func (n SqlNull[T]) MarshalXML(e *xml.Encoder, start xml.StartElement) error {
|
||||||
|
if !n.Valid {
|
||||||
|
return e.EncodeElement("", start)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check if T is []byte, and encode to base64 (mirrors MarshalJSON).
|
||||||
|
if b, ok := any(n.Val).([]byte); ok {
|
||||||
|
return e.EncodeElement(base64.StdEncoding.EncodeToString(b), start)
|
||||||
|
}
|
||||||
|
|
||||||
|
return e.EncodeElement(n.Val, start)
|
||||||
|
}
|
||||||
|
|
||||||
|
// UnmarshalXML implements xml.Unmarshaler.
|
||||||
|
//
|
||||||
|
// XML has no native null representation, so an empty element unmarshals to
|
||||||
|
// an invalid (null) value rather than a zero-value-but-valid one.
|
||||||
|
func (n *SqlNull[T]) UnmarshalXML(d *xml.Decoder, start xml.StartElement) error {
|
||||||
|
var s string
|
||||||
|
if err := d.DecodeElement(&s, &start); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if s == "" {
|
||||||
|
n.Valid = false
|
||||||
|
n.Val = *new(T)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
var zero T
|
||||||
|
if _, ok := any(zero).([]byte); ok {
|
||||||
|
if decoded, err := base64.StdEncoding.DecodeString(s); err == nil {
|
||||||
|
n.Val = any(decoded).(T)
|
||||||
|
n.Valid = true
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
n.Val = any([]byte(s)).(T)
|
||||||
|
n.Valid = true
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
return n.FromString(s)
|
||||||
|
}
|
||||||
|
|
||||||
|
// String implements fmt.Stringer.
|
||||||
|
func (n SqlNull[T]) String() string {
|
||||||
|
if !n.Valid {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
// Check if the type implements fmt.Stringer for better string representation
|
||||||
|
if stringer, ok := any(n.Val).(fmt.Stringer); ok {
|
||||||
|
return stringer.String()
|
||||||
|
}
|
||||||
|
return fmt.Sprintf("%v", n.Val)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Int64 converts to int64 or 0 if invalid.
|
||||||
|
func (n SqlNull[T]) Int64() int64 {
|
||||||
|
if !n.Valid {
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
v := reflect.ValueOf(any(n.Val))
|
||||||
|
switch v.Kind() {
|
||||||
|
case reflect.Int, reflect.Int8, reflect.Int16, reflect.Int32, reflect.Int64:
|
||||||
|
return v.Int()
|
||||||
|
case reflect.Uint, reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64:
|
||||||
|
return int64(v.Uint())
|
||||||
|
case reflect.Float32, reflect.Float64:
|
||||||
|
return int64(v.Float())
|
||||||
|
case reflect.String:
|
||||||
|
i, _ := strconv.ParseInt(v.String(), 10, 64)
|
||||||
|
return i
|
||||||
|
case reflect.Bool:
|
||||||
|
if v.Bool() {
|
||||||
|
return 1
|
||||||
|
}
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
|
||||||
|
// Float64 converts to float64 or 0.0 if invalid.
|
||||||
|
func (n SqlNull[T]) Float64() float64 {
|
||||||
|
if !n.Valid {
|
||||||
|
return 0.0
|
||||||
|
}
|
||||||
|
v := reflect.ValueOf(any(n.Val))
|
||||||
|
switch v.Kind() {
|
||||||
|
case reflect.Float32, reflect.Float64:
|
||||||
|
return v.Float()
|
||||||
|
case reflect.Int, reflect.Int8, reflect.Int16, reflect.Int32, reflect.Int64:
|
||||||
|
return float64(v.Int())
|
||||||
|
case reflect.Uint, reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64:
|
||||||
|
return float64(v.Uint())
|
||||||
|
case reflect.String:
|
||||||
|
f, _ := strconv.ParseFloat(v.String(), 64)
|
||||||
|
return f
|
||||||
|
}
|
||||||
|
return 0.0
|
||||||
|
}
|
||||||
|
|
||||||
|
// Bool converts to bool or false if invalid.
|
||||||
|
func (n SqlNull[T]) Bool() bool {
|
||||||
|
if !n.Valid {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
v := reflect.ValueOf(any(n.Val))
|
||||||
|
if v.Kind() == reflect.Bool {
|
||||||
|
return v.Bool()
|
||||||
|
}
|
||||||
|
s := strings.ToLower(strings.TrimSpace(fmt.Sprint(n.Val)))
|
||||||
|
return s == "true" || s == "t" || s == "1" || s == "yes" || s == "on"
|
||||||
|
}
|
||||||
|
|
||||||
|
// Time converts to time.Time or zero if invalid.
|
||||||
|
func (n SqlNull[T]) Time() time.Time {
|
||||||
|
if !n.Valid {
|
||||||
|
return time.Time{}
|
||||||
|
}
|
||||||
|
if t, ok := any(n.Val).(time.Time); ok {
|
||||||
|
return t
|
||||||
|
}
|
||||||
|
return time.Time{}
|
||||||
|
}
|
||||||
|
|
||||||
|
// UUID converts to uuid.UUID or Nil if invalid.
|
||||||
|
func (n SqlNull[T]) UUID() uuid.UUID {
|
||||||
|
if !n.Valid {
|
||||||
|
return uuid.Nil
|
||||||
|
}
|
||||||
|
if u, ok := any(n.Val).(uuid.UUID); ok {
|
||||||
|
return u
|
||||||
|
}
|
||||||
|
return uuid.Nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Type aliases for common types.
|
||||||
|
type (
|
||||||
|
SqlInt16 = SqlNull[int16]
|
||||||
|
SqlInt32 = SqlNull[int32]
|
||||||
|
SqlInt64 = SqlNull[int64]
|
||||||
|
SqlFloat32 = SqlNull[float32]
|
||||||
|
SqlFloat64 = SqlNull[float64]
|
||||||
|
SqlBool = SqlNull[bool]
|
||||||
|
SqlString = SqlNull[string]
|
||||||
|
SqlByteArray = SqlNull[[]byte]
|
||||||
|
SqlUUID = SqlNull[uuid.UUID]
|
||||||
|
)
|
||||||
|
|
||||||
|
// SqlTimeStamp - Timestamp with custom formatting (YYYY-MM-DDTHH:MM:SS).
|
||||||
|
type SqlTimeStamp struct{ SqlNull[time.Time] }
|
||||||
|
|
||||||
|
func (t SqlTimeStamp) MarshalJSON() ([]byte, error) {
|
||||||
|
if !t.Valid || t.Val.IsZero() || t.Val.Before(time.Date(0002, 1, 1, 0, 0, 0, 0, time.UTC)) {
|
||||||
|
return []byte("null"), nil
|
||||||
|
}
|
||||||
|
return fmt.Appendf(nil, `"%s"`, t.Val.Format("2006-01-02T15:04:05")), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *SqlTimeStamp) UnmarshalJSON(b []byte) error {
|
||||||
|
if err := t.SqlNull.UnmarshalJSON(b); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if t.Valid && (t.Val.IsZero() || t.Val.Format("2006-01-02T15:04:05") == "0001-01-01T00:00:00") {
|
||||||
|
t.Valid = false
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t SqlTimeStamp) Value() (driver.Value, error) {
|
||||||
|
if !t.Valid || t.Val.IsZero() || t.Val.Before(time.Date(0002, 1, 1, 0, 0, 0, 0, time.UTC)) {
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
return t.Val.Format("2006-01-02T15:04:05"), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t SqlTimeStamp) MarshalYAML() (any, error) {
|
||||||
|
if !t.Valid || t.Val.IsZero() || t.Val.Before(time.Date(0002, 1, 1, 0, 0, 0, 0, time.UTC)) {
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
return t.Val.Format("2006-01-02T15:04:05"), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *SqlTimeStamp) UnmarshalYAML(value *yaml.Node) error {
|
||||||
|
if err := t.SqlNull.UnmarshalYAML(value); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if t.Valid && (t.Val.IsZero() || t.Val.Format("2006-01-02T15:04:05") == "0001-01-01T00:00:00") {
|
||||||
|
t.Valid = false
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t SqlTimeStamp) MarshalXML(e *xml.Encoder, start xml.StartElement) error {
|
||||||
|
if !t.Valid || t.Val.IsZero() || t.Val.Before(time.Date(0002, 1, 1, 0, 0, 0, 0, time.UTC)) {
|
||||||
|
return e.EncodeElement("", start)
|
||||||
|
}
|
||||||
|
return e.EncodeElement(t.Val.Format("2006-01-02T15:04:05"), start)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *SqlTimeStamp) UnmarshalXML(d *xml.Decoder, start xml.StartElement) error {
|
||||||
|
var s string
|
||||||
|
if err := d.DecodeElement(&s, &start); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if s == "" {
|
||||||
|
t.Valid = false
|
||||||
|
t.Val = time.Time{}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
tm, err := tryParseDT(s)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
t.Val = tm
|
||||||
|
t.Valid = !tm.IsZero() && tm.Format("2006-01-02T15:04:05") != "0001-01-01T00:00:00"
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func SqlTimeStampNow() SqlTimeStamp {
|
||||||
|
return SqlTimeStamp{SqlNull: SqlNull[time.Time]{Val: time.Now(), Valid: true}}
|
||||||
|
}
|
||||||
|
|
||||||
|
// SqlDate - Date only (YYYY-MM-DD).
|
||||||
|
type SqlDate struct{ SqlNull[time.Time] }
|
||||||
|
|
||||||
|
func (d SqlDate) MarshalJSON() ([]byte, error) {
|
||||||
|
if !d.Valid || d.Val.IsZero() {
|
||||||
|
return []byte("null"), nil
|
||||||
|
}
|
||||||
|
s := d.Val.Format("2006-01-02")
|
||||||
|
if strings.HasPrefix(s, "0001-01-01") {
|
||||||
|
return []byte("null"), nil
|
||||||
|
}
|
||||||
|
return fmt.Appendf(nil, `"%s"`, s), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (d *SqlDate) UnmarshalJSON(b []byte) error {
|
||||||
|
if err := d.SqlNull.UnmarshalJSON(b); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if d.Valid && d.Val.Format("2006-01-02") <= "0001-01-01" {
|
||||||
|
d.Valid = false
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (d SqlDate) Value() (driver.Value, error) {
|
||||||
|
if !d.Valid || d.Val.IsZero() {
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
s := d.Val.Format("2006-01-02")
|
||||||
|
if s <= "0001-01-01" {
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
return s, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (d SqlDate) String() string {
|
||||||
|
if !d.Valid {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
s := d.Val.Format("2006-01-02")
|
||||||
|
if strings.HasPrefix(s, "0001-01-01") || strings.HasPrefix(s, "1800-12-31") {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
return s
|
||||||
|
}
|
||||||
|
|
||||||
|
func (d SqlDate) MarshalYAML() (any, error) {
|
||||||
|
if !d.Valid || d.Val.IsZero() {
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
s := d.Val.Format("2006-01-02")
|
||||||
|
if strings.HasPrefix(s, "0001-01-01") {
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
return s, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (d *SqlDate) UnmarshalYAML(value *yaml.Node) error {
|
||||||
|
if err := d.SqlNull.UnmarshalYAML(value); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if d.Valid && d.Val.Format("2006-01-02") <= "0001-01-01" {
|
||||||
|
d.Valid = false
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (d SqlDate) MarshalXML(e *xml.Encoder, start xml.StartElement) error {
|
||||||
|
if !d.Valid || d.Val.IsZero() {
|
||||||
|
return e.EncodeElement("", start)
|
||||||
|
}
|
||||||
|
s := d.Val.Format("2006-01-02")
|
||||||
|
if strings.HasPrefix(s, "0001-01-01") {
|
||||||
|
return e.EncodeElement("", start)
|
||||||
|
}
|
||||||
|
return e.EncodeElement(s, start)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (d *SqlDate) UnmarshalXML(dec *xml.Decoder, start xml.StartElement) error {
|
||||||
|
var s string
|
||||||
|
if err := dec.DecodeElement(&s, &start); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if s == "" {
|
||||||
|
d.Valid = false
|
||||||
|
d.Val = time.Time{}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
tm, err := tryParseDT(s)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
d.Val = tm
|
||||||
|
d.Valid = !tm.IsZero() && tm.Format("2006-01-02") > "0001-01-01"
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func SqlDateNow() SqlDate {
|
||||||
|
return SqlDate{SqlNull: SqlNull[time.Time]{Val: time.Now(), Valid: true}}
|
||||||
|
}
|
||||||
|
|
||||||
|
// SqlTime - Time only (HH:MM:SS).
|
||||||
|
type SqlTime struct{ SqlNull[time.Time] }
|
||||||
|
|
||||||
|
func (t SqlTime) MarshalJSON() ([]byte, error) {
|
||||||
|
if !t.Valid || t.Val.IsZero() {
|
||||||
|
return []byte("null"), nil
|
||||||
|
}
|
||||||
|
s := t.Val.Format("15:04:05")
|
||||||
|
if s == "00:00:00" {
|
||||||
|
return []byte("null"), nil
|
||||||
|
}
|
||||||
|
return fmt.Appendf(nil, `"%s"`, s), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *SqlTime) UnmarshalJSON(b []byte) error {
|
||||||
|
if err := t.SqlNull.UnmarshalJSON(b); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if t.Valid && t.Val.Format("15:04:05") == "00:00:00" {
|
||||||
|
t.Valid = false
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t SqlTime) Value() (driver.Value, error) {
|
||||||
|
if !t.Valid || t.Val.IsZero() {
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
return t.Val.Format("15:04:05"), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t SqlTime) String() string {
|
||||||
|
if !t.Valid {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
return t.Val.Format("15:04:05")
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t SqlTime) MarshalYAML() (any, error) {
|
||||||
|
if !t.Valid || t.Val.IsZero() {
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
s := t.Val.Format("15:04:05")
|
||||||
|
if s == "00:00:00" {
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
return s, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *SqlTime) UnmarshalYAML(value *yaml.Node) error {
|
||||||
|
if err := t.SqlNull.UnmarshalYAML(value); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if t.Valid && t.Val.Format("15:04:05") == "00:00:00" {
|
||||||
|
t.Valid = false
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t SqlTime) MarshalXML(e *xml.Encoder, start xml.StartElement) error {
|
||||||
|
if !t.Valid || t.Val.IsZero() {
|
||||||
|
return e.EncodeElement("", start)
|
||||||
|
}
|
||||||
|
s := t.Val.Format("15:04:05")
|
||||||
|
if s == "00:00:00" {
|
||||||
|
return e.EncodeElement("", start)
|
||||||
|
}
|
||||||
|
return e.EncodeElement(s, start)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *SqlTime) UnmarshalXML(d *xml.Decoder, start xml.StartElement) error {
|
||||||
|
var s string
|
||||||
|
if err := d.DecodeElement(&s, &start); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if s == "" {
|
||||||
|
t.Valid = false
|
||||||
|
t.Val = time.Time{}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
tm, err := tryParseDT(s)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
t.Val = tm
|
||||||
|
t.Valid = !tm.IsZero() && tm.Format("15:04:05") != "00:00:00"
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func SqlTimeNow() SqlTime {
|
||||||
|
return SqlTime{SqlNull: SqlNull[time.Time]{Val: time.Now(), Valid: true}}
|
||||||
|
}
|
||||||
|
|
||||||
|
// SqlJSONB - Nullable JSONB as []byte.
|
||||||
|
type SqlJSONB []byte
|
||||||
|
|
||||||
|
// SqlJSON - Nullable JSON as []byte. PostgreSQL's json and jsonb types share
|
||||||
|
// the same textual representation and Go marshalling behavior, differing only
|
||||||
|
// in server-side storage, so SqlJSON is an alias of SqlJSONB.
|
||||||
|
type SqlJSON = SqlJSONB
|
||||||
|
|
||||||
|
// Scan implements sql.Scanner.
|
||||||
|
func (n *SqlJSONB) Scan(value any) error {
|
||||||
|
if value == nil {
|
||||||
|
*n = nil
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
switch v := value.(type) {
|
||||||
|
case string:
|
||||||
|
*n = []byte(v)
|
||||||
|
case []byte:
|
||||||
|
*n = v
|
||||||
|
default:
|
||||||
|
dat, err := json.Marshal(value)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("failed to marshal value to JSON: %v", err)
|
||||||
|
}
|
||||||
|
*n = dat
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Value implements driver.Valuer.
|
||||||
|
func (n SqlJSONB) Value() (driver.Value, error) {
|
||||||
|
if len(n) == 0 {
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
var js any
|
||||||
|
if err := json.Unmarshal(n, &js); err != nil {
|
||||||
|
return nil, fmt.Errorf("invalid JSON: %v", err)
|
||||||
|
}
|
||||||
|
return string(n), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// MarshalJSON implements json.Marshaler.
|
||||||
|
func (n SqlJSONB) MarshalJSON() ([]byte, error) {
|
||||||
|
if len(n) == 0 {
|
||||||
|
return []byte("null"), nil
|
||||||
|
}
|
||||||
|
var obj any
|
||||||
|
if err := json.Unmarshal(n, &obj); err != nil {
|
||||||
|
return []byte("null"), nil
|
||||||
|
}
|
||||||
|
return n, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// UnmarshalJSON implements json.Unmarshaler.
|
||||||
|
func (n *SqlJSONB) UnmarshalJSON(b []byte) error {
|
||||||
|
s := strings.TrimSpace(string(b))
|
||||||
|
if s == "null" || s == "" || (!strings.HasPrefix(s, "{") && !strings.HasPrefix(s, "[")) {
|
||||||
|
*n = nil
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
*n = b
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// MarshalYAML implements yaml.Marshaler. The underlying JSON is decoded into
|
||||||
|
// a generic value first so it renders as native YAML mappings/sequences
|
||||||
|
// rather than an embedded JSON string.
|
||||||
|
func (n SqlJSONB) MarshalYAML() (any, error) {
|
||||||
|
if len(n) == 0 {
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
var v any
|
||||||
|
if err := json.Unmarshal(n, &v); err != nil {
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
return v, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// UnmarshalYAML implements yaml.Unmarshaler.
|
||||||
|
func (n *SqlJSONB) UnmarshalYAML(value *yaml.Node) error {
|
||||||
|
if value == nil || value.Tag == "!!null" {
|
||||||
|
*n = nil
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
var v any
|
||||||
|
if err := value.Decode(&v); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
b, err := json.Marshal(v)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
*n = b
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// MarshalXML implements xml.Marshaler. JSON has no clean structural mapping
|
||||||
|
// to XML, so the raw JSON text is emitted as the element's text content.
|
||||||
|
func (n SqlJSONB) MarshalXML(e *xml.Encoder, start xml.StartElement) error {
|
||||||
|
if len(n) == 0 {
|
||||||
|
return e.EncodeElement("", start)
|
||||||
|
}
|
||||||
|
var obj any
|
||||||
|
if err := json.Unmarshal(n, &obj); err != nil {
|
||||||
|
return e.EncodeElement("", start)
|
||||||
|
}
|
||||||
|
return e.EncodeElement(string(n), start)
|
||||||
|
}
|
||||||
|
|
||||||
|
// UnmarshalXML implements xml.Unmarshaler, reading back the raw JSON text
|
||||||
|
// written by MarshalXML.
|
||||||
|
func (n *SqlJSONB) UnmarshalXML(d *xml.Decoder, start xml.StartElement) error {
|
||||||
|
var s string
|
||||||
|
if err := d.DecodeElement(&s, &start); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
s = strings.TrimSpace(s)
|
||||||
|
if s == "" {
|
||||||
|
*n = nil
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
*n = []byte(s)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (n SqlJSONB) AsMap() (map[string]any, error) {
|
||||||
|
if len(n) == 0 {
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
js := make(map[string]any)
|
||||||
|
if err := json.Unmarshal(n, &js); err != nil {
|
||||||
|
return nil, fmt.Errorf("invalid JSON: %v", err)
|
||||||
|
}
|
||||||
|
return js, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (n SqlJSONB) AsSlice() ([]any, error) {
|
||||||
|
if len(n) == 0 {
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
js := make([]any, 0)
|
||||||
|
if err := json.Unmarshal(n, &js); err != nil {
|
||||||
|
return nil, fmt.Errorf("invalid JSON: %v", err)
|
||||||
|
}
|
||||||
|
return js, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// TryIfInt64 tries to parse any value to int64 with default.
|
||||||
|
func TryIfInt64(v any, def int64) int64 {
|
||||||
|
switch val := v.(type) {
|
||||||
|
case string:
|
||||||
|
i, err := strconv.ParseInt(val, 10, 64)
|
||||||
|
if err != nil {
|
||||||
|
return def
|
||||||
|
}
|
||||||
|
return i
|
||||||
|
case int:
|
||||||
|
return int64(val)
|
||||||
|
case int8:
|
||||||
|
return int64(val)
|
||||||
|
case int16:
|
||||||
|
return int64(val)
|
||||||
|
case int32:
|
||||||
|
return int64(val)
|
||||||
|
case int64:
|
||||||
|
return val
|
||||||
|
case uint:
|
||||||
|
return int64(val)
|
||||||
|
case uint8:
|
||||||
|
return int64(val)
|
||||||
|
case uint16:
|
||||||
|
return int64(val)
|
||||||
|
case uint32:
|
||||||
|
return int64(val)
|
||||||
|
case uint64:
|
||||||
|
return int64(val)
|
||||||
|
case float32:
|
||||||
|
return int64(val)
|
||||||
|
case float64:
|
||||||
|
return int64(val)
|
||||||
|
case []byte:
|
||||||
|
i, err := strconv.ParseInt(string(val), 10, 64)
|
||||||
|
if err != nil {
|
||||||
|
return def
|
||||||
|
}
|
||||||
|
return i
|
||||||
|
default:
|
||||||
|
return def
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Constructor helpers - clean and fast value creation
|
||||||
|
func Null[T any](v T, valid bool) SqlNull[T] {
|
||||||
|
return SqlNull[T]{Val: v, Valid: valid}
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewSql[T any](value any) SqlNull[T] {
|
||||||
|
n := SqlNull[T]{}
|
||||||
|
|
||||||
|
if value == nil {
|
||||||
|
return n
|
||||||
|
}
|
||||||
|
|
||||||
|
// Fast path: exact match
|
||||||
|
if v, ok := value.(T); ok {
|
||||||
|
n.Val = v
|
||||||
|
n.Valid = true
|
||||||
|
return n
|
||||||
|
}
|
||||||
|
|
||||||
|
// Try from another SqlNull
|
||||||
|
if sn, ok := value.(SqlNull[T]); ok {
|
||||||
|
return sn
|
||||||
|
}
|
||||||
|
|
||||||
|
// Convert via string
|
||||||
|
_ = n.FromString(fmt.Sprintf("%v", value))
|
||||||
|
return n
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewSqlInt16(v int16) SqlInt16 {
|
||||||
|
return SqlInt16{Val: v, Valid: true}
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewSqlInt32(v int32) SqlInt32 {
|
||||||
|
return SqlInt32{Val: v, Valid: true}
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewSqlInt64(v int64) SqlInt64 {
|
||||||
|
return SqlInt64{Val: v, Valid: true}
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewSqlFloat32(v float32) SqlFloat32 {
|
||||||
|
return SqlFloat32{Val: v, Valid: true}
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewSqlFloat64(v float64) SqlFloat64 {
|
||||||
|
return SqlFloat64{Val: v, Valid: true}
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewSqlBool(v bool) SqlBool {
|
||||||
|
return SqlBool{Val: v, Valid: true}
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewSqlString(v string) SqlString {
|
||||||
|
return SqlString{Val: v, Valid: true}
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewSqlByteArray(v []byte) SqlByteArray {
|
||||||
|
return SqlByteArray{Val: v, Valid: true}
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewSqlUUID(v uuid.UUID) SqlUUID {
|
||||||
|
return SqlUUID{Val: v, Valid: true}
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewSqlTimeStamp(v time.Time) SqlTimeStamp {
|
||||||
|
return SqlTimeStamp{SqlNull: SqlNull[time.Time]{Val: v, Valid: true}}
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewSqlDate(v time.Time) SqlDate {
|
||||||
|
return SqlDate{SqlNull: SqlNull[time.Time]{Val: v, Valid: true}}
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewSqlTime(v time.Time) SqlTime {
|
||||||
|
return SqlTime{SqlNull: SqlNull[time.Time]{Val: v, Valid: true}}
|
||||||
|
}
|
||||||
@@ -0,0 +1,134 @@
|
|||||||
|
package sqltypes
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
// TestSqlNull_FromString_UnsignedInt guards against a regression where
|
||||||
|
// FromString called reflect.Value.SetInt on unsigned-kind fields, which
|
||||||
|
// panics since SetInt only accepts Int/Int8/.../Int64 kinds.
|
||||||
|
func TestSqlNull_FromString_UnsignedInt(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
input string
|
||||||
|
expected uint32
|
||||||
|
}{
|
||||||
|
{"simple", "123", 123},
|
||||||
|
{"zero", "0", 0},
|
||||||
|
{"large", "4000000000", 4000000000},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
var n SqlNull[uint32]
|
||||||
|
if err := n.FromString(tt.input); err != nil {
|
||||||
|
t.Fatalf("FromString failed: %v", err)
|
||||||
|
}
|
||||||
|
if !n.Valid {
|
||||||
|
t.Fatalf("expected valid=true")
|
||||||
|
}
|
||||||
|
if n.Val != tt.expected {
|
||||||
|
t.Errorf("expected %d, got %d", tt.expected, n.Val)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestSqlNull_FromString_Int64Precision guards against a regression where
|
||||||
|
// FromString unconditionally re-parsed the string as float64 after a
|
||||||
|
// successful ParseInt, corrupting large int64 values due to float rounding
|
||||||
|
// (e.g. math.MaxInt64 became negative).
|
||||||
|
func TestSqlNull_FromString_Int64Precision(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
input string
|
||||||
|
expected int64
|
||||||
|
}{
|
||||||
|
{"max int64", "9223372036854775807", 9223372036854775807},
|
||||||
|
{"min int64", "-9223372036854775808", -9223372036854775808},
|
||||||
|
{"large safe value", "1234567890123456", 1234567890123456},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
var n SqlNull[int64]
|
||||||
|
if err := n.FromString(tt.input); err != nil {
|
||||||
|
t.Fatalf("FromString failed: %v", err)
|
||||||
|
}
|
||||||
|
if !n.Valid {
|
||||||
|
t.Fatalf("expected valid=true")
|
||||||
|
}
|
||||||
|
if n.Val != tt.expected {
|
||||||
|
t.Errorf("expected %d, got %d", tt.expected, n.Val)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestSqlNull_FromString_FloatFallback verifies that non-integral strings
|
||||||
|
// still parse via the float fallback path for integer-kind fields.
|
||||||
|
func TestSqlNull_FromString_FloatFallback(t *testing.T) {
|
||||||
|
var n SqlNull[int32]
|
||||||
|
if err := n.FromString("42.9"); err != nil {
|
||||||
|
t.Fatalf("FromString failed: %v", err)
|
||||||
|
}
|
||||||
|
if !n.Valid || n.Val != 42 {
|
||||||
|
t.Errorf("expected valid int32=42, got valid=%v val=%d", n.Valid, n.Val)
|
||||||
|
}
|
||||||
|
|
||||||
|
var u SqlNull[uint32]
|
||||||
|
if err := u.FromString("42.9"); err != nil {
|
||||||
|
t.Fatalf("FromString failed: %v", err)
|
||||||
|
}
|
||||||
|
if !u.Valid || u.Val != 42 {
|
||||||
|
t.Errorf("expected valid uint32=42, got valid=%v val=%d", u.Valid, u.Val)
|
||||||
|
}
|
||||||
|
|
||||||
|
var neg SqlNull[uint32]
|
||||||
|
if err := neg.FromString("-1.5"); err != nil {
|
||||||
|
t.Fatalf("FromString failed: %v", err)
|
||||||
|
}
|
||||||
|
if neg.Valid {
|
||||||
|
t.Errorf("expected invalid for negative value into unsigned type, got %v", neg.Val)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestSqlNull_UnsignedTypes_ScanAndJSON exercises Scan and JSON round-trip
|
||||||
|
// for every unsigned alias to make sure none of them panic.
|
||||||
|
func TestSqlNull_UnsignedTypes_ScanAndJSON(t *testing.T) {
|
||||||
|
t.Run("uint8 scan from string", func(t *testing.T) {
|
||||||
|
var n SqlNull[uint8]
|
||||||
|
if err := n.Scan("200"); err != nil {
|
||||||
|
t.Fatalf("Scan failed: %v", err)
|
||||||
|
}
|
||||||
|
if !n.Valid || n.Val != 200 {
|
||||||
|
t.Errorf("expected valid uint8=200, got valid=%v val=%d", n.Valid, n.Val)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("uint64 json round-trip", func(t *testing.T) {
|
||||||
|
n := Null(uint64(18446744073709551615), true)
|
||||||
|
data, err := json.Marshal(n)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Marshal failed: %v", err)
|
||||||
|
}
|
||||||
|
var n2 SqlNull[uint64]
|
||||||
|
if err := json.Unmarshal(data, &n2); err != nil {
|
||||||
|
t.Fatalf("Unmarshal failed: %v", err)
|
||||||
|
}
|
||||||
|
if n2.Val != n.Val {
|
||||||
|
t.Errorf("expected %d, got %d", n.Val, n2.Val)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("uint scan from numeric string fallback", func(t *testing.T) {
|
||||||
|
var n SqlNull[uint]
|
||||||
|
if err := n.Scan([]byte("77")); err != nil {
|
||||||
|
t.Fatalf("Scan failed: %v", err)
|
||||||
|
}
|
||||||
|
if !n.Valid || n.Val != 77 {
|
||||||
|
t.Errorf("expected valid uint=77, got valid=%v val=%d", n.Valid, n.Val)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
@@ -0,0 +1,958 @@
|
|||||||
|
package sqltypes
|
||||||
|
|
||||||
|
import (
|
||||||
|
"database/sql/driver"
|
||||||
|
"encoding/json"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/google/uuid"
|
||||||
|
)
|
||||||
|
|
||||||
|
// TestNewSqlInt16 tests NewSqlInt16 type
|
||||||
|
func TestNewSqlInt16(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
input interface{}
|
||||||
|
expected SqlInt16
|
||||||
|
}{
|
||||||
|
{"int", 42, Null(int16(42), true)},
|
||||||
|
{"int32", int32(100), NewSqlInt16(100)},
|
||||||
|
{"int64", int64(200), NewSqlInt16(200)},
|
||||||
|
{"string", "123", NewSqlInt16(123)},
|
||||||
|
{"nil", nil, Null(int16(0), false)},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
var n SqlInt16
|
||||||
|
if err := n.Scan(tt.input); err != nil {
|
||||||
|
t.Fatalf("Scan failed: %v", err)
|
||||||
|
}
|
||||||
|
if n != tt.expected {
|
||||||
|
t.Errorf("expected %v, got %v", tt.expected, n)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNewSqlInt16_Value(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
input SqlInt16
|
||||||
|
expected driver.Value
|
||||||
|
}{
|
||||||
|
{"zero", Null(int16(0), false), nil},
|
||||||
|
{"positive", NewSqlInt16(42), int16(42)},
|
||||||
|
{"negative", NewSqlInt16(-10), int16(-10)},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
val, err := tt.input.Value()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Value failed: %v", err)
|
||||||
|
}
|
||||||
|
if val != tt.expected {
|
||||||
|
t.Errorf("expected %v, got %v", tt.expected, val)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNewSqlInt16_JSON(t *testing.T) {
|
||||||
|
n := NewSqlInt16(42)
|
||||||
|
|
||||||
|
// Marshal
|
||||||
|
data, err := json.Marshal(n)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Marshal failed: %v", err)
|
||||||
|
}
|
||||||
|
expected := "42"
|
||||||
|
if string(data) != expected {
|
||||||
|
t.Errorf("expected %s, got %s", expected, string(data))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Unmarshal
|
||||||
|
var n2 SqlInt16
|
||||||
|
if err := json.Unmarshal([]byte("123"), &n2); err != nil {
|
||||||
|
t.Fatalf("Unmarshal failed: %v", err)
|
||||||
|
}
|
||||||
|
if n2.Int64() != 123 {
|
||||||
|
t.Errorf("expected 123, got %d", n2.Int64())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestNewSqlInt64 tests NewSqlInt64 type
|
||||||
|
func TestNewSqlInt64(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
input interface{}
|
||||||
|
expected SqlInt64
|
||||||
|
}{
|
||||||
|
{"int", 42, NewSqlInt64(42)},
|
||||||
|
{"int32", int32(100), NewSqlInt64(100)},
|
||||||
|
{"int64", int64(9223372036854775807), NewSqlInt64(9223372036854775807)},
|
||||||
|
{"uint32", uint32(100), NewSqlInt64(100)},
|
||||||
|
{"uint64", uint64(200), NewSqlInt64(200)},
|
||||||
|
{"nil", nil, SqlInt64{}},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
var n SqlInt64
|
||||||
|
if err := n.Scan(tt.input); err != nil {
|
||||||
|
t.Fatalf("Scan failed: %v", err)
|
||||||
|
}
|
||||||
|
if n != tt.expected {
|
||||||
|
t.Errorf("expected %v, got %v", tt.expected, n)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestSqlFloat64 tests SqlFloat64 type
|
||||||
|
func TestSqlFloat64(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
input interface{}
|
||||||
|
expected float64
|
||||||
|
valid bool
|
||||||
|
}{
|
||||||
|
{"float64", float64(3.14), 3.14, true},
|
||||||
|
{"float32", float32(2.5), 2.5, true},
|
||||||
|
{"int", 42, 42.0, true},
|
||||||
|
{"int64", int64(100), 100.0, true},
|
||||||
|
{"nil", nil, 0, false},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
var n SqlFloat64
|
||||||
|
if err := n.Scan(tt.input); err != nil {
|
||||||
|
t.Fatalf("Scan failed: %v", err)
|
||||||
|
}
|
||||||
|
if n.Valid != tt.valid {
|
||||||
|
t.Errorf("expected valid=%v, got valid=%v", tt.valid, n.Valid)
|
||||||
|
}
|
||||||
|
if tt.valid && n.Float64() != tt.expected {
|
||||||
|
t.Errorf("expected %v, got %v", tt.expected, n.Float64())
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestSqlTimeStamp tests SqlTimeStamp type
|
||||||
|
func TestSqlTimeStamp(t *testing.T) {
|
||||||
|
now := time.Now()
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
input interface{}
|
||||||
|
}{
|
||||||
|
{"time.Time", now},
|
||||||
|
{"string RFC3339", now.Format(time.RFC3339)},
|
||||||
|
{"string date", "2024-01-15"},
|
||||||
|
{"string datetime", "2024-01-15T10:30:00"},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
var ts SqlTimeStamp
|
||||||
|
if err := ts.Scan(tt.input); err != nil {
|
||||||
|
t.Fatalf("Scan failed: %v", err)
|
||||||
|
}
|
||||||
|
if ts.Time().IsZero() {
|
||||||
|
t.Error("expected non-zero time")
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSqlTimeStamp_JSON(t *testing.T) {
|
||||||
|
now := time.Date(2024, 1, 15, 10, 30, 45, 0, time.UTC)
|
||||||
|
ts := NewSqlTimeStamp(now)
|
||||||
|
|
||||||
|
// Marshal
|
||||||
|
data, err := json.Marshal(ts)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Marshal failed: %v", err)
|
||||||
|
}
|
||||||
|
expected := `"2024-01-15T10:30:45"`
|
||||||
|
if string(data) != expected {
|
||||||
|
t.Errorf("expected %s, got %s", expected, string(data))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Unmarshal
|
||||||
|
var ts2 SqlTimeStamp
|
||||||
|
if err := json.Unmarshal([]byte(`"2024-01-15T10:30:45"`), &ts2); err != nil {
|
||||||
|
t.Fatalf("Unmarshal failed: %v", err)
|
||||||
|
}
|
||||||
|
if ts2.Time().Year() != 2024 {
|
||||||
|
t.Errorf("expected year 2024, got %d", ts2.Time().Year())
|
||||||
|
}
|
||||||
|
|
||||||
|
// Test null
|
||||||
|
var ts3 SqlTimeStamp
|
||||||
|
if err := json.Unmarshal([]byte("null"), &ts3); err != nil {
|
||||||
|
t.Fatalf("Unmarshal null failed: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestSqlDate tests SqlDate type
|
||||||
|
func TestSqlDate(t *testing.T) {
|
||||||
|
now := time.Now()
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
input interface{}
|
||||||
|
}{
|
||||||
|
{"time.Time", now},
|
||||||
|
{"string date", "2024-01-15"},
|
||||||
|
{"string UK format", "15/01/2024"},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
var d SqlDate
|
||||||
|
if err := d.Scan(tt.input); err != nil {
|
||||||
|
t.Fatalf("Scan failed: %v", err)
|
||||||
|
}
|
||||||
|
if d.String() == "0" {
|
||||||
|
t.Error("expected non-zero date")
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSqlDate_JSON(t *testing.T) {
|
||||||
|
date := NewSqlDate(time.Date(2024, 1, 15, 0, 0, 0, 0, time.UTC))
|
||||||
|
|
||||||
|
// Marshal
|
||||||
|
data, err := json.Marshal(date)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Marshal failed: %v", err)
|
||||||
|
}
|
||||||
|
expected := `"2024-01-15"`
|
||||||
|
if string(data) != expected {
|
||||||
|
t.Errorf("expected %s, got %s", expected, string(data))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Unmarshal
|
||||||
|
var d2 SqlDate
|
||||||
|
if err := json.Unmarshal([]byte(`"2024-01-15"`), &d2); err != nil {
|
||||||
|
t.Fatalf("Unmarshal failed: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestSqlTime tests SqlTime type
|
||||||
|
func TestSqlTime(t *testing.T) {
|
||||||
|
now := time.Now()
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
input interface{}
|
||||||
|
expected string
|
||||||
|
}{
|
||||||
|
{"time.Time", now, now.Format("15:04:05")},
|
||||||
|
{"string time", "10:30:45", "10:30:45"},
|
||||||
|
{"string short time", "10:30", "10:30:00"},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
var tm SqlTime
|
||||||
|
if err := tm.Scan(tt.input); err != nil {
|
||||||
|
t.Fatalf("Scan failed: %v", err)
|
||||||
|
}
|
||||||
|
if tm.String() != tt.expected {
|
||||||
|
t.Errorf("expected %s, got %s", tt.expected, tm.String())
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestSqlJSONB tests SqlJSONB type
|
||||||
|
func TestSqlJSONB_Scan(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
input interface{}
|
||||||
|
expected string
|
||||||
|
}{
|
||||||
|
{"string JSON object", `{"key":"value"}`, `{"key":"value"}`},
|
||||||
|
{"string JSON array", `[1,2,3]`, `[1,2,3]`},
|
||||||
|
{"bytes", []byte(`{"test":true}`), `{"test":true}`},
|
||||||
|
{"nil", nil, ""},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
var j SqlJSONB
|
||||||
|
if err := j.Scan(tt.input); err != nil {
|
||||||
|
t.Fatalf("Scan failed: %v", err)
|
||||||
|
}
|
||||||
|
if tt.expected == "" && j == nil {
|
||||||
|
return // nil case
|
||||||
|
}
|
||||||
|
if string(j) != tt.expected {
|
||||||
|
t.Errorf("expected %s, got %s", tt.expected, string(j))
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSqlJSONB_Value(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
input SqlJSONB
|
||||||
|
expected string
|
||||||
|
wantErr bool
|
||||||
|
}{
|
||||||
|
{"valid object", SqlJSONB(`{"key":"value"}`), `{"key":"value"}`, false},
|
||||||
|
{"valid array", SqlJSONB(`[1,2,3]`), `[1,2,3]`, false},
|
||||||
|
{"empty", SqlJSONB{}, "", false},
|
||||||
|
{"nil", nil, "", false},
|
||||||
|
{"invalid JSON", SqlJSONB(`{invalid`), "", true},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
val, err := tt.input.Value()
|
||||||
|
if tt.wantErr {
|
||||||
|
if err == nil {
|
||||||
|
t.Error("expected error, got nil")
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Value failed: %v", err)
|
||||||
|
}
|
||||||
|
if tt.expected == "" && val == nil {
|
||||||
|
return // nil case
|
||||||
|
}
|
||||||
|
if val.(string) != tt.expected {
|
||||||
|
t.Errorf("expected %s, got %s", tt.expected, val)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSqlJSONB_JSON(t *testing.T) {
|
||||||
|
// Marshal
|
||||||
|
j := SqlJSONB(`{"name":"test","count":42}`)
|
||||||
|
data, err := json.Marshal(j)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Marshal failed: %v", err)
|
||||||
|
}
|
||||||
|
var result map[string]interface{}
|
||||||
|
if err := json.Unmarshal(data, &result); err != nil {
|
||||||
|
t.Fatalf("Unmarshal result failed: %v", err)
|
||||||
|
}
|
||||||
|
if result["name"] != "test" {
|
||||||
|
t.Errorf("expected name=test, got %v", result["name"])
|
||||||
|
}
|
||||||
|
|
||||||
|
// Unmarshal
|
||||||
|
var j2 SqlJSONB
|
||||||
|
if err := json.Unmarshal([]byte(`{"key":"value"}`), &j2); err != nil {
|
||||||
|
t.Fatalf("Unmarshal failed: %v", err)
|
||||||
|
}
|
||||||
|
if string(j2) != `{"key":"value"}` {
|
||||||
|
t.Errorf("expected {\"key\":\"value\"}, got %s", string(j2))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Test null
|
||||||
|
var j3 SqlJSONB
|
||||||
|
if err := json.Unmarshal([]byte("null"), &j3); err != nil {
|
||||||
|
t.Fatalf("Unmarshal null failed: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSqlJSONB_AsMap(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
input SqlJSONB
|
||||||
|
wantErr bool
|
||||||
|
wantNil bool
|
||||||
|
}{
|
||||||
|
{"valid object", SqlJSONB(`{"name":"test","age":30}`), false, false},
|
||||||
|
{"empty", SqlJSONB{}, false, true},
|
||||||
|
{"nil", nil, false, true},
|
||||||
|
{"invalid JSON", SqlJSONB(`{invalid`), true, false},
|
||||||
|
{"array not object", SqlJSONB(`[1,2,3]`), true, false},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
m, err := tt.input.AsMap()
|
||||||
|
if tt.wantErr {
|
||||||
|
if err == nil {
|
||||||
|
t.Error("expected error, got nil")
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("AsMap failed: %v", err)
|
||||||
|
}
|
||||||
|
if tt.wantNil {
|
||||||
|
if m != nil {
|
||||||
|
t.Errorf("expected nil, got %v", m)
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if m == nil {
|
||||||
|
t.Error("expected non-nil map")
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSqlJSONB_AsSlice(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
input SqlJSONB
|
||||||
|
wantErr bool
|
||||||
|
wantNil bool
|
||||||
|
}{
|
||||||
|
{"valid array", SqlJSONB(`[1,2,3]`), false, false},
|
||||||
|
{"empty", SqlJSONB{}, false, true},
|
||||||
|
{"nil", nil, false, true},
|
||||||
|
{"invalid JSON", SqlJSONB(`[invalid`), true, false},
|
||||||
|
{"object not array", SqlJSONB(`{"key":"value"}`), true, false},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
s, err := tt.input.AsSlice()
|
||||||
|
if tt.wantErr {
|
||||||
|
if err == nil {
|
||||||
|
t.Error("expected error, got nil")
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("AsSlice failed: %v", err)
|
||||||
|
}
|
||||||
|
if tt.wantNil {
|
||||||
|
if s != nil {
|
||||||
|
t.Errorf("expected nil, got %v", s)
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if s == nil {
|
||||||
|
t.Error("expected non-nil slice")
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestSqlUUID tests SqlUUID type
|
||||||
|
func TestSqlUUID_Scan(t *testing.T) {
|
||||||
|
testUUID := uuid.New()
|
||||||
|
testUUIDStr := testUUID.String()
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
input interface{}
|
||||||
|
expected string
|
||||||
|
valid bool
|
||||||
|
}{
|
||||||
|
{"string UUID", testUUIDStr, testUUIDStr, true},
|
||||||
|
{"bytes UUID", []byte(testUUIDStr), testUUIDStr, true},
|
||||||
|
{"nil", nil, "", false},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
var u SqlUUID
|
||||||
|
if err := u.Scan(tt.input); err != nil {
|
||||||
|
t.Fatalf("Scan failed: %v", err)
|
||||||
|
}
|
||||||
|
if u.Valid != tt.valid {
|
||||||
|
t.Errorf("expected valid=%v, got valid=%v", tt.valid, u.Valid)
|
||||||
|
}
|
||||||
|
if tt.valid && u.String() != tt.expected {
|
||||||
|
t.Errorf("expected %s, got %s", tt.expected, u.String())
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSqlUUID_Value(t *testing.T) {
|
||||||
|
testUUID := uuid.New()
|
||||||
|
u := NewSqlUUID(testUUID)
|
||||||
|
|
||||||
|
val, err := u.Value()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Value failed: %v", err)
|
||||||
|
}
|
||||||
|
// Value() should return a string for driver compatibility
|
||||||
|
if val != testUUID.String() {
|
||||||
|
t.Errorf("expected %s, got %s", testUUID.String(), val)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Test invalid UUID
|
||||||
|
u2 := SqlUUID{Valid: false}
|
||||||
|
val2, err := u2.Value()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Value failed: %v", err)
|
||||||
|
}
|
||||||
|
if val2 != nil {
|
||||||
|
t.Errorf("expected nil, got %v", val2)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSqlUUID_JSON(t *testing.T) {
|
||||||
|
testUUID := uuid.New()
|
||||||
|
u := NewSqlUUID(testUUID)
|
||||||
|
|
||||||
|
// Marshal
|
||||||
|
data, err := json.Marshal(u)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Marshal failed: %v", err)
|
||||||
|
}
|
||||||
|
expected := `"` + testUUID.String() + `"`
|
||||||
|
if string(data) != expected {
|
||||||
|
t.Errorf("expected %s, got %s", expected, string(data))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Unmarshal
|
||||||
|
var u2 SqlUUID
|
||||||
|
if err := json.Unmarshal([]byte(`"`+testUUID.String()+`"`), &u2); err != nil {
|
||||||
|
t.Fatalf("Unmarshal failed: %v", err)
|
||||||
|
}
|
||||||
|
if u2.String() != testUUID.String() {
|
||||||
|
t.Errorf("expected %s, got %s", testUUID.String(), u2.String())
|
||||||
|
}
|
||||||
|
|
||||||
|
// Test null
|
||||||
|
var u3 SqlUUID
|
||||||
|
if err := json.Unmarshal([]byte("null"), &u3); err != nil {
|
||||||
|
t.Fatalf("Unmarshal null failed: %v", err)
|
||||||
|
}
|
||||||
|
if u3.Valid {
|
||||||
|
t.Error("expected invalid UUID")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestTryIfInt64 tests the TryIfInt64 helper function
|
||||||
|
func TestTryIfInt64(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
input interface{}
|
||||||
|
def int64
|
||||||
|
expected int64
|
||||||
|
}{
|
||||||
|
{"string valid", "123", 0, 123},
|
||||||
|
{"string invalid", "abc", 99, 99},
|
||||||
|
{"int", 42, 0, 42},
|
||||||
|
{"int32", int32(100), 0, 100},
|
||||||
|
{"int64", int64(200), 0, 200},
|
||||||
|
{"uint32", uint32(50), 0, 50},
|
||||||
|
{"uint64", uint64(75), 0, 75},
|
||||||
|
{"float32", float32(3.14), 0, 3},
|
||||||
|
{"float64", float64(2.71), 0, 2},
|
||||||
|
{"bytes", []byte("456"), 0, 456},
|
||||||
|
{"unknown type", struct{}{}, 999, 999},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
result := TryIfInt64(tt.input, tt.def)
|
||||||
|
if result != tt.expected {
|
||||||
|
t.Errorf("expected %d, got %d", tt.expected, result)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestSqlString tests SqlString without base64 (plain text)
|
||||||
|
func TestSqlString_Scan(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
input interface{}
|
||||||
|
expected string
|
||||||
|
valid bool
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "plain string",
|
||||||
|
input: "hello world",
|
||||||
|
expected: "hello world",
|
||||||
|
valid: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "plain text",
|
||||||
|
input: "plain text",
|
||||||
|
expected: "plain text",
|
||||||
|
valid: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "bytes as string",
|
||||||
|
input: []byte("raw bytes"),
|
||||||
|
expected: "raw bytes",
|
||||||
|
valid: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "nil value",
|
||||||
|
input: nil,
|
||||||
|
expected: "",
|
||||||
|
valid: false,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
var s SqlString
|
||||||
|
if err := s.Scan(tt.input); err != nil {
|
||||||
|
t.Fatalf("Scan failed: %v", err)
|
||||||
|
}
|
||||||
|
if s.Valid != tt.valid {
|
||||||
|
t.Errorf("expected valid=%v, got valid=%v", tt.valid, s.Valid)
|
||||||
|
}
|
||||||
|
if tt.valid && s.String() != tt.expected {
|
||||||
|
t.Errorf("expected %q, got %q", tt.expected, s.String())
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSqlString_JSON(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
inputValue string
|
||||||
|
expectedJSON string
|
||||||
|
expectedDecode string
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "simple string",
|
||||||
|
inputValue: "hello world",
|
||||||
|
expectedJSON: `"hello world"`, // plain text, not base64
|
||||||
|
expectedDecode: "hello world",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "special characters",
|
||||||
|
inputValue: "test@#$%",
|
||||||
|
expectedJSON: `"test@#$%"`, // plain text, not base64
|
||||||
|
expectedDecode: "test@#$%",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "unicode string",
|
||||||
|
inputValue: "Hello 世界",
|
||||||
|
expectedJSON: `"Hello 世界"`, // plain text, not base64
|
||||||
|
expectedDecode: "Hello 世界",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "empty string",
|
||||||
|
inputValue: "",
|
||||||
|
expectedJSON: `""`,
|
||||||
|
expectedDecode: "",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
// Test MarshalJSON
|
||||||
|
s := NewSqlString(tt.inputValue)
|
||||||
|
data, err := json.Marshal(s)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Marshal failed: %v", err)
|
||||||
|
}
|
||||||
|
if string(data) != tt.expectedJSON {
|
||||||
|
t.Errorf("Marshal: expected %s, got %s", tt.expectedJSON, string(data))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Test UnmarshalJSON
|
||||||
|
var s2 SqlString
|
||||||
|
if err := json.Unmarshal(data, &s2); err != nil {
|
||||||
|
t.Fatalf("Unmarshal failed: %v", err)
|
||||||
|
}
|
||||||
|
if !s2.Valid {
|
||||||
|
t.Error("expected valid=true after unmarshal")
|
||||||
|
}
|
||||||
|
if s2.String() != tt.expectedDecode {
|
||||||
|
t.Errorf("Unmarshal: expected %q, got %q", tt.expectedDecode, s2.String())
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSqlString_JSON_Null(t *testing.T) {
|
||||||
|
// Test null handling
|
||||||
|
var s SqlString
|
||||||
|
if err := json.Unmarshal([]byte("null"), &s); err != nil {
|
||||||
|
t.Fatalf("Unmarshal null failed: %v", err)
|
||||||
|
}
|
||||||
|
if s.Valid {
|
||||||
|
t.Error("expected invalid after unmarshaling null")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Test marshal null
|
||||||
|
data, err := json.Marshal(s)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Marshal failed: %v", err)
|
||||||
|
}
|
||||||
|
if string(data) != "null" {
|
||||||
|
t.Errorf("expected null, got %s", string(data))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestSqlByteArray_Base64 tests SqlByteArray with base64 encoding/decoding
|
||||||
|
func TestSqlByteArray_Base64_Scan(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
input interface{}
|
||||||
|
expected []byte
|
||||||
|
valid bool
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "base64 encoded bytes from SQL",
|
||||||
|
input: "aGVsbG8gd29ybGQ=", // "hello world" in base64
|
||||||
|
expected: []byte("hello world"),
|
||||||
|
valid: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "plain bytes fallback",
|
||||||
|
input: "plain text",
|
||||||
|
expected: []byte("plain text"),
|
||||||
|
valid: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "bytes base64 encoded",
|
||||||
|
input: []byte("SGVsbG8gR29waGVy"), // "Hello Gopher" in base64
|
||||||
|
expected: []byte("Hello Gopher"),
|
||||||
|
valid: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "bytes plain fallback",
|
||||||
|
input: []byte("raw bytes"),
|
||||||
|
expected: []byte("raw bytes"),
|
||||||
|
valid: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "binary data",
|
||||||
|
input: "AQIDBA==", // []byte{1, 2, 3, 4} in base64
|
||||||
|
expected: []byte{1, 2, 3, 4},
|
||||||
|
valid: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "nil value",
|
||||||
|
input: nil,
|
||||||
|
expected: nil,
|
||||||
|
valid: false,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
var b SqlByteArray
|
||||||
|
if err := b.Scan(tt.input); err != nil {
|
||||||
|
t.Fatalf("Scan failed: %v", err)
|
||||||
|
}
|
||||||
|
if b.Valid != tt.valid {
|
||||||
|
t.Errorf("expected valid=%v, got valid=%v", tt.valid, b.Valid)
|
||||||
|
}
|
||||||
|
if tt.valid {
|
||||||
|
if string(b.Val) != string(tt.expected) {
|
||||||
|
t.Errorf("expected %q, got %q", tt.expected, b.Val)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSqlByteArray_Base64_JSON(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
inputValue []byte
|
||||||
|
expectedJSON string
|
||||||
|
expectedDecode []byte
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "text bytes",
|
||||||
|
inputValue: []byte("hello world"),
|
||||||
|
expectedJSON: `"aGVsbG8gd29ybGQ="`, // base64 encoded
|
||||||
|
expectedDecode: []byte("hello world"),
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "binary data",
|
||||||
|
inputValue: []byte{0x01, 0x02, 0x03, 0x04, 0xFF},
|
||||||
|
expectedJSON: `"AQIDBP8="`, // base64 encoded
|
||||||
|
expectedDecode: []byte{0x01, 0x02, 0x03, 0x04, 0xFF},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "empty bytes",
|
||||||
|
inputValue: []byte{},
|
||||||
|
expectedJSON: `""`, // base64 of empty bytes
|
||||||
|
expectedDecode: []byte{},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "unicode bytes",
|
||||||
|
inputValue: []byte("Hello 世界"),
|
||||||
|
expectedJSON: `"SGVsbG8g5LiW55WM"`, // base64 encoded
|
||||||
|
expectedDecode: []byte("Hello 世界"),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
// Test MarshalJSON
|
||||||
|
b := NewSqlByteArray(tt.inputValue)
|
||||||
|
data, err := json.Marshal(b)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Marshal failed: %v", err)
|
||||||
|
}
|
||||||
|
if string(data) != tt.expectedJSON {
|
||||||
|
t.Errorf("Marshal: expected %s, got %s", tt.expectedJSON, string(data))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Test UnmarshalJSON
|
||||||
|
var b2 SqlByteArray
|
||||||
|
if err := json.Unmarshal(data, &b2); err != nil {
|
||||||
|
t.Fatalf("Unmarshal failed: %v", err)
|
||||||
|
}
|
||||||
|
if !b2.Valid {
|
||||||
|
t.Error("expected valid=true after unmarshal")
|
||||||
|
}
|
||||||
|
if string(b2.Val) != string(tt.expectedDecode) {
|
||||||
|
t.Errorf("Unmarshal: expected %v, got %v", tt.expectedDecode, b2.Val)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSqlByteArray_Base64_JSON_Null(t *testing.T) {
|
||||||
|
// Test null handling
|
||||||
|
var b SqlByteArray
|
||||||
|
if err := json.Unmarshal([]byte("null"), &b); err != nil {
|
||||||
|
t.Fatalf("Unmarshal null failed: %v", err)
|
||||||
|
}
|
||||||
|
if b.Valid {
|
||||||
|
t.Error("expected invalid after unmarshaling null")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Test marshal null
|
||||||
|
data, err := json.Marshal(b)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Marshal failed: %v", err)
|
||||||
|
}
|
||||||
|
if string(data) != "null" {
|
||||||
|
t.Errorf("expected null, got %s", string(data))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSqlByteArray_Value(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
input SqlByteArray
|
||||||
|
expected interface{}
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "valid bytes",
|
||||||
|
input: NewSqlByteArray([]byte("test data")),
|
||||||
|
expected: []byte("test data"),
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "empty bytes",
|
||||||
|
input: NewSqlByteArray([]byte{}),
|
||||||
|
expected: []byte{},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "invalid",
|
||||||
|
input: SqlByteArray{Valid: false},
|
||||||
|
expected: nil,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
val, err := tt.input.Value()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Value failed: %v", err)
|
||||||
|
}
|
||||||
|
if tt.expected == nil && val != nil {
|
||||||
|
t.Errorf("expected nil, got %v", val)
|
||||||
|
}
|
||||||
|
if tt.expected != nil && val == nil {
|
||||||
|
t.Errorf("expected %v, got nil", tt.expected)
|
||||||
|
}
|
||||||
|
if tt.expected != nil && val != nil {
|
||||||
|
if string(val.([]byte)) != string(tt.expected.([]byte)) {
|
||||||
|
t.Errorf("expected %v, got %v", tt.expected, val)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestSqlString_RoundTrip tests complete round-trip: Go -> JSON -> Go -> SQL -> Go
|
||||||
|
func TestSqlString_RoundTrip(t *testing.T) {
|
||||||
|
original := "Test String with Special Chars: @#$%^&*()"
|
||||||
|
|
||||||
|
// Go -> JSON
|
||||||
|
s1 := NewSqlString(original)
|
||||||
|
jsonData, err := json.Marshal(s1)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Marshal failed: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// JSON -> Go
|
||||||
|
var s2 SqlString
|
||||||
|
if err := json.Unmarshal(jsonData, &s2); err != nil {
|
||||||
|
t.Fatalf("Unmarshal failed: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Go -> SQL (Value)
|
||||||
|
_, err = s2.Value()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Value failed: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// SQL -> Go (Scan plain text)
|
||||||
|
var s3 SqlString
|
||||||
|
// Simulate SQL driver returning plain text value
|
||||||
|
if err := s3.Scan(original); err != nil {
|
||||||
|
t.Fatalf("Scan failed: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Verify round-trip
|
||||||
|
if s3.String() != original {
|
||||||
|
t.Errorf("Round-trip failed: expected %q, got %q", original, s3.String())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestSqlByteArray_Base64_RoundTrip tests complete round-trip: Go -> JSON -> Go -> SQL -> Go
|
||||||
|
func TestSqlByteArray_Base64_RoundTrip(t *testing.T) {
|
||||||
|
original := []byte{0x48, 0x65, 0x6C, 0x6C, 0x6F, 0x20, 0xFF, 0xFE} // "Hello " + binary data
|
||||||
|
|
||||||
|
// Go -> JSON
|
||||||
|
b1 := NewSqlByteArray(original)
|
||||||
|
jsonData, err := json.Marshal(b1)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Marshal failed: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// JSON -> Go
|
||||||
|
var b2 SqlByteArray
|
||||||
|
if err := json.Unmarshal(jsonData, &b2); err != nil {
|
||||||
|
t.Fatalf("Unmarshal failed: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Go -> SQL (Value)
|
||||||
|
_, err = b2.Value()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Value failed: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// SQL -> Go (Scan with base64)
|
||||||
|
var b3 SqlByteArray
|
||||||
|
// Simulate SQL driver returning base64 encoded value
|
||||||
|
if err := b3.Scan("SGVsbG8g//4="); err != nil {
|
||||||
|
t.Fatalf("Scan failed: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Verify round-trip
|
||||||
|
if string(b3.Val) != string(original) {
|
||||||
|
t.Errorf("Round-trip failed: expected %v, got %v", original, b3.Val)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
@@ -0,0 +1,678 @@
|
|||||||
|
package sqltypes
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/xml"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/google/uuid"
|
||||||
|
"gopkg.in/yaml.v3"
|
||||||
|
)
|
||||||
|
|
||||||
|
// ── SqlNull: YAML ────────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
func TestSqlNull_YAML_Int(t *testing.T) {
|
||||||
|
n := NewSqlInt32(42)
|
||||||
|
data, err := yaml.Marshal(n)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Marshal failed: %v", err)
|
||||||
|
}
|
||||||
|
if string(data) != "42\n" {
|
||||||
|
t.Errorf("expected \"42\\n\", got %q", string(data))
|
||||||
|
}
|
||||||
|
|
||||||
|
var n2 SqlInt32
|
||||||
|
if err := yaml.Unmarshal(data, &n2); err != nil {
|
||||||
|
t.Fatalf("Unmarshal failed: %v", err)
|
||||||
|
}
|
||||||
|
if n2.Int64() != 42 {
|
||||||
|
t.Errorf("expected 42, got %d", n2.Int64())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSqlNull_YAML_Null(t *testing.T) {
|
||||||
|
var n SqlInt32
|
||||||
|
data, err := yaml.Marshal(n)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Marshal failed: %v", err)
|
||||||
|
}
|
||||||
|
if string(data) != "null\n" {
|
||||||
|
t.Errorf("expected \"null\\n\", got %q", string(data))
|
||||||
|
}
|
||||||
|
|
||||||
|
var n2 SqlInt32
|
||||||
|
if err := yaml.Unmarshal([]byte("null\n"), &n2); err != nil {
|
||||||
|
t.Fatalf("Unmarshal failed: %v", err)
|
||||||
|
}
|
||||||
|
if n2.Valid {
|
||||||
|
t.Error("expected invalid after unmarshaling null")
|
||||||
|
}
|
||||||
|
|
||||||
|
// ~ is also a YAML null.
|
||||||
|
var n3 SqlInt32
|
||||||
|
if err := yaml.Unmarshal([]byte("~\n"), &n3); err != nil {
|
||||||
|
t.Fatalf("Unmarshal failed: %v", err)
|
||||||
|
}
|
||||||
|
if n3.Valid {
|
||||||
|
t.Error("expected invalid after unmarshaling ~")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSqlNull_YAML_String(t *testing.T) {
|
||||||
|
s := NewSqlString("hello world")
|
||||||
|
data, err := yaml.Marshal(s)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Marshal failed: %v", err)
|
||||||
|
}
|
||||||
|
var s2 SqlString
|
||||||
|
if err := yaml.Unmarshal(data, &s2); err != nil {
|
||||||
|
t.Fatalf("Unmarshal failed: %v", err)
|
||||||
|
}
|
||||||
|
if s2.String() != "hello world" {
|
||||||
|
t.Errorf("expected %q, got %q", "hello world", s2.String())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSqlNull_YAML_UUID(t *testing.T) {
|
||||||
|
id := uuid.New()
|
||||||
|
u := NewSqlUUID(id)
|
||||||
|
data, err := yaml.Marshal(u)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Marshal failed: %v", err)
|
||||||
|
}
|
||||||
|
if string(data) != id.String()+"\n" {
|
||||||
|
t.Errorf("expected %q, got %q", id.String()+"\n", string(data))
|
||||||
|
}
|
||||||
|
var u2 SqlUUID
|
||||||
|
if err := yaml.Unmarshal(data, &u2); err != nil {
|
||||||
|
t.Fatalf("Unmarshal failed: %v", err)
|
||||||
|
}
|
||||||
|
if u2.UUID() != id {
|
||||||
|
t.Errorf("expected %v, got %v", id, u2.UUID())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSqlNull_YAML_ByteArray_Base64(t *testing.T) {
|
||||||
|
orig := []byte{0x01, 0x02, 0xFF}
|
||||||
|
b := NewSqlByteArray(orig)
|
||||||
|
data, err := yaml.Marshal(b)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Marshal failed: %v", err)
|
||||||
|
}
|
||||||
|
var b2 SqlByteArray
|
||||||
|
if err := yaml.Unmarshal(data, &b2); err != nil {
|
||||||
|
t.Fatalf("Unmarshal failed: %v", err)
|
||||||
|
}
|
||||||
|
if string(b2.Val) != string(orig) {
|
||||||
|
t.Errorf("expected %v, got %v", orig, b2.Val)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ── SqlNull: XML ─────────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
type xmlIntWrapper struct {
|
||||||
|
XMLName xml.Name `xml:"root"`
|
||||||
|
Value SqlInt32 `xml:"value"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSqlNull_XML_Int(t *testing.T) {
|
||||||
|
w := xmlIntWrapper{Value: NewSqlInt32(99)}
|
||||||
|
data, err := xml.Marshal(w)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Marshal failed: %v", err)
|
||||||
|
}
|
||||||
|
var w2 xmlIntWrapper
|
||||||
|
if err := xml.Unmarshal(data, &w2); err != nil {
|
||||||
|
t.Fatalf("Unmarshal failed: %v", err)
|
||||||
|
}
|
||||||
|
if w2.Value.Int64() != 99 {
|
||||||
|
t.Errorf("expected 99, got %d", w2.Value.Int64())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSqlNull_XML_Null(t *testing.T) {
|
||||||
|
w := xmlIntWrapper{}
|
||||||
|
data, err := xml.Marshal(w)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Marshal failed: %v", err)
|
||||||
|
}
|
||||||
|
var w2 xmlIntWrapper
|
||||||
|
if err := xml.Unmarshal(data, &w2); err != nil {
|
||||||
|
t.Fatalf("Unmarshal failed: %v", err)
|
||||||
|
}
|
||||||
|
if w2.Value.Valid {
|
||||||
|
t.Errorf("expected invalid, got %v", w2.Value)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
type xmlUUIDWrapper struct {
|
||||||
|
XMLName xml.Name `xml:"root"`
|
||||||
|
ID SqlUUID `xml:"id"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSqlNull_XML_UUID(t *testing.T) {
|
||||||
|
id := uuid.New()
|
||||||
|
w := xmlUUIDWrapper{ID: NewSqlUUID(id)}
|
||||||
|
data, err := xml.Marshal(w)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Marshal failed: %v", err)
|
||||||
|
}
|
||||||
|
var w2 xmlUUIDWrapper
|
||||||
|
if err := xml.Unmarshal(data, &w2); err != nil {
|
||||||
|
t.Fatalf("Unmarshal failed: %v", err)
|
||||||
|
}
|
||||||
|
if w2.ID.UUID() != id {
|
||||||
|
t.Errorf("expected %v, got %v", id, w2.ID.UUID())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
type xmlByteWrapper struct {
|
||||||
|
XMLName xml.Name `xml:"root"`
|
||||||
|
Data SqlByteArray `xml:"data"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSqlNull_XML_ByteArray_Base64(t *testing.T) {
|
||||||
|
orig := []byte{0xDE, 0xAD, 0xBE, 0xEF}
|
||||||
|
w := xmlByteWrapper{Data: NewSqlByteArray(orig)}
|
||||||
|
data, err := xml.Marshal(w)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Marshal failed: %v", err)
|
||||||
|
}
|
||||||
|
var w2 xmlByteWrapper
|
||||||
|
if err := xml.Unmarshal(data, &w2); err != nil {
|
||||||
|
t.Fatalf("Unmarshal failed: %v", err)
|
||||||
|
}
|
||||||
|
if string(w2.Data.Val) != string(orig) {
|
||||||
|
t.Errorf("expected %v, got %v", orig, w2.Data.Val)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ── SqlTimeStamp / SqlDate / SqlTime: YAML + XML ────────────────────────────
|
||||||
|
|
||||||
|
func TestSqlTimeStamp_YAML(t *testing.T) {
|
||||||
|
ts := NewSqlTimeStamp(time.Date(2024, 6, 15, 9, 30, 0, 0, time.UTC))
|
||||||
|
data, err := yaml.Marshal(ts)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Marshal failed: %v", err)
|
||||||
|
}
|
||||||
|
if string(data) != "2024-06-15T09:30:00\n" {
|
||||||
|
t.Errorf("unexpected YAML: %q", string(data))
|
||||||
|
}
|
||||||
|
var ts2 SqlTimeStamp
|
||||||
|
if err := yaml.Unmarshal(data, &ts2); err != nil {
|
||||||
|
t.Fatalf("Unmarshal failed: %v", err)
|
||||||
|
}
|
||||||
|
if ts2.Time().Format("2006-01-02T15:04:05") != "2024-06-15T09:30:00" {
|
||||||
|
t.Errorf("expected 2024-06-15T09:30:00, got %v", ts2.Time())
|
||||||
|
}
|
||||||
|
|
||||||
|
var zero SqlTimeStamp
|
||||||
|
zdata, err := yaml.Marshal(zero)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Marshal failed: %v", err)
|
||||||
|
}
|
||||||
|
if string(zdata) != "null\n" {
|
||||||
|
t.Errorf("expected null, got %q", zdata)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
type xmlTimeStampWrapper struct {
|
||||||
|
XMLName xml.Name `xml:"root"`
|
||||||
|
At SqlTimeStamp `xml:"at"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSqlTimeStamp_XML(t *testing.T) {
|
||||||
|
w := xmlTimeStampWrapper{At: NewSqlTimeStamp(time.Date(2024, 6, 15, 9, 30, 0, 0, time.UTC))}
|
||||||
|
data, err := xml.Marshal(w)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Marshal failed: %v", err)
|
||||||
|
}
|
||||||
|
var w2 xmlTimeStampWrapper
|
||||||
|
if err := xml.Unmarshal(data, &w2); err != nil {
|
||||||
|
t.Fatalf("Unmarshal failed: %v", err)
|
||||||
|
}
|
||||||
|
if w2.At.Time().Format("2006-01-02T15:04:05") != "2024-06-15T09:30:00" {
|
||||||
|
t.Errorf("expected 2024-06-15T09:30:00, got %v", w2.At.Time())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSqlDate_YAML(t *testing.T) {
|
||||||
|
d := NewSqlDate(time.Date(2024, 6, 15, 0, 0, 0, 0, time.UTC))
|
||||||
|
data, err := yaml.Marshal(d)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Marshal failed: %v", err)
|
||||||
|
}
|
||||||
|
if string(data) != "\"2024-06-15\"\n" {
|
||||||
|
t.Errorf("unexpected YAML: %q", string(data))
|
||||||
|
}
|
||||||
|
var d2 SqlDate
|
||||||
|
if err := yaml.Unmarshal(data, &d2); err != nil {
|
||||||
|
t.Fatalf("Unmarshal failed: %v", err)
|
||||||
|
}
|
||||||
|
if d2.String() != "2024-06-15" {
|
||||||
|
t.Errorf("expected 2024-06-15, got %q", d2.String())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
type xmlDateWrapper struct {
|
||||||
|
XMLName xml.Name `xml:"root"`
|
||||||
|
Day SqlDate `xml:"day"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSqlDate_XML(t *testing.T) {
|
||||||
|
w := xmlDateWrapper{Day: NewSqlDate(time.Date(2024, 6, 15, 0, 0, 0, 0, time.UTC))}
|
||||||
|
data, err := xml.Marshal(w)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Marshal failed: %v", err)
|
||||||
|
}
|
||||||
|
var w2 xmlDateWrapper
|
||||||
|
if err := xml.Unmarshal(data, &w2); err != nil {
|
||||||
|
t.Fatalf("Unmarshal failed: %v", err)
|
||||||
|
}
|
||||||
|
if w2.Day.String() != "2024-06-15" {
|
||||||
|
t.Errorf("expected 2024-06-15, got %q", w2.Day.String())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSqlTime_YAML(t *testing.T) {
|
||||||
|
tm := NewSqlTime(time.Date(0, 1, 1, 14, 5, 9, 0, time.UTC))
|
||||||
|
data, err := yaml.Marshal(tm)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Marshal failed: %v", err)
|
||||||
|
}
|
||||||
|
if string(data) != "\"14:05:09\"\n" {
|
||||||
|
t.Errorf("unexpected YAML: %q", string(data))
|
||||||
|
}
|
||||||
|
var tm2 SqlTime
|
||||||
|
if err := yaml.Unmarshal(data, &tm2); err != nil {
|
||||||
|
t.Fatalf("Unmarshal failed: %v", err)
|
||||||
|
}
|
||||||
|
if tm2.String() != "14:05:09" {
|
||||||
|
t.Errorf("expected 14:05:09, got %q", tm2.String())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
type xmlTimeWrapper struct {
|
||||||
|
XMLName xml.Name `xml:"root"`
|
||||||
|
At SqlTime `xml:"at"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSqlTime_XML(t *testing.T) {
|
||||||
|
w := xmlTimeWrapper{At: NewSqlTime(time.Date(0, 1, 1, 14, 5, 9, 0, time.UTC))}
|
||||||
|
data, err := xml.Marshal(w)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Marshal failed: %v", err)
|
||||||
|
}
|
||||||
|
var w2 xmlTimeWrapper
|
||||||
|
if err := xml.Unmarshal(data, &w2); err != nil {
|
||||||
|
t.Fatalf("Unmarshal failed: %v", err)
|
||||||
|
}
|
||||||
|
if w2.At.String() != "14:05:09" {
|
||||||
|
t.Errorf("expected 14:05:09, got %q", w2.At.String())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ── SqlJSONB: YAML + XML ─────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
func TestSqlJSONB_YAML(t *testing.T) {
|
||||||
|
j := SqlJSONB(`{"name":"test","count":42}`)
|
||||||
|
data, err := yaml.Marshal(j)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Marshal failed: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var decoded map[string]any
|
||||||
|
if err := yaml.Unmarshal(data, &decoded); err != nil {
|
||||||
|
t.Fatalf("failed decoding produced YAML: %v", err)
|
||||||
|
}
|
||||||
|
if decoded["name"] != "test" {
|
||||||
|
t.Errorf("expected name=test, got %v", decoded["name"])
|
||||||
|
}
|
||||||
|
|
||||||
|
var j2 SqlJSONB
|
||||||
|
if err := yaml.Unmarshal(data, &j2); err != nil {
|
||||||
|
t.Fatalf("Unmarshal failed: %v", err)
|
||||||
|
}
|
||||||
|
m, err := j2.AsMap()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("AsMap failed: %v", err)
|
||||||
|
}
|
||||||
|
if m["name"] != "test" {
|
||||||
|
t.Errorf("expected name=test, got %v", m["name"])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSqlJSONB_YAML_Null(t *testing.T) {
|
||||||
|
var j SqlJSONB
|
||||||
|
data, err := yaml.Marshal(j)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Marshal failed: %v", err)
|
||||||
|
}
|
||||||
|
if string(data) != "null\n" {
|
||||||
|
t.Errorf("expected null, got %q", data)
|
||||||
|
}
|
||||||
|
var j2 SqlJSONB
|
||||||
|
if err := yaml.Unmarshal([]byte("null\n"), &j2); err != nil {
|
||||||
|
t.Fatalf("Unmarshal failed: %v", err)
|
||||||
|
}
|
||||||
|
if j2 != nil {
|
||||||
|
t.Errorf("expected nil, got %v", j2)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
type xmlJSONBWrapper struct {
|
||||||
|
XMLName xml.Name `xml:"root"`
|
||||||
|
Meta SqlJSONB `xml:"meta"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSqlJSONB_XML(t *testing.T) {
|
||||||
|
w := xmlJSONBWrapper{Meta: SqlJSONB(`{"key":"value"}`)}
|
||||||
|
data, err := xml.Marshal(w)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Marshal failed: %v", err)
|
||||||
|
}
|
||||||
|
var w2 xmlJSONBWrapper
|
||||||
|
if err := xml.Unmarshal(data, &w2); err != nil {
|
||||||
|
t.Fatalf("Unmarshal failed: %v", err)
|
||||||
|
}
|
||||||
|
m, err := w2.Meta.AsMap()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("AsMap failed: %v", err)
|
||||||
|
}
|
||||||
|
if m["key"] != "value" {
|
||||||
|
t.Errorf("expected key=value, got %v", m)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ── Array types: YAML ────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
func TestSqlStringArray_YAML(t *testing.T) {
|
||||||
|
a := NewSqlStringArray([]string{"a", "b", "c"})
|
||||||
|
data, err := yaml.Marshal(a)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Marshal failed: %v", err)
|
||||||
|
}
|
||||||
|
var a2 SqlStringArray
|
||||||
|
if err := yaml.Unmarshal(data, &a2); err != nil {
|
||||||
|
t.Fatalf("Unmarshal failed: %v", err)
|
||||||
|
}
|
||||||
|
if len(a2.Val) != 3 || a2.Val[1] != "b" {
|
||||||
|
t.Errorf("unexpected value %v", a2.Val)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSqlStringArray_YAML_Null(t *testing.T) {
|
||||||
|
var a SqlStringArray
|
||||||
|
data, err := yaml.Marshal(a)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Marshal failed: %v", err)
|
||||||
|
}
|
||||||
|
if string(data) != "null\n" {
|
||||||
|
t.Errorf("expected null, got %q", data)
|
||||||
|
}
|
||||||
|
var a2 SqlStringArray
|
||||||
|
if err := yaml.Unmarshal([]byte("null\n"), &a2); err != nil {
|
||||||
|
t.Fatalf("Unmarshal failed: %v", err)
|
||||||
|
}
|
||||||
|
if a2.Valid {
|
||||||
|
t.Error("expected invalid")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSqlInt32Array_YAML(t *testing.T) {
|
||||||
|
a := NewSqlInt32Array([]int32{1, 2, 3})
|
||||||
|
data, err := yaml.Marshal(a)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Marshal failed: %v", err)
|
||||||
|
}
|
||||||
|
var a2 SqlInt32Array
|
||||||
|
if err := yaml.Unmarshal(data, &a2); err != nil {
|
||||||
|
t.Fatalf("Unmarshal failed: %v", err)
|
||||||
|
}
|
||||||
|
for i, v := range a.Val {
|
||||||
|
if a2.Val[i] != v {
|
||||||
|
t.Errorf("index %d: expected %d, got %d", i, v, a2.Val[i])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSqlUUIDArray_YAML(t *testing.T) {
|
||||||
|
ids := []uuid.UUID{uuid.New(), uuid.New()}
|
||||||
|
a := NewSqlUUIDArray(ids)
|
||||||
|
data, err := yaml.Marshal(a)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Marshal failed: %v", err)
|
||||||
|
}
|
||||||
|
var a2 SqlUUIDArray
|
||||||
|
if err := yaml.Unmarshal(data, &a2); err != nil {
|
||||||
|
t.Fatalf("Unmarshal failed: %v", err)
|
||||||
|
}
|
||||||
|
for i, v := range ids {
|
||||||
|
if a2.Val[i] != v {
|
||||||
|
t.Errorf("index %d: expected %v, got %v", i, v, a2.Val[i])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSqlVector_YAML(t *testing.T) {
|
||||||
|
v := NewSqlVector([]float32{0.1, 0.2, 0.3})
|
||||||
|
data, err := yaml.Marshal(v)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Marshal failed: %v", err)
|
||||||
|
}
|
||||||
|
var v2 SqlVector
|
||||||
|
if err := yaml.Unmarshal(data, &v2); err != nil {
|
||||||
|
t.Fatalf("Unmarshal failed: %v", err)
|
||||||
|
}
|
||||||
|
for i, f := range v.Val {
|
||||||
|
if v2.Val[i] != f {
|
||||||
|
t.Errorf("index %d: expected %v, got %v", i, f, v2.Val[i])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ── Array types: XML ─────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
type xmlStringArrayWrapper struct {
|
||||||
|
XMLName xml.Name `xml:"root"`
|
||||||
|
Tags SqlStringArray `xml:"tags"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSqlStringArray_XML(t *testing.T) {
|
||||||
|
w := xmlStringArrayWrapper{Tags: NewSqlStringArray([]string{"x", "y", "z"})}
|
||||||
|
data, err := xml.Marshal(w)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Marshal failed: %v", err)
|
||||||
|
}
|
||||||
|
var w2 xmlStringArrayWrapper
|
||||||
|
if err := xml.Unmarshal(data, &w2); err != nil {
|
||||||
|
t.Fatalf("Unmarshal failed: %v", err)
|
||||||
|
}
|
||||||
|
want := []string{"x", "y", "z"}
|
||||||
|
if len(w2.Tags.Val) != len(want) {
|
||||||
|
t.Fatalf("expected %v, got %v", want, w2.Tags.Val)
|
||||||
|
}
|
||||||
|
for i := range want {
|
||||||
|
if w2.Tags.Val[i] != want[i] {
|
||||||
|
t.Errorf("index %d: expected %q, got %q", i, want[i], w2.Tags.Val[i])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSqlStringArray_XML_Empty(t *testing.T) {
|
||||||
|
w := xmlStringArrayWrapper{}
|
||||||
|
data, err := xml.Marshal(w)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Marshal failed: %v", err)
|
||||||
|
}
|
||||||
|
var w2 xmlStringArrayWrapper
|
||||||
|
if err := xml.Unmarshal(data, &w2); err != nil {
|
||||||
|
t.Fatalf("Unmarshal failed: %v", err)
|
||||||
|
}
|
||||||
|
if len(w2.Tags.Val) != 0 {
|
||||||
|
t.Errorf("expected empty slice, got %v", w2.Tags.Val)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
type xmlInt64ArrayWrapper struct {
|
||||||
|
XMLName xml.Name `xml:"root"`
|
||||||
|
Scores SqlInt64Array `xml:"scores"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSqlInt64Array_XML(t *testing.T) {
|
||||||
|
w := xmlInt64ArrayWrapper{Scores: NewSqlInt64Array([]int64{10, 20, 30})}
|
||||||
|
data, err := xml.Marshal(w)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Marshal failed: %v", err)
|
||||||
|
}
|
||||||
|
var w2 xmlInt64ArrayWrapper
|
||||||
|
if err := xml.Unmarshal(data, &w2); err != nil {
|
||||||
|
t.Fatalf("Unmarshal failed: %v", err)
|
||||||
|
}
|
||||||
|
want := []int64{10, 20, 30}
|
||||||
|
for i := range want {
|
||||||
|
if w2.Scores.Val[i] != want[i] {
|
||||||
|
t.Errorf("index %d: expected %d, got %d", i, want[i], w2.Scores.Val[i])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
type xmlUUIDArrayWrapper struct {
|
||||||
|
XMLName xml.Name `xml:"root"`
|
||||||
|
IDs SqlUUIDArray `xml:"ids"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSqlUUIDArray_XML(t *testing.T) {
|
||||||
|
ids := []uuid.UUID{uuid.New(), uuid.New()}
|
||||||
|
w := xmlUUIDArrayWrapper{IDs: NewSqlUUIDArray(ids)}
|
||||||
|
data, err := xml.Marshal(w)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Marshal failed: %v", err)
|
||||||
|
}
|
||||||
|
var w2 xmlUUIDArrayWrapper
|
||||||
|
if err := xml.Unmarshal(data, &w2); err != nil {
|
||||||
|
t.Fatalf("Unmarshal failed: %v", err)
|
||||||
|
}
|
||||||
|
for i, id := range ids {
|
||||||
|
if w2.IDs.Val[i] != id {
|
||||||
|
t.Errorf("index %d: expected %v, got %v", i, id, w2.IDs.Val[i])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
type xmlVectorWrapper struct {
|
||||||
|
XMLName xml.Name `xml:"root"`
|
||||||
|
Embedding SqlVector `xml:"embedding"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSqlVector_XML(t *testing.T) {
|
||||||
|
w := xmlVectorWrapper{Embedding: NewSqlVector([]float32{1.5, -2.5})}
|
||||||
|
data, err := xml.Marshal(w)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Marshal failed: %v", err)
|
||||||
|
}
|
||||||
|
var w2 xmlVectorWrapper
|
||||||
|
if err := xml.Unmarshal(data, &w2); err != nil {
|
||||||
|
t.Fatalf("Unmarshal failed: %v", err)
|
||||||
|
}
|
||||||
|
want := []float32{1.5, -2.5}
|
||||||
|
for i := range want {
|
||||||
|
if w2.Embedding.Val[i] != want[i] {
|
||||||
|
t.Errorf("index %d: expected %v, got %v", i, want[i], w2.Embedding.Val[i])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ── Combined struct round-trip ───────────────────────────────────────────────
|
||||||
|
|
||||||
|
type yamlXMLRecord struct {
|
||||||
|
XMLName xml.Name `xml:"record" yaml:"-" json:"-"`
|
||||||
|
ID SqlUUID `xml:"id" yaml:"id"`
|
||||||
|
Name SqlString `xml:"name" yaml:"name"`
|
||||||
|
Age SqlInt32 `xml:"age" yaml:"age"`
|
||||||
|
Active SqlBool `xml:"active" yaml:"active"`
|
||||||
|
Created SqlTimeStamp `xml:"created" yaml:"created"`
|
||||||
|
Tags SqlStringArray `xml:"tags" yaml:"tags"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCombinedRecord_YAML_RoundTrip(t *testing.T) {
|
||||||
|
id := uuid.New()
|
||||||
|
original := yamlXMLRecord{
|
||||||
|
ID: NewSqlUUID(id),
|
||||||
|
Name: NewSqlString("Ada"),
|
||||||
|
Age: NewSqlInt32(36),
|
||||||
|
Active: NewSqlBool(true),
|
||||||
|
Created: NewSqlTimeStamp(time.Date(2024, 3, 10, 12, 30, 0, 0, time.UTC)),
|
||||||
|
Tags: NewSqlStringArray([]string{"engineer", "mathematician"}),
|
||||||
|
}
|
||||||
|
|
||||||
|
data, err := yaml.Marshal(original)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Marshal failed: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var decoded yamlXMLRecord
|
||||||
|
if err := yaml.Unmarshal(data, &decoded); err != nil {
|
||||||
|
t.Fatalf("Unmarshal failed: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if decoded.ID.UUID() != id {
|
||||||
|
t.Errorf("ID: expected %v, got %v", id, decoded.ID.UUID())
|
||||||
|
}
|
||||||
|
if decoded.Name.String() != "Ada" {
|
||||||
|
t.Errorf("Name: expected Ada, got %q", decoded.Name.String())
|
||||||
|
}
|
||||||
|
if decoded.Age.Int64() != 36 {
|
||||||
|
t.Errorf("Age: expected 36, got %d", decoded.Age.Int64())
|
||||||
|
}
|
||||||
|
if !decoded.Active.Bool() {
|
||||||
|
t.Error("Active: expected true")
|
||||||
|
}
|
||||||
|
if decoded.Created.Time().Format("2006-01-02T15:04:05") != "2024-03-10T12:30:00" {
|
||||||
|
t.Errorf("Created: expected 2024-03-10T12:30:00, got %v", decoded.Created.Time())
|
||||||
|
}
|
||||||
|
if len(decoded.Tags.Val) != 2 || decoded.Tags.Val[0] != "engineer" {
|
||||||
|
t.Errorf("Tags: unexpected value %v", decoded.Tags.Val)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCombinedRecord_XML_RoundTrip(t *testing.T) {
|
||||||
|
id := uuid.New()
|
||||||
|
original := yamlXMLRecord{
|
||||||
|
ID: NewSqlUUID(id),
|
||||||
|
Name: NewSqlString("Ada"),
|
||||||
|
Age: NewSqlInt32(36),
|
||||||
|
Active: NewSqlBool(true),
|
||||||
|
Created: NewSqlTimeStamp(time.Date(2024, 3, 10, 12, 30, 0, 0, time.UTC)),
|
||||||
|
Tags: NewSqlStringArray([]string{"engineer", "mathematician"}),
|
||||||
|
}
|
||||||
|
|
||||||
|
data, err := xml.Marshal(original)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Marshal failed: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var decoded yamlXMLRecord
|
||||||
|
if err := xml.Unmarshal(data, &decoded); err != nil {
|
||||||
|
t.Fatalf("Unmarshal failed: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if decoded.ID.UUID() != id {
|
||||||
|
t.Errorf("ID: expected %v, got %v", id, decoded.ID.UUID())
|
||||||
|
}
|
||||||
|
if decoded.Name.String() != "Ada" {
|
||||||
|
t.Errorf("Name: expected Ada, got %q", decoded.Name.String())
|
||||||
|
}
|
||||||
|
if decoded.Age.Int64() != 36 {
|
||||||
|
t.Errorf("Age: expected 36, got %d", decoded.Age.Int64())
|
||||||
|
}
|
||||||
|
if !decoded.Active.Bool() {
|
||||||
|
t.Error("Active: expected true")
|
||||||
|
}
|
||||||
|
if decoded.Created.Time().Format("2006-01-02T15:04:05") != "2024-03-10T12:30:00" {
|
||||||
|
t.Errorf("Created: expected 2024-03-10T12:30:00, got %v", decoded.Created.Time())
|
||||||
|
}
|
||||||
|
if len(decoded.Tags.Val) != 2 || decoded.Tags.Val[0] != "engineer" {
|
||||||
|
t.Errorf("Tags: unexpected value %v", decoded.Tags.Val)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,164 @@
|
|||||||
|
package sqltypes
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/google/uuid"
|
||||||
|
)
|
||||||
|
|
||||||
|
// record mimics a typical DB model composed of sqltypes fields, exercising
|
||||||
|
// marshalling/unmarshalling of the whole set together as encoding/json would
|
||||||
|
// when used on a real struct (not just the individual types in isolation).
|
||||||
|
type record struct {
|
||||||
|
ID SqlUUID `json:"id"`
|
||||||
|
Name SqlString `json:"name"`
|
||||||
|
Bio SqlString `json:"bio"`
|
||||||
|
Age SqlInt32 `json:"age"`
|
||||||
|
Score SqlFloat64 `json:"score"`
|
||||||
|
Active SqlBool `json:"active"`
|
||||||
|
CreatedAt SqlTimeStamp `json:"created_at"`
|
||||||
|
BirthDate SqlDate `json:"birth_date"`
|
||||||
|
Avatar SqlByteArray `json:"avatar"`
|
||||||
|
Tags SqlStringArray `json:"tags"`
|
||||||
|
Scores SqlInt32Array `json:"scores"`
|
||||||
|
Metadata SqlJSONB `json:"metadata"`
|
||||||
|
Embedding SqlVector `json:"embedding"`
|
||||||
|
Extras []uuid.UUID `json:"-"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestStruct_JSON_RoundTrip_AllFieldsPresent(t *testing.T) {
|
||||||
|
id := uuid.New()
|
||||||
|
createdAt := time.Date(2024, 3, 10, 12, 30, 0, 0, time.UTC)
|
||||||
|
birthDate := time.Date(1990, 5, 20, 0, 0, 0, 0, time.UTC)
|
||||||
|
|
||||||
|
original := record{
|
||||||
|
ID: NewSqlUUID(id),
|
||||||
|
Name: NewSqlString("Ada Lovelace"),
|
||||||
|
Bio: SqlString{}, // intentionally null
|
||||||
|
Age: NewSqlInt32(36),
|
||||||
|
Score: NewSqlFloat64(98.6),
|
||||||
|
Active: NewSqlBool(true),
|
||||||
|
CreatedAt: NewSqlTimeStamp(createdAt),
|
||||||
|
BirthDate: NewSqlDate(birthDate),
|
||||||
|
Avatar: NewSqlByteArray([]byte{0xDE, 0xAD, 0xBE, 0xEF}),
|
||||||
|
Tags: NewSqlStringArray([]string{"engineer", "mathematician"}),
|
||||||
|
Scores: NewSqlInt32Array([]int32{10, 20, 30}),
|
||||||
|
Metadata: SqlJSONB(`{"role":"admin"}`),
|
||||||
|
Embedding: NewSqlVector([]float32{0.1, 0.2, 0.3}),
|
||||||
|
}
|
||||||
|
|
||||||
|
data, err := json.Marshal(original)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Marshal failed: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var decoded record
|
||||||
|
if err := json.Unmarshal(data, &decoded); err != nil {
|
||||||
|
t.Fatalf("Unmarshal failed: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if decoded.ID.UUID() != id {
|
||||||
|
t.Errorf("ID: expected %v, got %v", id, decoded.ID.UUID())
|
||||||
|
}
|
||||||
|
if decoded.Name.String() != "Ada Lovelace" {
|
||||||
|
t.Errorf("Name: expected Ada Lovelace, got %q", decoded.Name.String())
|
||||||
|
}
|
||||||
|
if decoded.Bio.Valid {
|
||||||
|
t.Errorf("Bio: expected invalid/null, got %v", decoded.Bio)
|
||||||
|
}
|
||||||
|
if decoded.Age.Int64() != 36 {
|
||||||
|
t.Errorf("Age: expected 36, got %d", decoded.Age.Int64())
|
||||||
|
}
|
||||||
|
if decoded.Score.Float64() != 98.6 {
|
||||||
|
t.Errorf("Score: expected 98.6, got %v", decoded.Score.Float64())
|
||||||
|
}
|
||||||
|
if !decoded.Active.Bool() {
|
||||||
|
t.Errorf("Active: expected true")
|
||||||
|
}
|
||||||
|
if decoded.CreatedAt.Time().Format("2006-01-02T15:04:05") != createdAt.Format("2006-01-02T15:04:05") {
|
||||||
|
t.Errorf("CreatedAt: expected %v, got %v", createdAt, decoded.CreatedAt.Time())
|
||||||
|
}
|
||||||
|
if decoded.BirthDate.String() != "1990-05-20" {
|
||||||
|
t.Errorf("BirthDate: expected 1990-05-20, got %q", decoded.BirthDate.String())
|
||||||
|
}
|
||||||
|
if string(decoded.Avatar.Val) != string(original.Avatar.Val) {
|
||||||
|
t.Errorf("Avatar: expected %v, got %v", original.Avatar.Val, decoded.Avatar.Val)
|
||||||
|
}
|
||||||
|
if len(decoded.Tags.Val) != 2 || decoded.Tags.Val[0] != "engineer" {
|
||||||
|
t.Errorf("Tags: unexpected value %v", decoded.Tags.Val)
|
||||||
|
}
|
||||||
|
if len(decoded.Scores.Val) != 3 || decoded.Scores.Val[2] != 30 {
|
||||||
|
t.Errorf("Scores: unexpected value %v", decoded.Scores.Val)
|
||||||
|
}
|
||||||
|
m, err := decoded.Metadata.AsMap()
|
||||||
|
if err != nil || m["role"] != "admin" {
|
||||||
|
t.Errorf("Metadata: expected role=admin, got %v (err=%v)", m, err)
|
||||||
|
}
|
||||||
|
if len(decoded.Embedding.Val) != 3 || decoded.Embedding.Val[1] != 0.2 {
|
||||||
|
t.Errorf("Embedding: unexpected value %v", decoded.Embedding.Val)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestStruct_JSON_RoundTrip_AllNull(t *testing.T) {
|
||||||
|
var original record
|
||||||
|
|
||||||
|
data, err := json.Marshal(original)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Marshal failed: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var decoded record
|
||||||
|
if err := json.Unmarshal(data, &decoded); err != nil {
|
||||||
|
t.Fatalf("Unmarshal failed: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if decoded.ID.Valid || decoded.Name.Valid || decoded.Age.Valid || decoded.Score.Valid ||
|
||||||
|
decoded.Active.Valid || decoded.CreatedAt.Valid || decoded.BirthDate.Valid ||
|
||||||
|
decoded.Avatar.Valid || decoded.Tags.Valid || decoded.Scores.Valid || decoded.Embedding.Valid {
|
||||||
|
t.Errorf("expected all fields invalid/null after round-trip, got %+v", decoded)
|
||||||
|
}
|
||||||
|
if decoded.Metadata != nil {
|
||||||
|
t.Errorf("expected nil Metadata, got %v", decoded.Metadata)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestStruct_JSON_UnmarshalFromRawJSON(t *testing.T) {
|
||||||
|
raw := `{
|
||||||
|
"id": "3e843d4e-6b3c-4f2e-9a1e-6f0b2f3c9d10",
|
||||||
|
"name": "Grace Hopper",
|
||||||
|
"bio": null,
|
||||||
|
"age": 85,
|
||||||
|
"score": 100,
|
||||||
|
"active": false,
|
||||||
|
"created_at": "2023-12-01T08:00:00",
|
||||||
|
"birth_date": "1906-12-09",
|
||||||
|
"avatar": "AQIDBA==",
|
||||||
|
"tags": ["navy", "compiler"],
|
||||||
|
"scores": [1, 2, 3],
|
||||||
|
"metadata": {"key": "value"},
|
||||||
|
"embedding": [1.0, 2.0]
|
||||||
|
}`
|
||||||
|
|
||||||
|
var decoded record
|
||||||
|
if err := json.Unmarshal([]byte(raw), &decoded); err != nil {
|
||||||
|
t.Fatalf("Unmarshal failed: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if decoded.Name.String() != "Grace Hopper" {
|
||||||
|
t.Errorf("expected Grace Hopper, got %q", decoded.Name.String())
|
||||||
|
}
|
||||||
|
if decoded.Age.Int64() != 85 {
|
||||||
|
t.Errorf("expected age 85, got %d", decoded.Age.Int64())
|
||||||
|
}
|
||||||
|
if decoded.Active.Valid && decoded.Active.Bool() {
|
||||||
|
t.Errorf("expected active=false")
|
||||||
|
}
|
||||||
|
if string(decoded.Avatar.Val) != string([]byte{1, 2, 3, 4}) {
|
||||||
|
t.Errorf("expected avatar bytes [1 2 3 4], got %v", decoded.Avatar.Val)
|
||||||
|
}
|
||||||
|
if len(decoded.Tags.Val) != 2 || decoded.Tags.Val[1] != "compiler" {
|
||||||
|
t.Errorf("unexpected tags %v", decoded.Tags.Val)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,180 @@
|
|||||||
|
package sqltypes
|
||||||
|
|
||||||
|
import (
|
||||||
|
"database/sql"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/google/uuid"
|
||||||
|
_ "modernc.org/sqlite"
|
||||||
|
)
|
||||||
|
|
||||||
|
// TestUUIDWithRealDatabase tests that SqlUUID works with actual database operations
|
||||||
|
func TestUUIDWithRealDatabase(t *testing.T) {
|
||||||
|
// Open an in-memory SQLite database
|
||||||
|
db, err := sql.Open("sqlite", ":memory:")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Failed to open database: %v", err)
|
||||||
|
}
|
||||||
|
defer db.Close()
|
||||||
|
|
||||||
|
// Create a test table with UUID column
|
||||||
|
_, err = db.Exec(`
|
||||||
|
CREATE TABLE test_users (
|
||||||
|
id INTEGER PRIMARY KEY,
|
||||||
|
user_id TEXT,
|
||||||
|
name TEXT
|
||||||
|
)
|
||||||
|
`)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Failed to create table: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Test 1: Insert with UUID
|
||||||
|
testUUID1 := uuid.New()
|
||||||
|
sqlUUID1 := NewSqlUUID(testUUID1)
|
||||||
|
|
||||||
|
_, err = db.Exec("INSERT INTO test_users (id, user_id, name) VALUES (?, ?, ?)",
|
||||||
|
1, sqlUUID1, "Alice")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Failed to insert record: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Test 2: Update with UUID
|
||||||
|
testUUID2 := uuid.New()
|
||||||
|
sqlUUID2 := NewSqlUUID(testUUID2)
|
||||||
|
|
||||||
|
_, err = db.Exec("UPDATE test_users SET user_id = ? WHERE id = ?",
|
||||||
|
sqlUUID2, 1)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Failed to update record: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Test 3: Read back and verify
|
||||||
|
var retrievedID string
|
||||||
|
var name string
|
||||||
|
err = db.QueryRow("SELECT user_id, name FROM test_users WHERE id = ?", 1).Scan(&retrievedID, &name)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Failed to query record: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if retrievedID != testUUID2.String() {
|
||||||
|
t.Errorf("Expected UUID %s, got %s", testUUID2.String(), retrievedID)
|
||||||
|
}
|
||||||
|
|
||||||
|
if name != "Alice" {
|
||||||
|
t.Errorf("Expected name 'Alice', got '%s'", name)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Test 4: Insert with NULL UUID
|
||||||
|
nullUUID := SqlUUID{Valid: false}
|
||||||
|
_, err = db.Exec("INSERT INTO test_users (id, user_id, name) VALUES (?, ?, ?)",
|
||||||
|
2, nullUUID, "Bob")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Failed to insert record with NULL UUID: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Test 5: Read NULL UUID back
|
||||||
|
var retrievedNullID sql.NullString
|
||||||
|
err = db.QueryRow("SELECT user_id FROM test_users WHERE id = ?", 2).Scan(&retrievedNullID)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Failed to query NULL UUID record: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if retrievedNullID.Valid {
|
||||||
|
t.Errorf("Expected NULL UUID, got %s", retrievedNullID.String)
|
||||||
|
}
|
||||||
|
|
||||||
|
t.Logf("All database operations with UUID succeeded!")
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestUUIDValueReturnsString verifies that Value() returns string, not uuid.UUID
|
||||||
|
func TestUUIDValueReturnsString(t *testing.T) {
|
||||||
|
testUUID := uuid.New()
|
||||||
|
sqlUUID := NewSqlUUID(testUUID)
|
||||||
|
|
||||||
|
val, err := sqlUUID.Value()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Value() failed: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// The value should be a string, not a uuid.UUID
|
||||||
|
strVal, ok := val.(string)
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("Expected Value() to return string, got %T", val)
|
||||||
|
}
|
||||||
|
|
||||||
|
if strVal != testUUID.String() {
|
||||||
|
t.Errorf("Expected %s, got %s", testUUID.String(), strVal)
|
||||||
|
}
|
||||||
|
|
||||||
|
t.Logf("✓ Value() correctly returns string: %s", strVal)
|
||||||
|
}
|
||||||
|
|
||||||
|
// CustomStringableType is a custom type that implements fmt.Stringer
|
||||||
|
type CustomStringableType string
|
||||||
|
|
||||||
|
func (c CustomStringableType) String() string {
|
||||||
|
return "custom:" + string(c)
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestCustomStringableType verifies that any type implementing fmt.Stringer works
|
||||||
|
func TestCustomStringableType(t *testing.T) {
|
||||||
|
customVal := CustomStringableType("test-value")
|
||||||
|
sqlCustom := SqlNull[CustomStringableType]{
|
||||||
|
Val: customVal,
|
||||||
|
Valid: true,
|
||||||
|
}
|
||||||
|
|
||||||
|
val, err := sqlCustom.Value()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Value() failed: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Should return the result of String() method
|
||||||
|
strVal, ok := val.(string)
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("Expected Value() to return string, got %T", val)
|
||||||
|
}
|
||||||
|
|
||||||
|
expected := "custom:test-value"
|
||||||
|
if strVal != expected {
|
||||||
|
t.Errorf("Expected %s, got %s", expected, strVal)
|
||||||
|
}
|
||||||
|
|
||||||
|
t.Logf("✓ Custom Stringer type correctly converted to string: %s", strVal)
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestStringMethodUsesStringer verifies that String() method also uses fmt.Stringer
|
||||||
|
func TestStringMethodUsesStringer(t *testing.T) {
|
||||||
|
// Test with UUID
|
||||||
|
testUUID := uuid.New()
|
||||||
|
sqlUUID := NewSqlUUID(testUUID)
|
||||||
|
|
||||||
|
strResult := sqlUUID.String()
|
||||||
|
if strResult != testUUID.String() {
|
||||||
|
t.Errorf("Expected UUID String() to return %s, got %s", testUUID.String(), strResult)
|
||||||
|
}
|
||||||
|
t.Logf("✓ UUID String() method: %s", strResult)
|
||||||
|
|
||||||
|
// Test with custom Stringer type
|
||||||
|
customVal := CustomStringableType("test-value")
|
||||||
|
sqlCustom := SqlNull[CustomStringableType]{
|
||||||
|
Val: customVal,
|
||||||
|
Valid: true,
|
||||||
|
}
|
||||||
|
|
||||||
|
customStr := sqlCustom.String()
|
||||||
|
expected := "custom:test-value"
|
||||||
|
if customStr != expected {
|
||||||
|
t.Errorf("Expected custom String() to return %s, got %s", expected, customStr)
|
||||||
|
}
|
||||||
|
t.Logf("✓ Custom Stringer String() method: %s", customStr)
|
||||||
|
|
||||||
|
// Test with regular type (should use fmt.Sprintf)
|
||||||
|
sqlInt := NewSqlInt64(42)
|
||||||
|
intStr := sqlInt.String()
|
||||||
|
if intStr != "42" {
|
||||||
|
t.Errorf("Expected int String() to return '42', got '%s'", intStr)
|
||||||
|
}
|
||||||
|
t.Logf("✓ Regular type String() method: %s", intStr)
|
||||||
|
}
|
||||||
@@ -2,6 +2,7 @@ package ui
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"sort"
|
||||||
|
|
||||||
"github.com/rivo/tview"
|
"github.com/rivo/tview"
|
||||||
|
|
||||||
@@ -69,5 +70,6 @@ func getColumnNames(table *models.Table) []string {
|
|||||||
for name := range table.Columns {
|
for name := range table.Columns {
|
||||||
names = append(names, name)
|
names = append(names, name)
|
||||||
}
|
}
|
||||||
|
sort.Strings(names)
|
||||||
return names
|
return names
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,6 +1,10 @@
|
|||||||
package ui
|
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
|
// 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 {
|
for name := range table.Relationships {
|
||||||
names = append(names, name)
|
names = append(names, name)
|
||||||
}
|
}
|
||||||
|
sort.Strings(names)
|
||||||
return names
|
return names
|
||||||
}
|
}
|
||||||
|
|||||||
+118
-21
@@ -6,6 +6,8 @@ Generates Go source files with Bun model definitions from database schema inform
|
|||||||
|
|
||||||
The Bun Writer converts RelSpec's internal database model representation into Go source code with Bun struct definitions, complete with proper tags, relationships, and table configuration.
|
The Bun Writer converts RelSpec's internal database model representation into Go source code with Bun struct definitions, complete with proper tags, relationships, and table configuration.
|
||||||
|
|
||||||
|
With `--types sqltypes`, nullable fields use the [`pkg/sqltypes`](../../sqltypes/README.md) package.
|
||||||
|
|
||||||
## Features
|
## Features
|
||||||
|
|
||||||
- Generates Bun-compatible Go structs
|
- Generates Bun-compatible Go structs
|
||||||
@@ -46,45 +48,51 @@ func main() {
|
|||||||
### CLI Examples
|
### CLI Examples
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
# Generate Bun models from a DBML schema (default: resolvespec types)
|
# Generate Bun models from a DBML schema (default: baselib pointer types)
|
||||||
relspec convert --from dbml --from-path schema.dbml \
|
relspec convert --from dbml --from-path schema.dbml \
|
||||||
--to bun --to-path models.go --package models
|
--to bun --to-path models.go --package models
|
||||||
|
|
||||||
# Use standard library database/sql nullable types instead of resolvespec
|
# Use standard library database/sql nullable types instead
|
||||||
relspec convert --from dbml --from-path schema.dbml \
|
relspec convert --from dbml --from-path schema.dbml \
|
||||||
--to bun --to-path models.go --package models \
|
--to bun --to-path models.go --package models \
|
||||||
--types stdlib
|
--types stdlib
|
||||||
|
|
||||||
# Explicitly select resolvespec types (same as omitting --types)
|
# Select sqltypes package types (git.warky.dev/wdevs/relspecgo/pkg/sqltypes)
|
||||||
relspec convert --from pgsql --from-conn "postgres://localhost/mydb" \
|
relspec convert --from pgsql --from-conn "postgres://localhost/mydb" \
|
||||||
--to bun --to-path models.go --package models \
|
--to bun --to-path models.go --package models \
|
||||||
--types resolvespec
|
--types sqltypes
|
||||||
|
|
||||||
# Multi-file output (one file per table)
|
# Multi-file output (one file per table)
|
||||||
relspec convert --from json --from-path schema.json \
|
relspec convert --from json --from-path schema.json \
|
||||||
--to bun --to-path models/ --package models
|
--to bun --to-path models/ --package models
|
||||||
|
|
||||||
|
# Inject computed/scan-only fields that are not present in the database schema
|
||||||
|
# (extra-fields.json contains the JSON array shown below)
|
||||||
|
relspec convert --from dbml --from-path schema.dbml \
|
||||||
|
--to bun --to-path models.go --package models \
|
||||||
|
--extra-fields extra-fields.json
|
||||||
```
|
```
|
||||||
|
|
||||||
## Generated Code Examples
|
## Generated Code Examples
|
||||||
|
|
||||||
### Default — resolvespec types (`--types resolvespec`)
|
### sqltypes package types (`--types sqltypes`)
|
||||||
|
|
||||||
```go
|
```go
|
||||||
package models
|
package models
|
||||||
|
|
||||||
import (
|
import (
|
||||||
resolvespec_common "github.com/bitechdev/ResolveSpec/pkg/spectypes"
|
sql_types "git.warky.dev/wdevs/relspecgo/pkg/sqltypes"
|
||||||
"github.com/uptrace/bun"
|
"github.com/uptrace/bun"
|
||||||
)
|
)
|
||||||
|
|
||||||
type User struct {
|
type User struct {
|
||||||
bun.BaseModel `bun:"table:users,alias:u"`
|
bun.BaseModel `bun:"table:users,alias:u"`
|
||||||
|
|
||||||
ID int64 `bun:"id,type:uuid,pk," json:"id"`
|
ID int64 `bun:"id,type:uuid,pk," json:"id"`
|
||||||
Username string `bun:"username,type:text,notnull," json:"username"`
|
Username string `bun:"username,type:text,notnull," json:"username"`
|
||||||
Email resolvespec_common.SqlString `bun:"email,type:text,nullzero," json:"email"`
|
Email sql_types.SqlString `bun:"email,type:text,nullzero," json:"email"`
|
||||||
Tags resolvespec_common.SqlStringArray `bun:"tags,type:text[],default:'{}',notnull," json:"tags"`
|
Tags []string `bun:"tags,type:text[],array,default:'{}',notnull," json:"tags"`
|
||||||
CreatedAt resolvespec_common.SqlTimeStamp `bun:"created_at,type:timestamptz,default:now(),notnull," json:"created_at"`
|
CreatedAt sql_types.SqlTimeStamp `bun:"created_at,type:timestamptz,default:now(),notnull," json:"created_at"`
|
||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|
||||||
@@ -105,7 +113,7 @@ type User struct {
|
|||||||
ID string `bun:"id,type:uuid,pk," json:"id"`
|
ID string `bun:"id,type:uuid,pk," json:"id"`
|
||||||
Username string `bun:"username,type:text,notnull," json:"username"`
|
Username string `bun:"username,type:text,notnull," json:"username"`
|
||||||
Email sql.NullString `bun:"email,type:text,nullzero," json:"email"`
|
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"`
|
CreatedAt time.Time `bun:"created_at,type:timestamptz,default:now(),notnull," json:"created_at"`
|
||||||
}
|
}
|
||||||
```
|
```
|
||||||
@@ -126,7 +134,7 @@ type User struct {
|
|||||||
|
|
||||||
The nullable type package is selected with `--types` (or `WriterOptions.NullableTypes`).
|
The nullable type package is selected with `--types` (or `WriterOptions.NullableTypes`).
|
||||||
|
|
||||||
| SQL Type | NOT NULL (both) | Nullable — resolvespec | Nullable — stdlib |
|
| SQL Type | NOT NULL (both) | Nullable — sqltypes | Nullable — stdlib |
|
||||||
|---|---|---|---|
|
|---|---|---|---|
|
||||||
| `bigint` | `int64` | `SqlInt64` | `sql.NullInt64` |
|
| `bigint` | `int64` | `SqlInt64` | `sql.NullInt64` |
|
||||||
| `integer` | `int32` | `SqlInt32` | `sql.NullInt32` |
|
| `integer` | `int32` | `SqlInt32` | `sql.NullInt32` |
|
||||||
@@ -137,12 +145,18 @@ The nullable type package is selected with `--types` (or `WriterOptions.Nullable
|
|||||||
| `numeric`, `decimal` | `float64` | `SqlFloat64` | `sql.NullFloat64` |
|
| `numeric`, `decimal` | `float64` | `SqlFloat64` | `sql.NullFloat64` |
|
||||||
| `uuid` | `string` | `SqlUUID` | `sql.NullString` |
|
| `uuid` | `string` | `SqlUUID` | `sql.NullString` |
|
||||||
| `jsonb` | `string` | `SqlJSONB` | `sql.NullString` |
|
| `jsonb` | `string` | `SqlJSONB` | `sql.NullString` |
|
||||||
| `text[]` | `SqlStringArray` | `SqlStringArray` | `[]string` |
|
| `text[]` | `[]string`† | `[]string`† | `[]string`† |
|
||||||
| `integer[]` | `SqlInt32Array` | `SqlInt32Array` | `[]int32` |
|
| `integer[]` | `[]int32`† | `[]int32`† | `[]int32`† |
|
||||||
| `uuid[]` | `SqlUUIDArray` | `SqlUUIDArray` | `[]string` |
|
| `uuid[]` | `[]string`† | `[]string`† | `[]string`† |
|
||||||
| `vector` | `SqlVector` | `SqlVector` | `[]float32` |
|
| `vector` | `SqlVector` | `SqlVector` | `[]float32` |
|
||||||
|
|
||||||
\* In resolvespec 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`.
|
† 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
|
## Writer Options
|
||||||
|
|
||||||
@@ -151,11 +165,11 @@ The nullable type package is selected with `--types` (or `WriterOptions.Nullable
|
|||||||
Controls which Go package is used for nullable column types. Set via the `--types` CLI flag or `WriterOptions.NullableTypes`:
|
Controls which Go package is used for nullable column types. Set via the `--types` CLI flag or `WriterOptions.NullableTypes`:
|
||||||
|
|
||||||
```go
|
```go
|
||||||
// Use resolvespec types (default — omit NullableTypes or set to "resolvespec")
|
// Use sqltypes package types
|
||||||
options := &writers.WriterOptions{
|
options := &writers.WriterOptions{
|
||||||
OutputPath: "models.go",
|
OutputPath: "models.go",
|
||||||
PackageName: "models",
|
PackageName: "models",
|
||||||
NullableTypes: writers.NullableTypeResolveSpec,
|
NullableTypes: writers.NullableTypeSqlTypes,
|
||||||
}
|
}
|
||||||
|
|
||||||
// Use standard library database/sql types
|
// Use standard library database/sql types
|
||||||
@@ -166,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
|
### Metadata Options
|
||||||
|
|
||||||
```go
|
```go
|
||||||
@@ -180,12 +222,67 @@ options := &writers.WriterOptions{
|
|||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|
||||||
|
#### Extra fields
|
||||||
|
|
||||||
|
Use `Metadata["extra_fields"]` (or the CLI `--extra-fields <path>` flag) to inject
|
||||||
|
custom fields into generated Bun models without editing generated files. This is
|
||||||
|
intended for computed columns and scan-only query fields that Bun must scan into
|
||||||
|
but that are not real database table columns.
|
||||||
|
|
||||||
|
The CLI flag expects a path to a JSON file, not inline JSON. This avoids shell
|
||||||
|
quoting problems with struct tags and complex type names.
|
||||||
|
|
||||||
|
Each entry supports:
|
||||||
|
|
||||||
|
| Field | Required | Description |
|
||||||
|
|---|---:|---|
|
||||||
|
| `target_table` | No | Optional model scope. Accepts table name (`projects`), qualified table (`public.projects`), or model name (`ModelPublicProjects`). If omitted, the field is added to every generated model. |
|
||||||
|
| `name` | Yes | Go struct field name. |
|
||||||
|
| `type` | Yes | Go type to emit. |
|
||||||
|
| `bun_tag` | No | Raw Bun tag contents, e.g. `thought_count,scanonly`. |
|
||||||
|
| `json_tag` | No | Raw JSON tag contents. |
|
||||||
|
| `comment` | No | Optional line comment. |
|
||||||
|
|
||||||
|
```go
|
||||||
|
options := &writers.WriterOptions{
|
||||||
|
OutputPath: "models.go",
|
||||||
|
PackageName: "models",
|
||||||
|
Metadata: map[string]any{
|
||||||
|
"extra_fields": `[
|
||||||
|
{
|
||||||
|
"target_table": "projects",
|
||||||
|
"name": "ThoughtCount",
|
||||||
|
"type": "sql_types.SqlInt64",
|
||||||
|
"bun_tag": "thought_count,scanonly",
|
||||||
|
"json_tag": "thought_count",
|
||||||
|
"comment": "Computed by ResolveSpec queries"
|
||||||
|
}
|
||||||
|
]`,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
Example `extra-fields.json`:
|
||||||
|
|
||||||
|
```json
|
||||||
|
[
|
||||||
|
{
|
||||||
|
"target_table": "projects",
|
||||||
|
"name": "ThoughtCount",
|
||||||
|
"type": "sql_types.SqlInt64",
|
||||||
|
"bun_tag": "thought_count,scanonly",
|
||||||
|
"json_tag": "thought_count",
|
||||||
|
"comment": "Computed by ResolveSpec queries"
|
||||||
|
}
|
||||||
|
]
|
||||||
|
```
|
||||||
|
|
||||||
## Notes
|
## Notes
|
||||||
|
|
||||||
- Model names are derived from table names (singularized, PascalCase)
|
- Model names are derived from table names (singularized, PascalCase)
|
||||||
- Table aliases are auto-generated from table names
|
- Table aliases are auto-generated from table names
|
||||||
- Nullable columns use `resolvespec_common.SqlString`, `resolvespec_common.SqlTimeStamp`, etc. by default; pass `--types stdlib` to use `sql.NullString`, `sql.NullTime`, etc. instead
|
- 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 `resolvespec_common.SqlStringArray`, `resolvespec_common.SqlInt32Array`, etc. by default; `--types stdlib` produces plain Go slices (`[]string`, `[]int32`, …)
|
- 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`
|
- Multi-file mode: one file per table named `sql_{schema}_{table}.go`
|
||||||
- Generated code is auto-formatted
|
- Generated code is auto-formatted
|
||||||
- JSON tags are automatically added
|
- 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)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
@@ -1,6 +1,8 @@
|
|||||||
package bun
|
package bun
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
"sort"
|
"sort"
|
||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
@@ -32,6 +34,90 @@ type ModelData struct {
|
|||||||
PrimaryKeyIDType string // Helper method GetID/SetID/UpdateID type
|
PrimaryKeyIDType string // Helper method GetID/SetID/UpdateID type
|
||||||
IDColumnName string // Name of the ID column in database
|
IDColumnName string // Name of the ID column in database
|
||||||
Prefix string // 3-letter prefix
|
Prefix string // 3-letter prefix
|
||||||
|
|
||||||
|
// ExtraFields are user-defined fields added via generator config (issue #4)
|
||||||
|
ExtraFields []*FieldData `json:"extra_fields,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// ExtraFieldConfig represents a custom field to be added to generated models
|
||||||
|
type ExtraFieldConfig struct {
|
||||||
|
Name string `json:"name"`
|
||||||
|
Type string `json:"type"`
|
||||||
|
BunTag string `json:"bun_tag,omitempty"`
|
||||||
|
JSONTag string `json:"json_tag,omitempty"`
|
||||||
|
Comment string `json:"comment,omitempty"`
|
||||||
|
TargetTable string `json:"target_table,omitempty"` // optional: if set, field only applies to this table
|
||||||
|
}
|
||||||
|
|
||||||
|
// LoadExtraFieldsFromMetadata loads custom field configurations from metadata map.
|
||||||
|
// Accepts either a JSON-encoded array of fields (for CLI use) or a structured list
|
||||||
|
// (for programmatic/sidecar file use). Each entry must have "name" and "type";
|
||||||
|
// an optional "target_table" scopes the field to one table only.
|
||||||
|
func LoadExtraFieldsFromMetadata(metadata map[string]interface{}) []ExtraFieldConfig {
|
||||||
|
if metadata == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
extraFieldsRaw, ok := metadata["extra_fields"]
|
||||||
|
if !ok || extraFieldsRaw == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Try to parse as JSON string first (from CLI flag)
|
||||||
|
if strVal, ok := extraFieldsRaw.(string); ok {
|
||||||
|
var fields []ExtraFieldConfig
|
||||||
|
if err := json.Unmarshal([]byte(strVal), &fields); err == nil && len(fields) > 0 {
|
||||||
|
return filterValidFields(fields)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Try to parse as structured data (from sidecar file or programmatic API)
|
||||||
|
extraFieldsList, ok := extraFieldsRaw.([]interface{})
|
||||||
|
if !ok || len(extraFieldsList) == 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
fields := make([]ExtraFieldConfig, 0, len(extraFieldsList))
|
||||||
|
for _, item := range extraFieldsList {
|
||||||
|
if fieldMap, ok := item.(map[string]interface{}); ok {
|
||||||
|
field := ExtraFieldConfig{
|
||||||
|
Name: toString(fieldMap["name"]),
|
||||||
|
Type: toString(fieldMap["type"]),
|
||||||
|
BunTag: toString(fieldMap["bun_tag"]),
|
||||||
|
JSONTag: toString(fieldMap["json_tag"]),
|
||||||
|
Comment: toString(fieldMap["comment"]),
|
||||||
|
TargetTable: toString(fieldMap["target_table"]),
|
||||||
|
}
|
||||||
|
if field.Name != "" && field.Type != "" {
|
||||||
|
fields = append(fields, field)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return filterValidFields(fields)
|
||||||
|
}
|
||||||
|
|
||||||
|
// filterValidFields removes entries that are missing required fields.
|
||||||
|
func filterValidFields(fields []ExtraFieldConfig) []ExtraFieldConfig {
|
||||||
|
var valid []ExtraFieldConfig
|
||||||
|
for _, f := range fields {
|
||||||
|
if f.Name != "" && f.Type != "" {
|
||||||
|
valid = append(valid, f)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return valid
|
||||||
|
}
|
||||||
|
|
||||||
|
func toString(v interface{}) string {
|
||||||
|
if v == nil {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
switch val := v.(type) {
|
||||||
|
case string:
|
||||||
|
return val
|
||||||
|
default:
|
||||||
|
return fmt.Sprintf("%v", val)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// FieldData represents a single field in a struct
|
// FieldData represents a single field in a struct
|
||||||
@@ -134,17 +220,20 @@ func NewModelData(table *models.Table, schema string, typeMapper *TypeMapper, fl
|
|||||||
Prefix: GeneratePrefix(table.Name),
|
Prefix: GeneratePrefix(table.Name),
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Convert columns to fields (sorted by sequence or name)
|
||||||
|
columns := sortColumns(table.Columns)
|
||||||
|
|
||||||
// Find primary key
|
// Find primary key
|
||||||
for _, col := range table.Columns {
|
for _, col := range columns {
|
||||||
if col.IsPrimaryKey {
|
if col.IsPrimaryKey {
|
||||||
// Sanitize column name to remove backticks
|
// Sanitize column name to remove backticks
|
||||||
safeName := writers.SanitizeStructTagValue(col.Name)
|
safeName := writers.SanitizeStructTagValue(col.Name)
|
||||||
model.PrimaryKeyField = SnakeCaseToPascalCase(safeName)
|
model.PrimaryKeyField = SnakeCaseToPascalCase(safeName)
|
||||||
model.IDColumnName = safeName
|
model.IDColumnName = safeName
|
||||||
// Check if PK type is a SQL type (contains resolvespec_common or sql_types)
|
// Check if PK type is a SQL type (contains sql_types)
|
||||||
goType := typeMapper.SQLTypeToGoType(col.Type, col.NotNull)
|
goType := typeMapper.SQLTypeToGoType(col.Type, col.NotNull)
|
||||||
model.PrimaryKeyType = goType
|
model.PrimaryKeyType = goType
|
||||||
model.PrimaryKeyIsSQL = strings.Contains(goType, "resolvespec_common") || strings.Contains(goType, "sql_types")
|
model.PrimaryKeyIsSQL = strings.Contains(goType, "sql_types")
|
||||||
model.PrimaryKeyIsStr = isStringLikePrimaryKeyType(goType)
|
model.PrimaryKeyIsStr = isStringLikePrimaryKeyType(goType)
|
||||||
model.PrimaryKeyIDType = "int64"
|
model.PrimaryKeyIDType = "int64"
|
||||||
if model.PrimaryKeyIsStr {
|
if model.PrimaryKeyIsStr {
|
||||||
@@ -154,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 {
|
for _, col := range columns {
|
||||||
field := columnToField(col, table, typeMapper)
|
field := columnToField(col, table, typeMapper)
|
||||||
// Check for name collision with generated methods and rename if needed
|
// Check for name collision with generated methods and rename if needed
|
||||||
@@ -203,7 +290,7 @@ func formatComment(description, comment string) string {
|
|||||||
|
|
||||||
func isStringLikePrimaryKeyType(goType string) bool {
|
func isStringLikePrimaryKeyType(goType string) bool {
|
||||||
switch goType {
|
switch goType {
|
||||||
case "string", "sql.NullString", "resolvespec_common.SqlString", "resolvespec_common.SqlUUID":
|
case "string", "sql.NullString", "sql_types.SqlString", "sql_types.SqlUUID":
|
||||||
return true
|
return true
|
||||||
default:
|
default:
|
||||||
return false
|
return false
|
||||||
@@ -249,6 +336,21 @@ func sortConstraints(constraints map[string]*models.Constraint) []*models.Constr
|
|||||||
return result
|
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
|
// sortColumns sorts columns by sequence, then by name
|
||||||
func sortColumns(columns map[string]*models.Column) []*models.Column {
|
func sortColumns(columns map[string]*models.Column) []*models.Column {
|
||||||
result := make([]*models.Column, 0, len(columns))
|
result := make([]*models.Column, 0, len(columns))
|
||||||
|
|||||||
@@ -23,6 +23,13 @@ type {{.Name}} struct {
|
|||||||
{{- range .Fields}}
|
{{- range .Fields}}
|
||||||
{{.Name}} {{.Type}} ` + "`bun:\"{{.BunTag}}\" json:\"{{.JSONTag}}\"`" + `{{if .Comment}} // {{.Comment}}{{end}}
|
{{.Name}} {{.Type}} ` + "`bun:\"{{.BunTag}}\" json:\"{{.JSONTag}}\"`" + `{{if .Comment}} // {{.Comment}}{{end}}
|
||||||
{{- end}}
|
{{- end}}
|
||||||
|
{{- if .ExtraFields}}
|
||||||
|
|
||||||
|
// --- Custom fields (managed by relspecgo generator config, issue #4) ---
|
||||||
|
{{- range .ExtraFields}}
|
||||||
|
{{.Name}} {{.Type}} ` + "`bun:\"{{.BunTag}}\" json:\"{{.JSONTag}}\"`" + `{{if .Comment}} // {{.Comment}}{{end}}
|
||||||
|
{{- end}}
|
||||||
|
{{- end}}
|
||||||
}
|
}
|
||||||
{{if .Config.GenerateTableName}}
|
{{if .Config.GenerateTableName}}
|
||||||
// TableName returns the table name for {{.Name}}
|
// TableName returns the table name for {{.Name}}
|
||||||
|
|||||||
@@ -12,39 +12,56 @@ import (
|
|||||||
// TypeMapper handles type conversions between SQL and Go types for Bun
|
// TypeMapper handles type conversions between SQL and Go types for Bun
|
||||||
type TypeMapper struct {
|
type TypeMapper struct {
|
||||||
sqlTypesAlias string
|
sqlTypesAlias string
|
||||||
typeStyle string // writers.NullableTypeResolveSpec | writers.NullableTypeStdlib
|
typeStyle string // writers.NullableTypeSqlTypes | writers.NullableTypeStdlib | writers.NullableTypeBaselib
|
||||||
|
arrayNullable string // writers.NullableArraysSlice | writers.NullableArraysPointerSlice
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewTypeMapper creates a new TypeMapper.
|
// NewTypeMapper creates a new TypeMapper.
|
||||||
// typeStyle should be writers.NullableTypeResolveSpec or writers.NullableTypeStdlib;
|
// typeStyle should be writers.NullableTypeSqlTypes, writers.NullableTypeStdlib, or
|
||||||
// an empty string defaults to resolvespec.
|
// 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 == "" {
|
if typeStyle == "" {
|
||||||
typeStyle = writers.NullableTypeResolveSpec
|
typeStyle = writers.NullableTypeBaselib
|
||||||
|
}
|
||||||
|
if arrayNullable == "" {
|
||||||
|
arrayNullable = writers.NullableArraysSlice
|
||||||
}
|
}
|
||||||
return &TypeMapper{
|
return &TypeMapper{
|
||||||
sqlTypesAlias: "resolvespec_common",
|
sqlTypesAlias: "sql_types",
|
||||||
typeStyle: typeStyle,
|
typeStyle: typeStyle,
|
||||||
|
arrayNullable: arrayNullable,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// SQLTypeToGoType converts a SQL type to its Go equivalent.
|
// SQLTypeToGoType converts a SQL type to its Go equivalent.
|
||||||
func (tm *TypeMapper) SQLTypeToGoType(sqlType string, notNull bool) string {
|
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) {
|
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)
|
baseType := tm.extractBaseType(sqlType)
|
||||||
|
|
||||||
if tm.typeStyle == writers.NullableTypeStdlib {
|
switch tm.typeStyle {
|
||||||
|
case writers.NullableTypeStdlib:
|
||||||
if notNull {
|
if notNull {
|
||||||
return tm.rawGoType(baseType)
|
return tm.rawGoType(baseType)
|
||||||
}
|
}
|
||||||
return tm.stdlibNullableGoType(baseType)
|
return tm.stdlibNullableGoType(baseType)
|
||||||
|
case writers.NullableTypeBaselib:
|
||||||
|
if notNull {
|
||||||
|
return tm.rawGoType(baseType)
|
||||||
|
}
|
||||||
|
return tm.baselibNullableGoType(baseType)
|
||||||
}
|
}
|
||||||
|
|
||||||
// resolvespec (default): use base Go types only for simple NOT NULL fields.
|
// sqltypes: use base Go types only for simple NOT NULL fields.
|
||||||
if notNull && tm.isSimpleType(baseType) {
|
if notNull && tm.isSimpleType(baseType) {
|
||||||
return tm.baseGoType(baseType)
|
return tm.baseGoType(baseType)
|
||||||
}
|
}
|
||||||
@@ -100,11 +117,11 @@ func (tm *TypeMapper) baseGoType(sqlType string) string {
|
|||||||
return goType
|
return goType
|
||||||
}
|
}
|
||||||
|
|
||||||
// Default to resolvespec type
|
// Default to sqltypes type
|
||||||
return tm.bunGoType(sqlType)
|
return tm.bunGoType(sqlType)
|
||||||
}
|
}
|
||||||
|
|
||||||
// bunGoType returns the Bun/ResolveSpec common type
|
// bunGoType returns the Bun/sqltypes common type
|
||||||
func (tm *TypeMapper) bunGoType(sqlType string) string {
|
func (tm *TypeMapper) bunGoType(sqlType string) string {
|
||||||
typeMap := map[string]string{
|
typeMap := map[string]string{
|
||||||
// Integer types
|
// Integer types
|
||||||
@@ -182,34 +199,13 @@ func (tm *TypeMapper) bunGoType(sqlType string) string {
|
|||||||
|
|
||||||
// arrayGoType returns the Go type for a PostgreSQL array column.
|
// arrayGoType returns the Go type for a PostgreSQL array column.
|
||||||
// The baseElemType is the canonical base type (e.g. "text", "integer").
|
// 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 {
|
func (tm *TypeMapper) arrayGoType(baseElemType string) string {
|
||||||
if tm.typeStyle == writers.NullableTypeStdlib {
|
return tm.stdlibArrayGoType(baseElemType)
|
||||||
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"
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// rawGoType returns the plain Go type for a NOT NULL column in stdlib mode.
|
// rawGoType returns the plain Go type for a NOT NULL column in stdlib mode.
|
||||||
@@ -276,6 +272,38 @@ func (tm *TypeMapper) stdlibNullableGoType(sqlType string) string {
|
|||||||
return "sql.NullString"
|
return "sql.NullString"
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// baselibNullableGoType returns plain Go pointer types for nullable columns.
|
||||||
|
func (tm *TypeMapper) baselibNullableGoType(sqlType string) string {
|
||||||
|
typeMap := map[string]string{
|
||||||
|
"integer": "*int32", "int": "*int32", "int4": "*int32", "serial": "*int32",
|
||||||
|
"smallint": "*int16", "int2": "*int16", "smallserial": "*int16",
|
||||||
|
"bigint": "*int64", "int8": "*int64", "bigserial": "*int64",
|
||||||
|
"boolean": "*bool", "bool": "*bool",
|
||||||
|
"real": "*float32", "float4": "*float32",
|
||||||
|
"double precision": "*float64", "float8": "*float64",
|
||||||
|
"numeric": "*float64", "decimal": "*float64", "money": "*float64",
|
||||||
|
"text": "*string", "varchar": "*string", "char": "*string",
|
||||||
|
"character": "*string", "citext": "*string", "bpchar": "*string",
|
||||||
|
"inet": "*string", "cidr": "*string", "macaddr": "*string",
|
||||||
|
"uuid": "*string", "json": "*string", "jsonb": "*string",
|
||||||
|
"timestamp": "*time.Time",
|
||||||
|
"timestamp without time zone": "*time.Time",
|
||||||
|
"timestamp with time zone": "*time.Time",
|
||||||
|
"timestamptz": "*time.Time",
|
||||||
|
"date": "*time.Time",
|
||||||
|
"time": "*time.Time",
|
||||||
|
"time without time zone": "*time.Time",
|
||||||
|
"time with time zone": "*time.Time",
|
||||||
|
"timetz": "*time.Time",
|
||||||
|
"bytea": "[]byte",
|
||||||
|
"vector": "[]float32",
|
||||||
|
}
|
||||||
|
if goType, ok := typeMap[sqlType]; ok {
|
||||||
|
return goType
|
||||||
|
}
|
||||||
|
return "*string"
|
||||||
|
}
|
||||||
|
|
||||||
// stdlibArrayGoType returns a plain Go slice type for array columns in stdlib mode.
|
// stdlibArrayGoType returns a plain Go slice type for array columns in stdlib mode.
|
||||||
func (tm *TypeMapper) stdlibArrayGoType(baseElemType string) string {
|
func (tm *TypeMapper) stdlibArrayGoType(baseElemType string) string {
|
||||||
typeMap := map[string]string{
|
typeMap := map[string]string{
|
||||||
@@ -323,7 +351,7 @@ func (tm *TypeMapper) BuildBunTag(column *models.Column, table *models.Table) st
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
parts = append(parts, fmt.Sprintf("type:%s", typeStr))
|
parts = append(parts, fmt.Sprintf("type:%s", typeStr))
|
||||||
if isArray && tm.typeStyle == writers.NullableTypeStdlib {
|
if isArray {
|
||||||
parts = append(parts, "array")
|
parts = append(parts, "array")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -355,7 +383,7 @@ func (tm *TypeMapper) BuildBunTag(column *models.Column, table *models.Table) st
|
|||||||
|
|
||||||
// Check for indexes (unique indexes should be added to tag)
|
// Check for indexes (unique indexes should be added to tag)
|
||||||
if table != nil {
|
if table != nil {
|
||||||
for _, index := range table.Indexes {
|
for _, index := range sortIndexes(table.Indexes) {
|
||||||
if !index.Unique {
|
if !index.Unique {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
@@ -423,16 +451,19 @@ func (tm *TypeMapper) NeedsFmtImport(generateGetIDStr bool) bool {
|
|||||||
return generateGetIDStr
|
return generateGetIDStr
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetSQLTypesImport returns the import path for the ResolveSpec spectypes package.
|
// GetSQLTypesImport returns the import path for the sqltypes package.
|
||||||
func (tm *TypeMapper) GetSQLTypesImport() string {
|
func (tm *TypeMapper) GetSQLTypesImport() string {
|
||||||
return "github.com/bitechdev/ResolveSpec/pkg/spectypes"
|
return "git.warky.dev/wdevs/relspecgo/pkg/sqltypes"
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetNullableTypeImportLine returns the full Go import line for the nullable type
|
// GetNullableTypeImportLine returns the full Go import line for the nullable type
|
||||||
// package (ready to pass to AddImport). Returns empty string when no import is needed.
|
// package (ready to pass to AddImport). Returns empty string when no import is needed.
|
||||||
func (tm *TypeMapper) GetNullableTypeImportLine() string {
|
func (tm *TypeMapper) GetNullableTypeImportLine() string {
|
||||||
if tm.typeStyle == writers.NullableTypeStdlib {
|
switch tm.typeStyle {
|
||||||
|
case writers.NullableTypeStdlib:
|
||||||
return "\"database/sql\""
|
return "\"database/sql\""
|
||||||
|
case writers.NullableTypeBaselib:
|
||||||
|
return ""
|
||||||
}
|
}
|
||||||
return fmt.Sprintf("%s \"%s\"", tm.sqlTypesAlias, tm.GetSQLTypesImport())
|
return fmt.Sprintf("%s \"%s\"", tm.sqlTypesAlias, tm.GetSQLTypesImport())
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -24,7 +24,7 @@ type Writer struct {
|
|||||||
func NewWriter(options *writers.WriterOptions) *Writer {
|
func NewWriter(options *writers.WriterOptions) *Writer {
|
||||||
w := &Writer{
|
w := &Writer{
|
||||||
options: options,
|
options: options,
|
||||||
typeMapper: NewTypeMapper(options.NullableTypes),
|
typeMapper: NewTypeMapper(options.NullableTypes, options.NullableArrays),
|
||||||
config: LoadMethodConfigFromMetadata(options.Metadata),
|
config: LoadMethodConfigFromMetadata(options.Metadata),
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -51,6 +51,43 @@ func (w *Writer) WriteDatabase(db *models.Database) error {
|
|||||||
return w.writeSingleFile(db)
|
return w.writeSingleFile(db)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// addExtraFields appends user-defined extra fields to matching models.
|
||||||
|
// Fields may be scoped by target_table (table name, schema.table name, or model name).
|
||||||
|
// Unscoped fields apply to every generated model.
|
||||||
|
func (w *Writer) addExtraFields(templateData *TemplateData) {
|
||||||
|
extraFields := LoadExtraFieldsFromMetadata(w.options.Metadata)
|
||||||
|
if len(extraFields) == 0 {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, model := range templateData.Models {
|
||||||
|
for _, ef := range extraFields {
|
||||||
|
if !extraFieldMatchesModel(ef, model) {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
fieldData := &FieldData{
|
||||||
|
Name: resolveFieldNameCollision(ef.Name),
|
||||||
|
Type: ef.Type,
|
||||||
|
BunTag: ef.BunTag,
|
||||||
|
JSONTag: ef.JSONTag,
|
||||||
|
Comment: ef.Comment,
|
||||||
|
}
|
||||||
|
model.ExtraFields = append(model.ExtraFields, fieldData)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func extraFieldMatchesModel(field ExtraFieldConfig, model *ModelData) bool {
|
||||||
|
if field.TargetTable == "" {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
return field.TargetTable == model.TableNameOnly ||
|
||||||
|
field.TargetTable == model.TableName ||
|
||||||
|
field.TargetTable == model.Name
|
||||||
|
}
|
||||||
|
|
||||||
// WriteSchema writes a schema as Bun models
|
// WriteSchema writes a schema as Bun models
|
||||||
func (w *Writer) WriteSchema(schema *models.Schema) error {
|
func (w *Writer) WriteSchema(schema *models.Schema) error {
|
||||||
// Create a temporary database with just this schema
|
// Create a temporary database with just this schema
|
||||||
@@ -80,7 +117,7 @@ func (w *Writer) writeSingleFile(db *models.Database) error {
|
|||||||
// Add bun import (always needed)
|
// Add bun import (always needed)
|
||||||
templateData.AddImport(fmt.Sprintf("\"%s\"", w.typeMapper.GetBunImport()))
|
templateData.AddImport(fmt.Sprintf("\"%s\"", w.typeMapper.GetBunImport()))
|
||||||
|
|
||||||
// Add nullable types import (resolvespec or stdlib depending on options)
|
// Add nullable types import (sqltypes or stdlib depending on options)
|
||||||
templateData.AddImport(w.typeMapper.GetNullableTypeImportLine())
|
templateData.AddImport(w.typeMapper.GetNullableTypeImportLine())
|
||||||
|
|
||||||
// Collect all models
|
// Collect all models
|
||||||
@@ -107,6 +144,9 @@ func (w *Writer) writeSingleFile(db *models.Database) error {
|
|||||||
templateData.AddImport("\"fmt\"")
|
templateData.AddImport("\"fmt\"")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Apply extra fields from generator config (issue #4)
|
||||||
|
w.addExtraFields(templateData)
|
||||||
|
|
||||||
// Finalize imports
|
// Finalize imports
|
||||||
templateData.FinalizeImports()
|
templateData.FinalizeImports()
|
||||||
|
|
||||||
@@ -177,7 +217,7 @@ func (w *Writer) writeMultiFile(db *models.Database) error {
|
|||||||
// Add bun import
|
// Add bun import
|
||||||
templateData.AddImport(fmt.Sprintf("\"%s\"", w.typeMapper.GetBunImport()))
|
templateData.AddImport(fmt.Sprintf("\"%s\"", w.typeMapper.GetBunImport()))
|
||||||
|
|
||||||
// Add nullable types import (resolvespec or stdlib depending on options)
|
// Add nullable types import (sqltypes or stdlib depending on options)
|
||||||
templateData.AddImport(w.typeMapper.GetNullableTypeImportLine())
|
templateData.AddImport(w.typeMapper.GetNullableTypeImportLine())
|
||||||
|
|
||||||
// Create model data
|
// Create model data
|
||||||
@@ -200,6 +240,9 @@ func (w *Writer) writeMultiFile(db *models.Database) error {
|
|||||||
templateData.AddImport("\"fmt\"")
|
templateData.AddImport("\"fmt\"")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Apply extra fields from generator config (issue #4)
|
||||||
|
w.addExtraFields(templateData)
|
||||||
|
|
||||||
// Finalize imports
|
// Finalize imports
|
||||||
templateData.FinalizeImports()
|
templateData.FinalizeImports()
|
||||||
|
|
||||||
|
|||||||
+286
-40
@@ -73,9 +73,9 @@ func TestWriter_WriteTable(t *testing.T) {
|
|||||||
"ID",
|
"ID",
|
||||||
"int64",
|
"int64",
|
||||||
"Email",
|
"Email",
|
||||||
"resolvespec_common.SqlString",
|
"*string",
|
||||||
"CreatedAt",
|
"CreatedAt",
|
||||||
"resolvespec_common.SqlTime",
|
"time.Time",
|
||||||
"bun:\"id",
|
"bun:\"id",
|
||||||
"bun:\"email",
|
"bun:\"email",
|
||||||
"func (m ModelPublicUsers) TableName() string",
|
"func (m ModelPublicUsers) TableName() string",
|
||||||
@@ -550,13 +550,13 @@ func TestWriter_FieldNameCollision(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Verify NO field named just "TableName" (without underscore)
|
// Verify NO field named just "TableName" (without underscore)
|
||||||
if strings.Contains(generated, "TableName resolvespec_common") || strings.Contains(generated, "TableName string") {
|
if strings.Contains(generated, "TableName sql_types") || strings.Contains(generated, "TableName string") {
|
||||||
t.Errorf("Field 'TableName' without underscore should not exist (would conflict with method)\nGenerated:\n%s", generated)
|
t.Errorf("Field 'TableName' without underscore should not exist (would conflict with method)\nGenerated:\n%s", generated)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestTypeMapper_SQLTypeToGoType_Bun(t *testing.T) {
|
func TestTypeMapper_SQLTypeToGoType_Bun(t *testing.T) {
|
||||||
mapper := NewTypeMapper("")
|
mapper := NewTypeMapper("", "")
|
||||||
|
|
||||||
tests := []struct {
|
tests := []struct {
|
||||||
sqlType string
|
sqlType string
|
||||||
@@ -564,20 +564,20 @@ func TestTypeMapper_SQLTypeToGoType_Bun(t *testing.T) {
|
|||||||
want string
|
want string
|
||||||
}{
|
}{
|
||||||
{"bigint", true, "int64"},
|
{"bigint", true, "int64"},
|
||||||
{"bigint", false, "resolvespec_common.SqlInt64"},
|
{"bigint", false, "*int64"},
|
||||||
{"varchar", true, "resolvespec_common.SqlString"}, // Bun uses sql types even for NOT NULL strings
|
{"varchar", true, "string"},
|
||||||
{"varchar", false, "resolvespec_common.SqlString"},
|
{"varchar", false, "*string"},
|
||||||
{"timestamp", true, "resolvespec_common.SqlTimeStamp"},
|
{"timestamp", true, "time.Time"},
|
||||||
{"timestamp", false, "resolvespec_common.SqlTimeStamp"},
|
{"timestamp", false, "*time.Time"},
|
||||||
{"date", false, "resolvespec_common.SqlDate"},
|
{"date", false, "*time.Time"},
|
||||||
{"boolean", true, "bool"},
|
{"boolean", true, "bool"},
|
||||||
{"boolean", false, "resolvespec_common.SqlBool"},
|
{"boolean", false, "*bool"},
|
||||||
{"uuid", false, "resolvespec_common.SqlUUID"},
|
{"uuid", false, "*string"},
|
||||||
{"jsonb", false, "resolvespec_common.SqlJSONB"},
|
{"jsonb", false, "*string"},
|
||||||
{"text[]", true, "resolvespec_common.SqlStringArray"},
|
{"text[]", true, "[]string"},
|
||||||
{"text[]", false, "resolvespec_common.SqlStringArray"},
|
{"text[]", false, "[]string"},
|
||||||
{"integer[]", true, "resolvespec_common.SqlInt32Array"},
|
{"integer[]", true, "[]int32"},
|
||||||
{"bigint[]", false, "resolvespec_common.SqlInt64Array"},
|
{"bigint[]", false, "[]int64"},
|
||||||
}
|
}
|
||||||
|
|
||||||
for _, tt := range tests {
|
for _, tt := range tests {
|
||||||
@@ -599,7 +599,7 @@ func TestWriter_UpdateIDTypeSafety_Bun(t *testing.T) {
|
|||||||
forbidInt32 bool
|
forbidInt32 bool
|
||||||
}{
|
}{
|
||||||
{"int32_pk", "int", "int32", "m.ID = int32(newid)", false},
|
{"int32_pk", "int", "int32", "m.ID = int32(newid)", false},
|
||||||
{"sql_int16_pk", "smallint", "resolvespec_common.SqlInt16", "m.ID.FromString(fmt.Sprintf(\"%d\", newid))", true},
|
{"sql_int16_pk", "smallint", "int16", "m.ID = int16(newid)", true},
|
||||||
{"int64_pk", "bigint", "int64", "m.ID = int64(newid)", true},
|
{"int64_pk", "bigint", "int64", "m.ID = int64(newid)", true},
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -680,13 +680,13 @@ func TestWriter_StringPrimaryKeyHelpers_Bun(t *testing.T) {
|
|||||||
generated := string(content)
|
generated := string(content)
|
||||||
|
|
||||||
expectations := []string{
|
expectations := []string{
|
||||||
"resolvespec_common.SqlUUID",
|
"ID string",
|
||||||
"func (m ModelPublicAccounts) GetID() string",
|
"func (m ModelPublicAccounts) GetID() string",
|
||||||
"return m.ID.String()",
|
"return m.ID",
|
||||||
"func (m ModelPublicAccounts) GetIDStr() string",
|
"func (m ModelPublicAccounts) GetIDStr() string",
|
||||||
"func (m ModelPublicAccounts) SetID(newid string)",
|
"func (m ModelPublicAccounts) SetID(newid string)",
|
||||||
"func (m *ModelPublicAccounts) UpdateID(newid string)",
|
"func (m *ModelPublicAccounts) UpdateID(newid string)",
|
||||||
"m.ID.FromString(newid)",
|
"m.ID = newid",
|
||||||
}
|
}
|
||||||
|
|
||||||
for _, expected := range expectations {
|
for _, expected := range expectations {
|
||||||
@@ -701,7 +701,7 @@ func TestWriter_StringPrimaryKeyHelpers_Bun(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestTypeMapper_BuildBunTag(t *testing.T) {
|
func TestTypeMapper_BuildBunTag(t *testing.T) {
|
||||||
mapper := NewTypeMapper("")
|
mapper := NewTypeMapper("", "")
|
||||||
|
|
||||||
tests := []struct {
|
tests := []struct {
|
||||||
name string
|
name string
|
||||||
@@ -827,36 +827,282 @@ func TestTypeMapper_BuildBunTag(t *testing.T) {
|
|||||||
t.Errorf("BuildBunTag() = %q, missing %q", result, part)
|
t.Errorf("BuildBunTag() = %q, missing %q", result, part)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
// resolvespec mode must NOT add "array" — SqlXxxArray uses sql.Scanner
|
// Array columns always carry the "array" tag, telling bun's
|
||||||
if strings.Contains(result, ",array,") || strings.HasSuffix(result, ",array,") {
|
// pgdialect to scan/append the native Go slice as a PostgreSQL array.
|
||||||
t.Errorf("BuildBunTag() = %q, must not contain 'array' in resolvespec mode", result)
|
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) {
|
// TestTypeMapper_BuildBunTag_MultipleUniqueIndexesDeterministic verifies that
|
||||||
mapper := NewTypeMapper(writers.NullableTypeStdlib)
|
// when a column belongs to more than one unique index, the "unique:" tag
|
||||||
|
// fragments always appear in the same order across repeated calls, instead
|
||||||
cases := []struct {
|
// of following Go's randomized map iteration order over Table.Indexes.
|
||||||
name string
|
func TestTypeMapper_BuildBunTag_MultipleUniqueIndexesDeterministic(t *testing.T) {
|
||||||
column *models.Column
|
mapper := NewTypeMapper("", "")
|
||||||
}{
|
table := &models.Table{
|
||||||
{name: "text array", column: &models.Column{Name: "tags", Type: "text[]"}},
|
Name: "accounts",
|
||||||
{name: "integer array", column: &models.Column{Name: "scores", Type: "integer[]", NotNull: true}},
|
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,
|
||||||
|
},
|
||||||
|
},
|
||||||
}
|
}
|
||||||
for _, tt := range cases {
|
column := &models.Column{Name: "email", Type: "varchar", Length: 255, NotNull: true}
|
||||||
t.Run(tt.name, func(t *testing.T) {
|
|
||||||
result := mapper.BuildBunTag(tt.column, nil)
|
first := mapper.BuildBunTag(column, table)
|
||||||
if !strings.Contains(result, "array") {
|
for i := 0; i < 50; i++ {
|
||||||
t.Errorf("BuildBunTag() = %q, expected 'array' in stdlib mode", result)
|
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) {
|
||||||
|
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)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestExtraFields_LoadFromMetadata(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
metadata map[string]interface{}
|
||||||
|
expected int
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "nil metadata",
|
||||||
|
metadata: nil,
|
||||||
|
expected: 0,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "no extra_fields key",
|
||||||
|
metadata: map[string]interface{}{
|
||||||
|
"generate_table_name": true,
|
||||||
|
},
|
||||||
|
expected: 0,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "empty array",
|
||||||
|
metadata: map[string]interface{}{
|
||||||
|
"extra_fields": []interface{}{},
|
||||||
|
},
|
||||||
|
expected: 0,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "JSON string format",
|
||||||
|
metadata: map[string]interface{}{
|
||||||
|
"extra_fields": `[{"name":"ThoughtCount","type":"int64","bun_tag":"thought_count,scanonly","json_tag":"thought_count"}]`,
|
||||||
|
},
|
||||||
|
expected: 1,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "structured format",
|
||||||
|
metadata: map[string]interface{}{
|
||||||
|
"extra_fields": []interface{}{
|
||||||
|
map[string]interface{}{
|
||||||
|
"name": "ThoughtCount",
|
||||||
|
"type": "int64",
|
||||||
|
"bun_tag": "thought_count,scanonly",
|
||||||
|
"json_tag": "thought_count",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
expected: 1,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
fields := LoadExtraFieldsFromMetadata(tt.metadata)
|
||||||
|
if len(fields) != tt.expected {
|
||||||
|
t.Errorf("expected %d fields, got %d", tt.expected, len(fields))
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestExtraFields_InSingleFile(t *testing.T) {
|
||||||
|
table := models.InitTable("projects", "public")
|
||||||
|
table.Columns["id"] = &models.Column{
|
||||||
|
Name: "id",
|
||||||
|
Type: "bigint",
|
||||||
|
NotNull: true,
|
||||||
|
IsPrimaryKey: true,
|
||||||
|
}
|
||||||
|
table.Columns["name"] = &models.Column{
|
||||||
|
Name: "name",
|
||||||
|
Type: "varchar",
|
||||||
|
Length: 255,
|
||||||
|
Sequence: 1,
|
||||||
|
}
|
||||||
|
|
||||||
|
opts := &writers.WriterOptions{
|
||||||
|
PackageName: "models",
|
||||||
|
OutputPath: t.TempDir() + "/test.go",
|
||||||
|
Metadata: map[string]interface{}{
|
||||||
|
"extra_fields": `[{"name":"ThoughtCount","type":"int64","bun_tag":"thought_count,scanonly","json_tag":"thought_count"}]`,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
writer := NewWriter(opts)
|
||||||
|
err := writer.WriteTable(table)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("WriteTable failed: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
content, _ := os.ReadFile(opts.OutputPath)
|
||||||
|
generated := string(content)
|
||||||
|
|
||||||
|
// Verify extra field is present with separator comment
|
||||||
|
expectations := []string{
|
||||||
|
"// --- Custom fields (managed by relspecgo generator config, issue #4) ---",
|
||||||
|
"ThoughtCount int64 `bun:\"thought_count,scanonly\" json:\"thought_count\"`",
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, expected := range expectations {
|
||||||
|
if !strings.Contains(generated, expected) {
|
||||||
|
t.Errorf("Generated code missing: %q\nGenerated:\n%s", expected, generated)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestExtraFields_InMultiFile(t *testing.T) {
|
||||||
|
db := models.InitDatabase("testdb")
|
||||||
|
schema := models.InitSchema("public")
|
||||||
|
|
||||||
|
users := models.InitTable("users", "public")
|
||||||
|
users.Columns["id"] = &models.Column{
|
||||||
|
Name: "id",
|
||||||
|
Type: "bigint",
|
||||||
|
NotNull: true,
|
||||||
|
IsPrimaryKey: true,
|
||||||
|
}
|
||||||
|
schema.Tables = append(schema.Tables, users)
|
||||||
|
|
||||||
|
posts := models.InitTable("posts", "public")
|
||||||
|
posts.Columns["id"] = &models.Column{
|
||||||
|
Name: "id",
|
||||||
|
Type: "bigint",
|
||||||
|
NotNull: true,
|
||||||
|
IsPrimaryKey: true,
|
||||||
|
}
|
||||||
|
schema.Tables = append(schema.Tables, posts)
|
||||||
|
|
||||||
|
db.Schemas = append(db.Schemas, schema)
|
||||||
|
|
||||||
|
tmpDir := t.TempDir()
|
||||||
|
opts := &writers.WriterOptions{
|
||||||
|
PackageName: "models",
|
||||||
|
OutputPath: tmpDir,
|
||||||
|
Metadata: map[string]interface{}{
|
||||||
|
"multi_file": true,
|
||||||
|
"extra_fields": `[{"target_table":"users","name":"PostCount","type":"int64","bun_tag":"post_count,scanonly","json_tag":"post_count"}]`,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
writer := NewWriter(opts)
|
||||||
|
err := writer.WriteDatabase(db)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("WriteDatabase failed: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check users file has the extra field (matched by table name "users")
|
||||||
|
usersContent, _ := os.ReadFile(tmpDir + "/sql_public_users.go")
|
||||||
|
usersStr := string(usersContent)
|
||||||
|
|
||||||
|
if !strings.Contains(usersStr, "PostCount int64 `bun:\"post_count,scanonly\"") {
|
||||||
|
t.Errorf("Extra field not found in users file:\n%s", usersStr)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Posts file should NOT have the extra field (name mismatch)
|
||||||
|
postsContent, _ := os.ReadFile(tmpDir + "/sql_public_posts.go")
|
||||||
|
postsStr := string(postsContent)
|
||||||
|
|
||||||
|
if strings.Contains(postsStr, "PostCount") {
|
||||||
|
t.Errorf("Extra field should not be in posts file:\n%s", postsStr)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestTypeMapper_BuildBunTag_PreservesExplicitTypeModifiers(t *testing.T) {
|
func TestTypeMapper_BuildBunTag_PreservesExplicitTypeModifiers(t *testing.T) {
|
||||||
mapper := NewTypeMapper("")
|
mapper := NewTypeMapper("", "")
|
||||||
|
|
||||||
col := &models.Column{
|
col := &models.Column{
|
||||||
Name: "embedding",
|
Name: "embedding",
|
||||||
|
|||||||
@@ -3,6 +3,7 @@ package dbml
|
|||||||
import (
|
import (
|
||||||
"fmt"
|
"fmt"
|
||||||
"os"
|
"os"
|
||||||
|
"sort"
|
||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||||
@@ -78,7 +79,7 @@ func (w *Writer) databaseToDBML(d *models.Database) string {
|
|||||||
sb.WriteString("\n// Relationships\n")
|
sb.WriteString("\n// Relationships\n")
|
||||||
for _, schema := range d.Schemas {
|
for _, schema := range d.Schemas {
|
||||||
for _, table := range schema.Tables {
|
for _, table := range schema.Tables {
|
||||||
for _, constraint := range table.Constraints {
|
for _, constraint := range sortConstraints(table.Constraints) {
|
||||||
if constraint.Type == models.ForeignKeyConstraint {
|
if constraint.Type == models.ForeignKeyConstraint {
|
||||||
sb.WriteString(w.constraintToDBML(constraint, table))
|
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)
|
tableName := fmt.Sprintf("%s.%s", t.Schema, t.Name)
|
||||||
fmt.Fprintf(&sb, "Table %s {\n", tableName)
|
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)
|
fmt.Fprintf(&sb, " %s %s", column.Name, column.Type)
|
||||||
|
|
||||||
var attrs []string
|
var attrs []string
|
||||||
@@ -149,7 +150,7 @@ func (w *Writer) tableToDBML(t *models.Table) string {
|
|||||||
|
|
||||||
if len(t.Indexes) > 0 {
|
if len(t.Indexes) > 0 {
|
||||||
sb.WriteString("\n indexes {\n")
|
sb.WriteString("\n indexes {\n")
|
||||||
for _, index := range t.Indexes {
|
for _, index := range sortIndexes(t.Indexes) {
|
||||||
var indexAttrs []string
|
var indexAttrs []string
|
||||||
if index.Unique {
|
if index.Unique {
|
||||||
indexAttrs = append(indexAttrs, "unique")
|
indexAttrs = append(indexAttrs, "unique")
|
||||||
@@ -230,3 +231,48 @@ func (w *Writer) constraintToDBML(c *models.Constraint, t *models.Table) string
|
|||||||
|
|
||||||
return refLine + "\n"
|
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
|
// Add table-level relationships
|
||||||
for _, table := range tableSlice {
|
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)
|
// Check if this relationship is already in the list (avoid duplicates)
|
||||||
isDuplicate := false
|
isDuplicate := false
|
||||||
for _, existing := range allRelations {
|
for _, existing := range allRelations {
|
||||||
|
|||||||
@@ -4,6 +4,7 @@ import (
|
|||||||
"encoding/json"
|
"encoding/json"
|
||||||
"fmt"
|
"fmt"
|
||||||
"os"
|
"os"
|
||||||
|
"sort"
|
||||||
|
|
||||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||||
"git.warky.dev/wdevs/relspecgo/pkg/writers"
|
"git.warky.dev/wdevs/relspecgo/pkg/writers"
|
||||||
@@ -175,7 +176,7 @@ func (w *Writer) databaseToDrawDB(d *models.Database) *DrawDBSchema {
|
|||||||
// Add relationships
|
// Add relationships
|
||||||
for _, schemaModel := range d.Schemas {
|
for _, schemaModel := range d.Schemas {
|
||||||
for _, table := range schemaModel.Tables {
|
for _, table := range schemaModel.Tables {
|
||||||
for _, constraint := range table.Constraints {
|
for _, constraint := range sortConstraints(table.Constraints) {
|
||||||
if constraint.Type == models.ForeignKeyConstraint && constraint.ReferencedTable != "" {
|
if constraint.Type == models.ForeignKeyConstraint && constraint.ReferencedTable != "" {
|
||||||
startTableKey := fmt.Sprintf("%s.%s", schemaModel.Name, table.Name)
|
startTableKey := fmt.Sprintf("%s.%s", schemaModel.Name, table.Name)
|
||||||
endTableKey := fmt.Sprintf("%s.%s", constraint.ReferencedSchema, constraint.ReferencedTable)
|
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
|
// Add fields
|
||||||
for _, column := range table.Columns {
|
for _, column := range sortColumns(table.Columns) {
|
||||||
field := &DrawDBField{
|
field := &DrawDBField{
|
||||||
ID: fieldID,
|
ID: fieldID,
|
||||||
Name: column.Name,
|
Name: column.Name,
|
||||||
@@ -339,7 +340,7 @@ func (w *Writer) convertTableToDrawDB(table *models.Table, schemaName string, ta
|
|||||||
|
|
||||||
// Add indexes
|
// Add indexes
|
||||||
indexID := 0
|
indexID := 0
|
||||||
for _, index := range table.Indexes {
|
for _, index := range sortIndexes(table.Indexes) {
|
||||||
drawIndex := &DrawDBIndex{
|
drawIndex := &DrawDBIndex{
|
||||||
ID: indexID,
|
ID: indexID,
|
||||||
Name: index.Name,
|
Name: index.Name,
|
||||||
@@ -393,3 +394,48 @@ func getColorForIndex(index int) string {
|
|||||||
}
|
}
|
||||||
return colors[index%len(colors)]
|
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"
|
"fmt"
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
|
"sort"
|
||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
"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)
|
indexColumnFields := make(map[string]bool)
|
||||||
|
|
||||||
// Add indexes (excluding single-column unique indexes, which are handled inline)
|
// 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)
|
// Skip single-column unique indexes (handled by .unique() modifier)
|
||||||
if index.Unique && len(index.Columns) == 1 {
|
if index.Unique && len(index.Columns) == 1 {
|
||||||
continue
|
continue
|
||||||
@@ -270,7 +271,7 @@ func (w *Writer) buildTableData(table *models.Table, schema *models.Schema, db *
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Add multi-column unique constraints as unique indexes
|
// 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 {
|
if constraint.Type == models.UniqueConstraint && len(constraint.Columns) > 1 {
|
||||||
// Create a unique index for this constraint
|
// Create a unique index for this constraint
|
||||||
indexData := &IndexData{
|
indexData := &IndexData{
|
||||||
@@ -316,6 +317,36 @@ func (w *Writer) buildTableData(table *models.Table, schema *models.Schema, db *
|
|||||||
return tableData
|
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
|
// sortStrings sorts a slice of strings in place
|
||||||
func sortStrings(strs []string) {
|
func sortStrings(strs []string) {
|
||||||
for i := 0; i < len(strs); i++ {
|
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)
|
enumNames := make([]string, 0)
|
||||||
seen := make(map[string]bool)
|
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)] {
|
if enumMap[col.Type] || enumMap[strings.ToLower(col.Type)] {
|
||||||
// Find the enum in schema
|
// Find the enum in schema
|
||||||
for _, enum := range schema.Enums {
|
for _, enum := range schema.Enums {
|
||||||
|
|||||||
@@ -6,6 +6,8 @@ Generates Go source files with GORM model definitions from database schema infor
|
|||||||
|
|
||||||
The GORM Writer converts RelSpec's internal database model representation into Go source code with GORM struct definitions, complete with proper tags, relationships, and methods.
|
The GORM Writer converts RelSpec's internal database model representation into Go source code with GORM struct definitions, complete with proper tags, relationships, and methods.
|
||||||
|
|
||||||
|
With `--types sqltypes`, nullable fields use the [`pkg/sqltypes`](../../sqltypes/README.md) package.
|
||||||
|
|
||||||
## Features
|
## Features
|
||||||
|
|
||||||
- Generates GORM-compatible Go structs
|
- Generates GORM-compatible Go structs
|
||||||
@@ -48,19 +50,19 @@ func main() {
|
|||||||
### CLI Examples
|
### CLI Examples
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
# Generate GORM models from a DBML schema (default: resolvespec types)
|
# Generate GORM models from a DBML schema (default: baselib pointer types)
|
||||||
relspec convert --from dbml --from-path schema.dbml \
|
relspec convert --from dbml --from-path schema.dbml \
|
||||||
--to gorm --to-path models.go --package models
|
--to gorm --to-path models.go --package models
|
||||||
|
|
||||||
# Use standard library database/sql nullable types instead of resolvespec
|
# Use standard library database/sql nullable types instead
|
||||||
relspec convert --from dbml --from-path schema.dbml \
|
relspec convert --from dbml --from-path schema.dbml \
|
||||||
--to gorm --to-path models.go --package models \
|
--to gorm --to-path models.go --package models \
|
||||||
--types stdlib
|
--types stdlib
|
||||||
|
|
||||||
# Explicitly select resolvespec types (same as omitting --types)
|
# Select sqltypes package types (git.warky.dev/wdevs/relspecgo/pkg/sqltypes)
|
||||||
relspec convert --from pgsql --from-conn "postgres://localhost/mydb" \
|
relspec convert --from pgsql --from-conn "postgres://localhost/mydb" \
|
||||||
--to gorm --to-path models.go --package models \
|
--to gorm --to-path models.go --package models \
|
||||||
--types resolvespec
|
--types sqltypes
|
||||||
|
|
||||||
# Multi-file output (one file per table)
|
# Multi-file output (one file per table)
|
||||||
relspec convert --from json --from-path schema.json \
|
relspec convert --from json --from-path schema.json \
|
||||||
@@ -89,13 +91,13 @@ Files are named: `sql_{schema}_{table}.go`
|
|||||||
|
|
||||||
## Generated Code Examples
|
## Generated Code Examples
|
||||||
|
|
||||||
### Default — resolvespec types (`--types resolvespec`)
|
### sqltypes package types (`--types sqltypes`)
|
||||||
|
|
||||||
```go
|
```go
|
||||||
package models
|
package models
|
||||||
|
|
||||||
import (
|
import (
|
||||||
sql_types "github.com/bitechdev/ResolveSpec/pkg/spectypes"
|
sql_types "git.warky.dev/wdevs/relspecgo/pkg/sqltypes"
|
||||||
)
|
)
|
||||||
|
|
||||||
type ModelUser struct {
|
type ModelUser struct {
|
||||||
@@ -141,11 +143,11 @@ func (ModelUser) TableName() string {
|
|||||||
Controls which Go package is used for nullable column types. Set via the `--types` CLI flag or `WriterOptions.NullableTypes`:
|
Controls which Go package is used for nullable column types. Set via the `--types` CLI flag or `WriterOptions.NullableTypes`:
|
||||||
|
|
||||||
```go
|
```go
|
||||||
// Use resolvespec types (default — omit NullableTypes or set to "resolvespec")
|
// Use sqltypes package types
|
||||||
options := &writers.WriterOptions{
|
options := &writers.WriterOptions{
|
||||||
OutputPath: "models.go",
|
OutputPath: "models.go",
|
||||||
PackageName: "models",
|
PackageName: "models",
|
||||||
NullableTypes: writers.NullableTypeResolveSpec,
|
NullableTypes: writers.NullableTypeSqlTypes,
|
||||||
}
|
}
|
||||||
|
|
||||||
// Use standard library database/sql types
|
// Use standard library database/sql types
|
||||||
@@ -176,7 +178,7 @@ options := &writers.WriterOptions{
|
|||||||
|
|
||||||
The nullable type package is selected with `--types` (or `WriterOptions.NullableTypes`).
|
The nullable type package is selected with `--types` (or `WriterOptions.NullableTypes`).
|
||||||
|
|
||||||
| SQL Type | NOT NULL — both | Nullable — resolvespec | Nullable — stdlib |
|
| SQL Type | NOT NULL — both | Nullable — sqltypes | Nullable — stdlib |
|
||||||
|---|---|---|---|
|
|---|---|---|---|
|
||||||
| `bigint` | `int64` | `SqlInt64` | `sql.NullInt64` |
|
| `bigint` | `int64` | `SqlInt64` | `sql.NullInt64` |
|
||||||
| `integer` | `int32` | `SqlInt32` | `sql.NullInt32` |
|
| `integer` | `int32` | `SqlInt32` | `sql.NullInt32` |
|
||||||
|
|||||||
@@ -134,8 +134,11 @@ func NewModelData(table *models.Table, schema string, typeMapper *TypeMapper, fl
|
|||||||
Prefix: GeneratePrefix(table.Name),
|
Prefix: GeneratePrefix(table.Name),
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Convert columns to fields (sorted by sequence or name)
|
||||||
|
columns := sortColumns(table.Columns)
|
||||||
|
|
||||||
// Find primary key
|
// Find primary key
|
||||||
for _, col := range table.Columns {
|
for _, col := range columns {
|
||||||
if col.IsPrimaryKey {
|
if col.IsPrimaryKey {
|
||||||
// Sanitize column name to remove backticks
|
// Sanitize column name to remove backticks
|
||||||
safeName := writers.SanitizeStructTagValue(col.Name)
|
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 {
|
for _, col := range columns {
|
||||||
field := columnToField(col, table, typeMapper)
|
field := columnToField(col, table, typeMapper)
|
||||||
// Check for name collision with generated methods and rename if needed
|
// Check for name collision with generated methods and rename if needed
|
||||||
@@ -202,7 +203,7 @@ func formatComment(description, comment string) string {
|
|||||||
|
|
||||||
func isStringLikePrimaryKeyType(goType string) bool {
|
func isStringLikePrimaryKeyType(goType string) bool {
|
||||||
switch goType {
|
switch goType {
|
||||||
case "string", "sql.NullString", "sql_types.SqlString", "sql_types.SqlUUID":
|
case "string", "*string", "sql.NullString", "sql_types.SqlString", "sql_types.SqlUUID":
|
||||||
return true
|
return true
|
||||||
default:
|
default:
|
||||||
return false
|
return false
|
||||||
@@ -248,6 +249,21 @@ func sortConstraints(constraints map[string]*models.Constraint) []*models.Constr
|
|||||||
return result
|
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
|
// sortColumns sorts columns by sequence, then by name
|
||||||
func sortColumns(columns map[string]*models.Column) []*models.Column {
|
func sortColumns(columns map[string]*models.Column) []*models.Column {
|
||||||
result := make([]*models.Column, 0, len(columns))
|
result := make([]*models.Column, 0, len(columns))
|
||||||
|
|||||||
@@ -12,15 +12,15 @@ import (
|
|||||||
// TypeMapper handles type conversions between SQL and Go types
|
// TypeMapper handles type conversions between SQL and Go types
|
||||||
type TypeMapper struct {
|
type TypeMapper struct {
|
||||||
sqlTypesAlias string
|
sqlTypesAlias string
|
||||||
typeStyle string // writers.NullableTypeResolveSpec | writers.NullableTypeStdlib
|
typeStyle string // writers.NullableTypeSqlTypes | writers.NullableTypeStdlib | writers.NullableTypeBaselib
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewTypeMapper creates a new TypeMapper.
|
// NewTypeMapper creates a new TypeMapper.
|
||||||
// typeStyle should be writers.NullableTypeResolveSpec or writers.NullableTypeStdlib;
|
// typeStyle should be writers.NullableTypeSqlTypes, writers.NullableTypeStdlib, or
|
||||||
// an empty string defaults to resolvespec.
|
// writers.NullableTypeBaselib; an empty string defaults to baselib.
|
||||||
func NewTypeMapper(typeStyle string) *TypeMapper {
|
func NewTypeMapper(typeStyle string) *TypeMapper {
|
||||||
if typeStyle == "" {
|
if typeStyle == "" {
|
||||||
typeStyle = writers.NullableTypeResolveSpec
|
typeStyle = writers.NullableTypeBaselib
|
||||||
}
|
}
|
||||||
return &TypeMapper{
|
return &TypeMapper{
|
||||||
sqlTypesAlias: "sql_types",
|
sqlTypesAlias: "sql_types",
|
||||||
@@ -37,14 +37,20 @@ func (tm *TypeMapper) SQLTypeToGoType(sqlType string, notNull bool) string {
|
|||||||
|
|
||||||
baseType := tm.extractBaseType(sqlType)
|
baseType := tm.extractBaseType(sqlType)
|
||||||
|
|
||||||
if tm.typeStyle == writers.NullableTypeStdlib {
|
switch tm.typeStyle {
|
||||||
|
case writers.NullableTypeStdlib:
|
||||||
if notNull {
|
if notNull {
|
||||||
return tm.rawGoType(baseType)
|
return tm.rawGoType(baseType)
|
||||||
}
|
}
|
||||||
return tm.stdlibNullableGoType(baseType)
|
return tm.stdlibNullableGoType(baseType)
|
||||||
|
case writers.NullableTypeBaselib:
|
||||||
|
if notNull {
|
||||||
|
return tm.rawGoType(baseType)
|
||||||
|
}
|
||||||
|
return tm.baselibNullableGoType(baseType)
|
||||||
}
|
}
|
||||||
|
|
||||||
// resolvespec (default)
|
// sqltypes
|
||||||
if notNull {
|
if notNull {
|
||||||
return tm.baseGoType(baseType)
|
return tm.baseGoType(baseType)
|
||||||
}
|
}
|
||||||
@@ -212,7 +218,7 @@ func (tm *TypeMapper) nullableGoType(sqlType string) string {
|
|||||||
// arrayGoType returns the Go type for a PostgreSQL array column.
|
// arrayGoType returns the Go type for a PostgreSQL array column.
|
||||||
// The baseElemType is the canonical base type (e.g. "text", "integer").
|
// The baseElemType is the canonical base type (e.g. "text", "integer").
|
||||||
func (tm *TypeMapper) arrayGoType(baseElemType string) string {
|
func (tm *TypeMapper) arrayGoType(baseElemType string) string {
|
||||||
if tm.typeStyle == writers.NullableTypeStdlib {
|
if tm.typeStyle == writers.NullableTypeStdlib || tm.typeStyle == writers.NullableTypeBaselib {
|
||||||
return tm.stdlibArrayGoType(baseElemType)
|
return tm.stdlibArrayGoType(baseElemType)
|
||||||
}
|
}
|
||||||
typeMap := map[string]string{
|
typeMap := map[string]string{
|
||||||
@@ -305,6 +311,38 @@ func (tm *TypeMapper) stdlibNullableGoType(sqlType string) string {
|
|||||||
return "sql.NullString"
|
return "sql.NullString"
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// baselibNullableGoType returns plain Go pointer types for nullable columns.
|
||||||
|
func (tm *TypeMapper) baselibNullableGoType(sqlType string) string {
|
||||||
|
typeMap := map[string]string{
|
||||||
|
"integer": "*int32", "int": "*int32", "int4": "*int32", "serial": "*int32",
|
||||||
|
"smallint": "*int16", "int2": "*int16", "smallserial": "*int16",
|
||||||
|
"bigint": "*int64", "int8": "*int64", "bigserial": "*int64",
|
||||||
|
"boolean": "*bool", "bool": "*bool",
|
||||||
|
"real": "*float32", "float4": "*float32",
|
||||||
|
"double precision": "*float64", "float8": "*float64",
|
||||||
|
"numeric": "*float64", "decimal": "*float64", "money": "*float64",
|
||||||
|
"text": "*string", "varchar": "*string", "char": "*string",
|
||||||
|
"character": "*string", "citext": "*string", "bpchar": "*string",
|
||||||
|
"inet": "*string", "cidr": "*string", "macaddr": "*string",
|
||||||
|
"uuid": "*string", "json": "*string", "jsonb": "*string",
|
||||||
|
"timestamp": "*time.Time",
|
||||||
|
"timestamp without time zone": "*time.Time",
|
||||||
|
"timestamp with time zone": "*time.Time",
|
||||||
|
"timestamptz": "*time.Time",
|
||||||
|
"date": "*time.Time",
|
||||||
|
"time": "*time.Time",
|
||||||
|
"time without time zone": "*time.Time",
|
||||||
|
"time with time zone": "*time.Time",
|
||||||
|
"timetz": "*time.Time",
|
||||||
|
"bytea": "[]byte",
|
||||||
|
"vector": "[]float32",
|
||||||
|
}
|
||||||
|
if goType, ok := typeMap[sqlType]; ok {
|
||||||
|
return goType
|
||||||
|
}
|
||||||
|
return "*string"
|
||||||
|
}
|
||||||
|
|
||||||
// stdlibArrayGoType returns a plain Go slice type for array columns in stdlib mode.
|
// stdlibArrayGoType returns a plain Go slice type for array columns in stdlib mode.
|
||||||
func (tm *TypeMapper) stdlibArrayGoType(baseElemType string) string {
|
func (tm *TypeMapper) stdlibArrayGoType(baseElemType string) string {
|
||||||
typeMap := map[string]string{
|
typeMap := map[string]string{
|
||||||
@@ -377,7 +415,7 @@ func (tm *TypeMapper) BuildGormTag(column *models.Column, table *models.Table) s
|
|||||||
|
|
||||||
// Check for unique constraint
|
// Check for unique constraint
|
||||||
if table != nil {
|
if table != nil {
|
||||||
for _, constraint := range table.Constraints {
|
for _, constraint := range sortConstraints(table.Constraints) {
|
||||||
if constraint.Type == models.UniqueConstraint {
|
if constraint.Type == models.UniqueConstraint {
|
||||||
for _, col := range constraint.Columns {
|
for _, col := range constraint.Columns {
|
||||||
if col == column.Name {
|
if col == column.Name {
|
||||||
@@ -393,7 +431,7 @@ func (tm *TypeMapper) BuildGormTag(column *models.Column, table *models.Table) s
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Check for index
|
// Check for index
|
||||||
for _, index := range table.Indexes {
|
for _, index := range sortIndexes(table.Indexes) {
|
||||||
for _, col := range index.Columns {
|
for _, col := range index.Columns {
|
||||||
if col == column.Name {
|
if col == column.Name {
|
||||||
if index.Unique {
|
if index.Unique {
|
||||||
@@ -467,16 +505,19 @@ func (tm *TypeMapper) NeedsFmtImport(generateGetIDStr bool) bool {
|
|||||||
return generateGetIDStr
|
return generateGetIDStr
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetSQLTypesImport returns the import path for the ResolveSpec spectypes package.
|
// GetSQLTypesImport returns the import path for the sqltypes package.
|
||||||
func (tm *TypeMapper) GetSQLTypesImport() string {
|
func (tm *TypeMapper) GetSQLTypesImport() string {
|
||||||
return "github.com/bitechdev/ResolveSpec/pkg/spectypes"
|
return "git.warky.dev/wdevs/relspecgo/pkg/sqltypes"
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetNullableTypeImportLine returns the full Go import line for the nullable type
|
// GetNullableTypeImportLine returns the full Go import line for the nullable type
|
||||||
// package (ready to pass to AddImport). Returns empty string when no import is needed.
|
// package (ready to pass to AddImport). Returns empty string when no import is needed.
|
||||||
func (tm *TypeMapper) GetNullableTypeImportLine() string {
|
func (tm *TypeMapper) GetNullableTypeImportLine() string {
|
||||||
if tm.typeStyle == writers.NullableTypeStdlib {
|
switch tm.typeStyle {
|
||||||
|
case writers.NullableTypeStdlib:
|
||||||
return "\"database/sql\""
|
return "\"database/sql\""
|
||||||
|
case writers.NullableTypeBaselib:
|
||||||
|
return ""
|
||||||
}
|
}
|
||||||
return fmt.Sprintf("%s \"%s\"", tm.sqlTypesAlias, tm.GetSQLTypesImport())
|
return fmt.Sprintf("%s \"%s\"", tm.sqlTypesAlias, tm.GetSQLTypesImport())
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -77,7 +77,7 @@ func (w *Writer) writeSingleFile(db *models.Database) error {
|
|||||||
packageName := w.getPackageName()
|
packageName := w.getPackageName()
|
||||||
templateData := NewTemplateData(packageName, w.config)
|
templateData := NewTemplateData(packageName, w.config)
|
||||||
|
|
||||||
// Add nullable types import (resolvespec or stdlib depending on options)
|
// Add nullable types import (sqltypes or stdlib depending on options)
|
||||||
templateData.AddImport(w.typeMapper.GetNullableTypeImportLine())
|
templateData.AddImport(w.typeMapper.GetNullableTypeImportLine())
|
||||||
|
|
||||||
// Collect all models
|
// Collect all models
|
||||||
@@ -171,7 +171,7 @@ func (w *Writer) writeMultiFile(db *models.Database) error {
|
|||||||
// Create template data for this single table
|
// Create template data for this single table
|
||||||
templateData := NewTemplateData(packageName, w.config)
|
templateData := NewTemplateData(packageName, w.config)
|
||||||
|
|
||||||
// Add nullable types import (resolvespec or stdlib depending on options)
|
// Add nullable types import (sqltypes or stdlib depending on options)
|
||||||
templateData.AddImport(w.typeMapper.GetNullableTypeImportLine())
|
templateData.AddImport(w.typeMapper.GetNullableTypeImportLine())
|
||||||
|
|
||||||
// Create model data
|
// Create model data
|
||||||
|
|||||||
@@ -70,7 +70,7 @@ func TestWriter_WriteTable(t *testing.T) {
|
|||||||
"ID",
|
"ID",
|
||||||
"int64",
|
"int64",
|
||||||
"Email",
|
"Email",
|
||||||
"sql_types.SqlString",
|
"*string",
|
||||||
"CreatedAt",
|
"CreatedAt",
|
||||||
"time.Time",
|
"time.Time",
|
||||||
"gorm:\"column:id",
|
"gorm:\"column:id",
|
||||||
@@ -700,17 +700,17 @@ func TestTypeMapper_SQLTypeToGoType(t *testing.T) {
|
|||||||
want string
|
want string
|
||||||
}{
|
}{
|
||||||
{"bigint", true, "int64"},
|
{"bigint", true, "int64"},
|
||||||
{"bigint", false, "sql_types.SqlInt64"},
|
{"bigint", false, "*int64"},
|
||||||
{"varchar", true, "string"},
|
{"varchar", true, "string"},
|
||||||
{"varchar", false, "sql_types.SqlString"},
|
{"varchar", false, "*string"},
|
||||||
{"timestamp", true, "time.Time"},
|
{"timestamp", true, "time.Time"},
|
||||||
{"timestamp", false, "sql_types.SqlTimeStamp"},
|
{"timestamp", false, "*time.Time"},
|
||||||
{"boolean", true, "bool"},
|
{"boolean", true, "bool"},
|
||||||
{"boolean", false, "sql_types.SqlBool"},
|
{"boolean", false, "*bool"},
|
||||||
{"text[]", true, "sql_types.SqlStringArray"},
|
{"text[]", true, "[]string"},
|
||||||
{"text[]", false, "sql_types.SqlStringArray"},
|
{"text[]", false, "[]string"},
|
||||||
{"integer[]", true, "sql_types.SqlInt32Array"},
|
{"integer[]", true, "[]int32"},
|
||||||
{"bigint[]", false, "sql_types.SqlInt64Array"},
|
{"bigint[]", false, "[]int64"},
|
||||||
}
|
}
|
||||||
|
|
||||||
for _, tt := range tests {
|
for _, tt := range tests {
|
||||||
@@ -757,3 +757,48 @@ func TestTypeMapper_BuildGormTag_PreservesExplicitTypeModifiers(t *testing.T) {
|
|||||||
t.Fatalf("type modifier appears duplicated in %q", tag)
|
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
|
// Add relation fields
|
||||||
relationFields = w.generateRelationFields(table, db, schema)
|
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 {
|
for _, field := range idFields {
|
||||||
sb.WriteString(field + "\n")
|
sb.WriteString(field + "\n")
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -239,7 +239,8 @@ func (w *MigrationWriter) generateDropScripts(model *models.Schema, current *mod
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Check each constraint in current database
|
// 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]
|
modelConstraint, existsInModel := modelTable.Constraints[constraintName]
|
||||||
|
|
||||||
shouldDrop := false
|
shouldDrop := false
|
||||||
@@ -252,7 +253,8 @@ func (w *MigrationWriter) generateDropScripts(model *models.Schema, current *mod
|
|||||||
if shouldDrop && currentConstraint.Type == models.PrimaryKeyConstraint {
|
if shouldDrop && currentConstraint.Type == models.PrimaryKeyConstraint {
|
||||||
// Drop FK constraints that depend on this PK before dropping the PK itself.
|
// Drop FK constraints that depend on this PK before dropping the PK itself.
|
||||||
for _, otherTable := range current.Tables {
|
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 {
|
if fkConstraint.Type != models.ForeignKeyConstraint {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
@@ -310,7 +312,8 @@ func (w *MigrationWriter) generateDropScripts(model *models.Schema, current *mod
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Check indexes
|
// Check indexes
|
||||||
for indexName, currentIndex := range currentTable.Indexes {
|
for _, currentIndex := range sortIndexes(currentTable.Indexes) {
|
||||||
|
indexName := currentIndex.Name
|
||||||
modelIndex, existsInModel := modelTable.Indexes[indexName]
|
modelIndex, existsInModel := modelTable.Indexes[indexName]
|
||||||
|
|
||||||
shouldDrop := false
|
shouldDrop := false
|
||||||
@@ -401,7 +404,7 @@ func (w *MigrationWriter) generateAlterTableScripts(schema *models.Schema, model
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Check each model column
|
// Check each model column
|
||||||
for _, modelCol := range modelTable.Columns {
|
for _, modelCol := range sortColumns(modelTable.Columns) {
|
||||||
currentCol, exists := currentColumns[strings.ToLower(modelCol.Name)]
|
currentCol, exists := currentColumns[strings.ToLower(modelCol.Name)]
|
||||||
|
|
||||||
if !exists {
|
if !exists {
|
||||||
@@ -518,7 +521,8 @@ func (w *MigrationWriter) generateIndexScripts(model *models.Schema, current *mo
|
|||||||
|
|
||||||
// Process primary keys first - check explicit constraints
|
// Process primary keys first - check explicit constraints
|
||||||
foundExplicitPK := false
|
foundExplicitPK := false
|
||||||
for constraintName, constraint := range modelTable.Constraints {
|
for _, constraint := range sortConstraints(modelTable.Constraints) {
|
||||||
|
constraintName := constraint.Name
|
||||||
if constraint.Type == models.PrimaryKeyConstraint {
|
if constraint.Type == models.PrimaryKeyConstraint {
|
||||||
foundExplicitPK = true
|
foundExplicitPK = true
|
||||||
shouldCreate := true
|
shouldCreate := true
|
||||||
@@ -603,7 +607,8 @@ func (w *MigrationWriter) generateIndexScripts(model *models.Schema, current *mo
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Process indexes
|
// Process indexes
|
||||||
for indexName, modelIndex := range modelTable.Indexes {
|
for _, modelIndex := range sortIndexes(modelTable.Indexes) {
|
||||||
|
indexName := modelIndex.Name
|
||||||
// Skip primary key indexes
|
// Skip primary key indexes
|
||||||
if strings.HasPrefix(strings.ToLower(indexName), "pk_") {
|
if strings.HasPrefix(strings.ToLower(indexName), "pk_") {
|
||||||
continue
|
continue
|
||||||
@@ -697,7 +702,8 @@ func (w *MigrationWriter) generateForeignKeyScripts(model *models.Schema, curren
|
|||||||
currentTable := currentTables[strings.ToLower(modelTable.Name)]
|
currentTable := currentTables[strings.ToLower(modelTable.Name)]
|
||||||
|
|
||||||
// Process each constraint
|
// Process each constraint
|
||||||
for constraintName, constraint := range modelTable.Constraints {
|
for _, constraint := range sortConstraints(modelTable.Constraints) {
|
||||||
|
constraintName := constraint.Name
|
||||||
if constraint.Type != models.ForeignKeyConstraint {
|
if constraint.Type != models.ForeignKeyConstraint {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
@@ -787,7 +793,7 @@ func (w *MigrationWriter) generateCommentScripts(model *models.Schema, current *
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Column comments
|
// Column comments
|
||||||
for _, col := range modelTable.Columns {
|
for _, col := range sortColumns(modelTable.Columns) {
|
||||||
if col.Description != "" {
|
if col.Description != "" {
|
||||||
sql, err := w.executor.ExecuteCommentColumn(CommentColumnData{
|
sql, err := w.executor.ExecuteCommentColumn(CommentColumnData{
|
||||||
SchemaName: model.Name,
|
SchemaName: model.Name,
|
||||||
|
|||||||
@@ -545,7 +545,7 @@ func BuildAuditFunctionData(
|
|||||||
|
|
||||||
// Build list of audited columns
|
// Build list of audited columns
|
||||||
auditedColumns := make([]*models.Column, 0)
|
auditedColumns := make([]*models.Column, 0)
|
||||||
for _, col := range table.Columns {
|
for _, col := range sortColumns(table.Columns) {
|
||||||
if col.Name == pk.Name {
|
if col.Name == pk.Name {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -199,7 +199,7 @@ func (w *Writer) GenerateSchemaStatements(schema *models.Schema) ([]string, erro
|
|||||||
for _, table := range schema.Tables {
|
for _, table := range schema.Tables {
|
||||||
// First check for explicit PrimaryKeyConstraint
|
// First check for explicit PrimaryKeyConstraint
|
||||||
var pkConstraint *models.Constraint
|
var pkConstraint *models.Constraint
|
||||||
for _, constraint := range table.Constraints {
|
for _, constraint := range sortConstraints(table.Constraints) {
|
||||||
if constraint.Type == models.PrimaryKeyConstraint {
|
if constraint.Type == models.PrimaryKeyConstraint {
|
||||||
pkConstraint = constraint
|
pkConstraint = constraint
|
||||||
break
|
break
|
||||||
@@ -255,7 +255,7 @@ func (w *Writer) GenerateSchemaStatements(schema *models.Schema) ([]string, erro
|
|||||||
|
|
||||||
// Phase 5: Indexes
|
// Phase 5: Indexes
|
||||||
for _, table := range schema.Tables {
|
for _, table := range schema.Tables {
|
||||||
for _, index := range table.Indexes {
|
for _, index := range sortIndexes(table.Indexes) {
|
||||||
// Skip primary key indexes
|
// Skip primary key indexes
|
||||||
if strings.HasSuffix(index.Name, "_pkey") {
|
if strings.HasSuffix(index.Name, "_pkey") {
|
||||||
continue
|
continue
|
||||||
@@ -298,7 +298,7 @@ func (w *Writer) GenerateSchemaStatements(schema *models.Schema) ([]string, erro
|
|||||||
|
|
||||||
// Phase 5.5: Unique constraints
|
// Phase 5.5: Unique constraints
|
||||||
for _, table := range schema.Tables {
|
for _, table := range schema.Tables {
|
||||||
for _, constraint := range table.Constraints {
|
for _, constraint := range sortConstraints(table.Constraints) {
|
||||||
if constraint.Type != models.UniqueConstraint {
|
if constraint.Type != models.UniqueConstraint {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
@@ -321,7 +321,7 @@ func (w *Writer) GenerateSchemaStatements(schema *models.Schema) ([]string, erro
|
|||||||
|
|
||||||
// Phase 5.7: Check constraints
|
// Phase 5.7: Check constraints
|
||||||
for _, table := range schema.Tables {
|
for _, table := range schema.Tables {
|
||||||
for _, constraint := range table.Constraints {
|
for _, constraint := range sortConstraints(table.Constraints) {
|
||||||
if constraint.Type != models.CheckConstraint {
|
if constraint.Type != models.CheckConstraint {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
@@ -344,7 +344,7 @@ func (w *Writer) GenerateSchemaStatements(schema *models.Schema) ([]string, erro
|
|||||||
|
|
||||||
// Phase 6: Foreign keys
|
// Phase 6: Foreign keys
|
||||||
for _, table := range schema.Tables {
|
for _, table := range schema.Tables {
|
||||||
for _, constraint := range table.Constraints {
|
for _, constraint := range sortConstraints(table.Constraints) {
|
||||||
if constraint.Type != models.ForeignKeyConstraint {
|
if constraint.Type != models.ForeignKeyConstraint {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
@@ -394,7 +394,7 @@ func (w *Writer) GenerateSchemaStatements(schema *models.Schema) ([]string, erro
|
|||||||
statements = append(statements, stmt)
|
statements = append(statements, stmt)
|
||||||
}
|
}
|
||||||
|
|
||||||
for _, column := range table.Columns {
|
for _, column := range sortColumns(table.Columns) {
|
||||||
if column.Comment != "" {
|
if column.Comment != "" {
|
||||||
stmt := fmt.Sprintf("COMMENT ON COLUMN %s.%s IS '%s'",
|
stmt := fmt.Sprintf("COMMENT ON COLUMN %s.%s IS '%s'",
|
||||||
w.qualTable(schema.SQLName(), table.SQLName()), column.SQLName(), escapeQuote(column.Comment))
|
w.qualTable(schema.SQLName(), table.SQLName()), column.SQLName(), escapeQuote(column.Comment))
|
||||||
@@ -866,10 +866,9 @@ func (w *Writer) writePrimaryKeys(schema *models.Schema) error {
|
|||||||
for _, table := range schema.Tables {
|
for _, table := range schema.Tables {
|
||||||
// Find primary key constraint
|
// Find primary key constraint
|
||||||
var pkConstraint *models.Constraint
|
var pkConstraint *models.Constraint
|
||||||
for name, constraint := range table.Constraints {
|
for _, constraint := range sortConstraints(table.Constraints) {
|
||||||
if constraint.Type == models.PrimaryKeyConstraint {
|
if constraint.Type == models.PrimaryKeyConstraint {
|
||||||
pkConstraint = constraint
|
pkConstraint = constraint
|
||||||
_ = name // Use the name variable
|
|
||||||
break
|
break
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -1475,6 +1474,51 @@ func resolveIndexColumn(table *models.Table, colName string) (*models.Column, bo
|
|||||||
return nil, false
|
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
|
// formatStringList formats a list of strings as a SQL-safe comma-separated quoted list
|
||||||
func formatStringList(items []string) string {
|
func formatStringList(items []string) string {
|
||||||
quoted := make([]string, len(items))
|
quoted := make([]string, len(items))
|
||||||
|
|||||||
@@ -549,14 +549,14 @@ func (w *Writer) generateBlockAttributes(table *models.Table) string {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// @@unique for multi-column unique constraints
|
// @@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 {
|
if constraint.Type == models.UniqueConstraint && len(constraint.Columns) > 1 {
|
||||||
fmt.Fprintf(&sb, " @@unique([%s])\n", strings.Join(constraint.Columns, ", "))
|
fmt.Fprintf(&sb, " @@unique([%s])\n", strings.Join(constraint.Columns, ", "))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// @@index for indexes
|
// @@index for indexes
|
||||||
for _, index := range table.Indexes {
|
for _, index := range sortIndexes(table.Indexes) {
|
||||||
if !index.Unique { // Unique indexes are handled by @@unique
|
if !index.Unique { // Unique indexes are handled by @@unique
|
||||||
fmt.Fprintf(&sb, " @@index([%s])\n", strings.Join(index.Columns, ", "))
|
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()
|
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"
|
"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/models"
|
||||||
"git.warky.dev/wdevs/relspecgo/pkg/pgsql"
|
"git.warky.dev/wdevs/relspecgo/pkg/pgsql"
|
||||||
"git.warky.dev/wdevs/relspecgo/pkg/writers"
|
"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)
|
script.Name, script.Priority, script.Sequence)
|
||||||
|
|
||||||
// Execute the SQL script
|
// 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 err != nil {
|
||||||
if ignoreErrors {
|
if ignoreErrors {
|
||||||
fmt.Printf("⚠ Error executing %s: %v (continuing due to --ignore-errors)\n", script.Name, err)
|
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
|
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
|
package sqlexec
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/assetloader"
|
||||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||||
"git.warky.dev/wdevs/relspecgo/pkg/writers"
|
"git.warky.dev/wdevs/relspecgo/pkg/writers"
|
||||||
)
|
)
|
||||||
@@ -216,3 +220,36 @@ func TestWriter_WriteSchema_EmptyScripts(t *testing.T) {
|
|||||||
// // Verify results
|
// // Verify results
|
||||||
// // Cleanup
|
// // 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,6 +4,7 @@ import (
|
|||||||
"bytes"
|
"bytes"
|
||||||
"embed"
|
"embed"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"sort"
|
||||||
"text/template"
|
"text/template"
|
||||||
|
|
||||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||||
@@ -133,15 +134,11 @@ func (te *TemplateExecutor) ExecuteCreateForeignKey(data ConstraintTemplateData)
|
|||||||
|
|
||||||
// BuildTableTemplateData builds TableTemplateData from a models.Table
|
// BuildTableTemplateData builds TableTemplateData from a models.Table
|
||||||
func BuildTableTemplateData(schema string, table *models.Table) TableTemplateData {
|
func BuildTableTemplateData(schema string, table *models.Table) TableTemplateData {
|
||||||
// Get sorted columns
|
columns := sortColumns(table.Columns)
|
||||||
columns := make([]*models.Column, 0, len(table.Columns))
|
|
||||||
for _, col := range table.Columns {
|
|
||||||
columns = append(columns, col)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Find primary key constraint
|
// Find primary key constraint
|
||||||
var pk *models.Constraint
|
var pk *models.Constraint
|
||||||
for _, constraint := range table.Constraints {
|
for _, constraint := range sortConstraints(table.Constraints) {
|
||||||
if constraint.Type == models.PrimaryKeyConstraint {
|
if constraint.Type == models.PrimaryKeyConstraint {
|
||||||
pk = constraint
|
pk = constraint
|
||||||
break
|
break
|
||||||
@@ -151,7 +148,7 @@ func BuildTableTemplateData(schema string, table *models.Table) TableTemplateDat
|
|||||||
// If no explicit primary key constraint, build one from columns with IsPrimaryKey=true
|
// If no explicit primary key constraint, build one from columns with IsPrimaryKey=true
|
||||||
if pk == nil {
|
if pk == nil {
|
||||||
pkCols := []string{}
|
pkCols := []string{}
|
||||||
for _, col := range table.Columns {
|
for _, col := range columns {
|
||||||
if col.IsPrimaryKey {
|
if col.IsPrimaryKey {
|
||||||
pkCols = append(pkCols, col.Name)
|
pkCols = append(pkCols, col.Name)
|
||||||
}
|
}
|
||||||
@@ -172,3 +169,48 @@ func BuildTableTemplateData(schema string, table *models.Table) TableTemplateDat
|
|||||||
PrimaryKey: pk,
|
PrimaryKey: pk,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// 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
|
||||||
|
}
|
||||||
|
|||||||
@@ -143,7 +143,7 @@ func (w *Writer) writeTable(schema string, table *models.Table) error {
|
|||||||
|
|
||||||
// writeIndexes writes indexes for a table
|
// writeIndexes writes indexes for a table
|
||||||
func (w *Writer) writeIndexes(schema string, table *models.Table) error {
|
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
|
// Skip primary key indexes
|
||||||
if strings.HasSuffix(index.Name, "_pkey") {
|
if strings.HasSuffix(index.Name, "_pkey") {
|
||||||
continue
|
continue
|
||||||
@@ -174,7 +174,7 @@ func (w *Writer) writeIndexes(schema string, table *models.Table) error {
|
|||||||
|
|
||||||
// writeUniqueConstraints writes unique constraints as unique indexes
|
// writeUniqueConstraints writes unique constraints as unique indexes
|
||||||
func (w *Writer) writeUniqueConstraints(schema string, table *models.Table) error {
|
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 {
|
if constraint.Type != models.UniqueConstraint {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
@@ -195,7 +195,7 @@ func (w *Writer) writeUniqueConstraints(schema string, table *models.Table) erro
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Also handle unique indexes from the Indexes map
|
// Also handle unique indexes from the Indexes map
|
||||||
for _, index := range table.Indexes {
|
for _, index := range sortIndexes(table.Indexes) {
|
||||||
if !index.Unique {
|
if !index.Unique {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
@@ -232,7 +232,7 @@ func (w *Writer) writeUniqueConstraints(schema string, table *models.Table) erro
|
|||||||
|
|
||||||
// writeCheckConstraints writes check constraints as comments
|
// writeCheckConstraints writes check constraints as comments
|
||||||
func (w *Writer) writeCheckConstraints(schema string, table *models.Table) error {
|
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 {
|
if constraint.Type != models.CheckConstraint {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
@@ -257,7 +257,7 @@ func (w *Writer) writeCheckConstraints(schema string, table *models.Table) error
|
|||||||
|
|
||||||
// writeForeignKeys writes foreign keys as comments
|
// writeForeignKeys writes foreign keys as comments
|
||||||
func (w *Writer) writeForeignKeys(schema string, table *models.Table) error {
|
func (w *Writer) writeForeignKeys(schema string, table *models.Table) error {
|
||||||
for _, constraint := range table.Constraints {
|
for _, constraint := range sortConstraints(table.Constraints) {
|
||||||
if constraint.Type != models.ForeignKeyConstraint {
|
if constraint.Type != models.ForeignKeyConstraint {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -60,7 +60,7 @@ func Has(m interface{}, key interface{}) bool {
|
|||||||
v := reflect.ValueOf(m)
|
v := reflect.ValueOf(m)
|
||||||
|
|
||||||
// Dereference pointers
|
// Dereference pointers
|
||||||
for v.Kind() == reflect.Ptr {
|
for v.Kind() == reflect.Pointer {
|
||||||
if v.IsNil() {
|
if v.IsNil() {
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
@@ -102,7 +102,7 @@ func Merge(maps ...interface{}) map[interface{}]interface{} {
|
|||||||
v := reflect.ValueOf(m)
|
v := reflect.ValueOf(m)
|
||||||
|
|
||||||
// Dereference pointers
|
// Dereference pointers
|
||||||
for v.Kind() == reflect.Ptr {
|
for v.Kind() == reflect.Pointer {
|
||||||
if v.IsNil() {
|
if v.IsNil() {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
@@ -129,7 +129,7 @@ func Pick(m interface{}, keys ...interface{}) map[interface{}]interface{} {
|
|||||||
v := reflect.ValueOf(m)
|
v := reflect.ValueOf(m)
|
||||||
|
|
||||||
// Dereference pointers
|
// Dereference pointers
|
||||||
for v.Kind() == reflect.Ptr {
|
for v.Kind() == reflect.Pointer {
|
||||||
if v.IsNil() {
|
if v.IsNil() {
|
||||||
return result
|
return result
|
||||||
}
|
}
|
||||||
@@ -158,7 +158,7 @@ func Omit(m interface{}, keys ...interface{}) map[interface{}]interface{} {
|
|||||||
v := reflect.ValueOf(m)
|
v := reflect.ValueOf(m)
|
||||||
|
|
||||||
// Dereference pointers
|
// Dereference pointers
|
||||||
for v.Kind() == reflect.Ptr {
|
for v.Kind() == reflect.Pointer {
|
||||||
if v.IsNil() {
|
if v.IsNil() {
|
||||||
return result
|
return result
|
||||||
}
|
}
|
||||||
@@ -237,7 +237,7 @@ func Pluck(slice interface{}, field string) []interface{} {
|
|||||||
v := reflect.ValueOf(slice)
|
v := reflect.ValueOf(slice)
|
||||||
|
|
||||||
// Dereference pointers
|
// Dereference pointers
|
||||||
for v.Kind() == reflect.Ptr {
|
for v.Kind() == reflect.Pointer {
|
||||||
if v.IsNil() {
|
if v.IsNil() {
|
||||||
return []interface{}{}
|
return []interface{}{}
|
||||||
}
|
}
|
||||||
@@ -253,13 +253,18 @@ func Pluck(slice interface{}, field string) []interface{} {
|
|||||||
item := v.Index(i)
|
item := v.Index(i)
|
||||||
|
|
||||||
// Dereference item pointers
|
// Dereference item pointers
|
||||||
for item.Kind() == reflect.Ptr {
|
nilItem := false
|
||||||
|
for item.Kind() == reflect.Pointer {
|
||||||
if item.IsNil() {
|
if item.IsNil() {
|
||||||
result = append(result, nil)
|
result = append(result, nil)
|
||||||
continue
|
nilItem = true
|
||||||
|
break
|
||||||
}
|
}
|
||||||
item = item.Elem()
|
item = item.Elem()
|
||||||
}
|
}
|
||||||
|
if nilItem {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
switch item.Kind() {
|
switch item.Kind() {
|
||||||
case reflect.Struct:
|
case reflect.Struct:
|
||||||
|
|||||||
@@ -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
|
// generateManyToManyRelations generates @ManyToMany fields
|
||||||
func (w *Writer) generateManyToManyRelations(table *models.Table, schema *models.Schema, joinTables map[string]bool, sb *strings.Builder) {
|
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)
|
joinTable := w.findTable(joinTableName, schema)
|
||||||
if joinTable == nil {
|
if joinTable == nil {
|
||||||
continue
|
continue
|
||||||
|
|||||||
+36
-4
@@ -23,13 +23,35 @@ type Writer interface {
|
|||||||
// NullableType constants control which Go package is used for nullable column types
|
// NullableType constants control which Go package is used for nullable column types
|
||||||
// in code-generation writers (Bun, GORM).
|
// in code-generation writers (Bun, GORM).
|
||||||
const (
|
const (
|
||||||
// NullableTypeResolveSpec uses github.com/bitechdev/ResolveSpec/pkg/spectypes
|
// NullableTypeSqlTypes uses git.warky.dev/wdevs/relspecgo/pkg/sqltypes
|
||||||
// (SqlString, SqlInt32, SqlVector, SqlStringArray, …). This is the default.
|
// (SqlString, SqlInt32, SqlVector, SqlStringArray, …).
|
||||||
NullableTypeResolveSpec = "resolvespec"
|
NullableTypeSqlTypes = "sqltypes"
|
||||||
|
|
||||||
// NullableTypeStdlib uses the standard library database/sql nullable types
|
// NullableTypeStdlib uses the standard library database/sql nullable types
|
||||||
// (sql.NullString, sql.NullInt32, …) and plain Go slices for arrays.
|
// (sql.NullString, sql.NullInt32, …) and plain Go slices for arrays.
|
||||||
NullableTypeStdlib = "stdlib"
|
NullableTypeStdlib = "stdlib"
|
||||||
|
|
||||||
|
// NullableTypeBaselib uses plain Go pointer types for nullable columns
|
||||||
|
// (*string, *int32, *time.Time, …) and plain Go slices for arrays.
|
||||||
|
// No external imports are required beyond the standard library. This is the default.
|
||||||
|
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
|
// WriterOptions contains common options for writers
|
||||||
@@ -47,10 +69,20 @@ type WriterOptions struct {
|
|||||||
|
|
||||||
// NullableTypes selects the Go type package used for nullable columns in
|
// NullableTypes selects the Go type package used for nullable columns in
|
||||||
// code-generation writers (bun, gorm). Accepted values:
|
// code-generation writers (bun, gorm). Accepted values:
|
||||||
// "resolvespec" (default) — github.com/bitechdev/ResolveSpec/pkg/spectypes
|
// "sqltypes" — git.warky.dev/wdevs/relspecgo/pkg/sqltypes
|
||||||
// "stdlib" — database/sql (sql.NullString, sql.NullInt32, …)
|
// "stdlib" — database/sql (sql.NullString, sql.NullInt32, …)
|
||||||
|
// "baselib" (default) — plain Go pointer types (*string, *int32, …)
|
||||||
NullableTypes string
|
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 enables Prisma 7-specific output for Prisma writers.
|
||||||
Prisma7 bool
|
Prisma7 bool
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
@@ -0,0 +1 @@
|
|||||||
|
PNG_PLACEHOLDER
|
||||||
@@ -0,0 +1,4 @@
|
|||||||
|
- file: banner.txt
|
||||||
|
call: INSERT INTO banners (data) VALUES (:bytes)
|
||||||
|
params:
|
||||||
|
owner_id: "1"
|
||||||
@@ -0,0 +1 @@
|
|||||||
|
BANNER CONTENT
|
||||||
+52
@@ -0,0 +1,52 @@
|
|||||||
|
package pgxpool
|
||||||
|
|
||||||
|
import (
|
||||||
|
"github.com/jackc/pgx/v5"
|
||||||
|
"github.com/jackc/pgx/v5/pgconn"
|
||||||
|
)
|
||||||
|
|
||||||
|
type errBatchResults struct {
|
||||||
|
err error
|
||||||
|
}
|
||||||
|
|
||||||
|
func (br errBatchResults) Exec() (pgconn.CommandTag, error) {
|
||||||
|
return pgconn.CommandTag{}, br.err
|
||||||
|
}
|
||||||
|
|
||||||
|
func (br errBatchResults) Query() (pgx.Rows, error) {
|
||||||
|
return errRows{err: br.err}, br.err
|
||||||
|
}
|
||||||
|
|
||||||
|
func (br errBatchResults) QueryRow() pgx.Row {
|
||||||
|
return errRow{err: br.err}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (br errBatchResults) Close() error {
|
||||||
|
return br.err
|
||||||
|
}
|
||||||
|
|
||||||
|
type poolBatchResults struct {
|
||||||
|
br pgx.BatchResults
|
||||||
|
c *Conn
|
||||||
|
}
|
||||||
|
|
||||||
|
func (br *poolBatchResults) Exec() (pgconn.CommandTag, error) {
|
||||||
|
return br.br.Exec()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (br *poolBatchResults) Query() (pgx.Rows, error) {
|
||||||
|
return br.br.Query()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (br *poolBatchResults) QueryRow() pgx.Row {
|
||||||
|
return br.br.QueryRow()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (br *poolBatchResults) Close() error {
|
||||||
|
err := br.br.Close()
|
||||||
|
if br.c != nil {
|
||||||
|
br.c.Release()
|
||||||
|
br.c = nil
|
||||||
|
}
|
||||||
|
return err
|
||||||
|
}
|
||||||
+134
@@ -0,0 +1,134 @@
|
|||||||
|
package pgxpool
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"sync/atomic"
|
||||||
|
|
||||||
|
"github.com/jackc/pgx/v5"
|
||||||
|
"github.com/jackc/pgx/v5/pgconn"
|
||||||
|
"github.com/jackc/puddle/v2"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Conn is an acquired *pgx.Conn from a Pool.
|
||||||
|
type Conn struct {
|
||||||
|
res *puddle.Resource[*connResource]
|
||||||
|
p *Pool
|
||||||
|
}
|
||||||
|
|
||||||
|
// Release returns c to the pool it was acquired from. Once Release has been called, other methods must not be called.
|
||||||
|
// However, it is safe to call Release multiple times. Subsequent calls after the first will be ignored.
|
||||||
|
func (c *Conn) Release() {
|
||||||
|
if c.res == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
conn := c.Conn()
|
||||||
|
res := c.res
|
||||||
|
c.res = nil
|
||||||
|
|
||||||
|
if c.p.releaseTracer != nil {
|
||||||
|
c.p.releaseTracer.TraceRelease(c.p, TraceReleaseData{Conn: conn})
|
||||||
|
}
|
||||||
|
|
||||||
|
if conn.IsClosed() || conn.PgConn().IsBusy() || conn.PgConn().TxStatus() != 'I' {
|
||||||
|
res.Destroy()
|
||||||
|
// Signal to the health check to run since we just destroyed a connections
|
||||||
|
// and we might be below minConns now
|
||||||
|
c.p.triggerHealthCheck()
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// If the pool is consistently being used, we might never get to check the
|
||||||
|
// lifetime of a connection since we only check idle connections in checkConnsHealth
|
||||||
|
// so we also check the lifetime here and force a health check
|
||||||
|
if c.p.isExpired(res) {
|
||||||
|
atomic.AddInt64(&c.p.lifetimeDestroyCount, 1)
|
||||||
|
res.Destroy()
|
||||||
|
// Signal to the health check to run since we just destroyed a connections
|
||||||
|
// and we might be below minConns now
|
||||||
|
c.p.triggerHealthCheck()
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if c.p.afterRelease == nil {
|
||||||
|
res.Release()
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
go func() {
|
||||||
|
if c.p.afterRelease(conn) {
|
||||||
|
res.Release()
|
||||||
|
} else {
|
||||||
|
res.Destroy()
|
||||||
|
// Signal to the health check to run since we just destroyed a connections
|
||||||
|
// and we might be below minConns now
|
||||||
|
c.p.triggerHealthCheck()
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Hijack assumes ownership of the connection from the pool. Caller is responsible for closing the connection. Hijack
|
||||||
|
// will panic if called on an already released or hijacked connection.
|
||||||
|
func (c *Conn) Hijack() *pgx.Conn {
|
||||||
|
if c.res == nil {
|
||||||
|
panic("cannot hijack already released or hijacked connection")
|
||||||
|
}
|
||||||
|
|
||||||
|
conn := c.Conn()
|
||||||
|
res := c.res
|
||||||
|
c.res = nil
|
||||||
|
|
||||||
|
res.Hijack()
|
||||||
|
|
||||||
|
return conn
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *Conn) Exec(ctx context.Context, sql string, arguments ...any) (pgconn.CommandTag, error) {
|
||||||
|
return c.Conn().Exec(ctx, sql, arguments...)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *Conn) Query(ctx context.Context, sql string, args ...any) (pgx.Rows, error) {
|
||||||
|
return c.Conn().Query(ctx, sql, args...)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *Conn) QueryRow(ctx context.Context, sql string, args ...any) pgx.Row {
|
||||||
|
return c.Conn().QueryRow(ctx, sql, args...)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *Conn) SendBatch(ctx context.Context, b *pgx.Batch) pgx.BatchResults {
|
||||||
|
return c.Conn().SendBatch(ctx, b)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *Conn) CopyFrom(ctx context.Context, tableName pgx.Identifier, columnNames []string, rowSrc pgx.CopyFromSource) (int64, error) {
|
||||||
|
return c.Conn().CopyFrom(ctx, tableName, columnNames, rowSrc)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Begin starts a transaction block from the *Conn without explicitly setting a transaction mode (see BeginTx with TxOptions if transaction mode is required).
|
||||||
|
func (c *Conn) Begin(ctx context.Context) (pgx.Tx, error) {
|
||||||
|
return c.Conn().Begin(ctx)
|
||||||
|
}
|
||||||
|
|
||||||
|
// BeginTx starts a transaction block from the *Conn with txOptions determining the transaction mode.
|
||||||
|
func (c *Conn) BeginTx(ctx context.Context, txOptions pgx.TxOptions) (pgx.Tx, error) {
|
||||||
|
return c.Conn().BeginTx(ctx, txOptions)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *Conn) Ping(ctx context.Context) error {
|
||||||
|
return c.Conn().Ping(ctx)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *Conn) Conn() *pgx.Conn {
|
||||||
|
return c.connResource().conn
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *Conn) connResource() *connResource {
|
||||||
|
return c.res.Value()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *Conn) getPoolRow(r pgx.Row) *poolRow {
|
||||||
|
return c.connResource().getPoolRow(c, r)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *Conn) getPoolRows(r pgx.Rows) *poolRows {
|
||||||
|
return c.connResource().getPoolRows(c, r)
|
||||||
|
}
|
||||||
+27
@@ -0,0 +1,27 @@
|
|||||||
|
// Package pgxpool is a concurrency-safe connection pool for pgx.
|
||||||
|
/*
|
||||||
|
pgxpool implements a nearly identical interface to pgx connections.
|
||||||
|
|
||||||
|
Creating a Pool
|
||||||
|
|
||||||
|
The primary way of creating a pool is with [pgxpool.New]:
|
||||||
|
|
||||||
|
pool, err := pgxpool.New(context.Background(), os.Getenv("DATABASE_URL"))
|
||||||
|
|
||||||
|
The database connection string can be in URL or keyword/value format. PostgreSQL settings, pgx settings, and pool settings can be
|
||||||
|
specified here. In addition, a config struct can be created by [ParseConfig].
|
||||||
|
|
||||||
|
config, err := pgxpool.ParseConfig(os.Getenv("DATABASE_URL"))
|
||||||
|
if err != nil {
|
||||||
|
// ...
|
||||||
|
}
|
||||||
|
config.AfterConnect = func(ctx context.Context, conn *pgx.Conn) error {
|
||||||
|
// do something with every new connection
|
||||||
|
}
|
||||||
|
|
||||||
|
pool, err := pgxpool.NewWithConfig(context.Background(), config)
|
||||||
|
|
||||||
|
A pool returns without waiting for any connections to be established. Acquire a connection immediately after creating
|
||||||
|
the pool to check if a connection can successfully be established.
|
||||||
|
*/
|
||||||
|
package pgxpool
|
||||||
+832
@@ -0,0 +1,832 @@
|
|||||||
|
package pgxpool
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"math/rand/v2"
|
||||||
|
"runtime"
|
||||||
|
"strconv"
|
||||||
|
"sync"
|
||||||
|
"sync/atomic"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/jackc/pgx/v5"
|
||||||
|
"github.com/jackc/pgx/v5/pgconn"
|
||||||
|
"github.com/jackc/puddle/v2"
|
||||||
|
)
|
||||||
|
|
||||||
|
var (
|
||||||
|
defaultMaxConns = int32(4)
|
||||||
|
defaultMinConns = int32(0)
|
||||||
|
defaultMinIdleConns = int32(0)
|
||||||
|
defaultMaxConnLifetime = time.Hour
|
||||||
|
defaultMaxConnIdleTime = time.Minute * 30
|
||||||
|
defaultHealthCheckPeriod = time.Minute
|
||||||
|
)
|
||||||
|
|
||||||
|
type connResource struct {
|
||||||
|
conn *pgx.Conn
|
||||||
|
conns []Conn
|
||||||
|
poolRows []poolRow
|
||||||
|
poolRowss []poolRows
|
||||||
|
maxAgeTime time.Time
|
||||||
|
}
|
||||||
|
|
||||||
|
func (cr *connResource) getConn(p *Pool, res *puddle.Resource[*connResource]) *Conn {
|
||||||
|
if len(cr.conns) == 0 {
|
||||||
|
cr.conns = make([]Conn, 128)
|
||||||
|
}
|
||||||
|
|
||||||
|
c := &cr.conns[len(cr.conns)-1]
|
||||||
|
cr.conns = cr.conns[0 : len(cr.conns)-1]
|
||||||
|
|
||||||
|
c.res = res
|
||||||
|
c.p = p
|
||||||
|
|
||||||
|
return c
|
||||||
|
}
|
||||||
|
|
||||||
|
func (cr *connResource) getPoolRow(c *Conn, r pgx.Row) *poolRow {
|
||||||
|
if len(cr.poolRows) == 0 {
|
||||||
|
cr.poolRows = make([]poolRow, 128)
|
||||||
|
}
|
||||||
|
|
||||||
|
pr := &cr.poolRows[len(cr.poolRows)-1]
|
||||||
|
cr.poolRows = cr.poolRows[0 : len(cr.poolRows)-1]
|
||||||
|
|
||||||
|
pr.c = c
|
||||||
|
pr.r = r
|
||||||
|
|
||||||
|
return pr
|
||||||
|
}
|
||||||
|
|
||||||
|
func (cr *connResource) getPoolRows(c *Conn, r pgx.Rows) *poolRows {
|
||||||
|
if len(cr.poolRowss) == 0 {
|
||||||
|
cr.poolRowss = make([]poolRows, 128)
|
||||||
|
}
|
||||||
|
|
||||||
|
pr := &cr.poolRowss[len(cr.poolRowss)-1]
|
||||||
|
cr.poolRowss = cr.poolRowss[0 : len(cr.poolRowss)-1]
|
||||||
|
|
||||||
|
pr.c = c
|
||||||
|
pr.r = r
|
||||||
|
|
||||||
|
return pr
|
||||||
|
}
|
||||||
|
|
||||||
|
// Pool allows for connection reuse.
|
||||||
|
type Pool struct {
|
||||||
|
// 64 bit fields accessed with atomics must be at beginning of struct to guarantee alignment for certain 32-bit
|
||||||
|
// architectures. See BUGS section of https://pkg.go.dev/sync/atomic and https://github.com/jackc/pgx/issues/1288.
|
||||||
|
newConnsCount int64
|
||||||
|
lifetimeDestroyCount int64
|
||||||
|
idleDestroyCount int64
|
||||||
|
|
||||||
|
p *puddle.Pool[*connResource]
|
||||||
|
config *Config
|
||||||
|
beforeConnect func(context.Context, *pgx.ConnConfig) error
|
||||||
|
afterConnect func(context.Context, *pgx.Conn) error
|
||||||
|
prepareConn func(context.Context, *pgx.Conn) (bool, error)
|
||||||
|
afterRelease func(*pgx.Conn) bool
|
||||||
|
beforeClose func(*pgx.Conn)
|
||||||
|
shouldPing func(context.Context, ShouldPingParams) bool
|
||||||
|
minConns int32
|
||||||
|
minIdleConns int32
|
||||||
|
maxConns int32
|
||||||
|
maxConnLifetime time.Duration
|
||||||
|
maxConnLifetimeJitter time.Duration
|
||||||
|
maxConnIdleTime time.Duration
|
||||||
|
healthCheckPeriod time.Duration
|
||||||
|
pingTimeout time.Duration
|
||||||
|
|
||||||
|
healthCheckMu sync.Mutex
|
||||||
|
healthCheckTimer *time.Timer
|
||||||
|
|
||||||
|
healthCheckChan chan struct{}
|
||||||
|
|
||||||
|
acquireTracer AcquireTracer
|
||||||
|
releaseTracer ReleaseTracer
|
||||||
|
|
||||||
|
closeOnce sync.Once
|
||||||
|
closeChan chan struct{}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ShouldPingParams are the parameters passed to ShouldPing.
|
||||||
|
type ShouldPingParams struct {
|
||||||
|
Conn *pgx.Conn
|
||||||
|
IdleDuration time.Duration
|
||||||
|
}
|
||||||
|
|
||||||
|
// Config is the configuration struct for creating a pool. It must be created by [ParseConfig] and then it can be
|
||||||
|
// modified.
|
||||||
|
type Config struct {
|
||||||
|
ConnConfig *pgx.ConnConfig
|
||||||
|
|
||||||
|
// BeforeConnect is called before a new connection is made. It is passed a copy of the underlying [pgx.ConnConfig] and
|
||||||
|
// will not impact any existing open connections.
|
||||||
|
BeforeConnect func(context.Context, *pgx.ConnConfig) error
|
||||||
|
|
||||||
|
// AfterConnect is called after a connection is established, but before it is added to the pool.
|
||||||
|
AfterConnect func(context.Context, *pgx.Conn) error
|
||||||
|
|
||||||
|
// BeforeAcquire is called before a connection is acquired from the pool. It must return true to allow the
|
||||||
|
// acquisition or false to indicate that the connection should be destroyed and a different connection should be
|
||||||
|
// acquired.
|
||||||
|
//
|
||||||
|
// Deprecated: Use PrepareConn instead. If both PrepareConn and BeforeAcquire are set, PrepareConn will take
|
||||||
|
// precedence, ignoring BeforeAcquire.
|
||||||
|
BeforeAcquire func(context.Context, *pgx.Conn) bool
|
||||||
|
|
||||||
|
// PrepareConn is called before a connection is acquired from the pool. If this function returns true, the connection
|
||||||
|
// is considered valid, otherwise the connection is destroyed. If the function returns a non-nil error, the instigating
|
||||||
|
// query will fail with the returned error.
|
||||||
|
//
|
||||||
|
// Specifically, this means that:
|
||||||
|
//
|
||||||
|
// - If it returns true and a nil error, the query proceeds as normal.
|
||||||
|
// - If it returns true and an error, the connection will be returned to the pool, and the instigating query will fail with the returned error.
|
||||||
|
// - If it returns false, and an error, the connection will be destroyed, and the query will fail with the returned error.
|
||||||
|
// - If it returns false and a nil error, the connection will be destroyed, and the instigating query will be retried on a new connection.
|
||||||
|
PrepareConn func(context.Context, *pgx.Conn) (bool, error)
|
||||||
|
|
||||||
|
// AfterRelease is called after a connection is released, but before it is returned to the pool. It must return true to
|
||||||
|
// return the connection to the pool or false to destroy the connection.
|
||||||
|
AfterRelease func(*pgx.Conn) bool
|
||||||
|
|
||||||
|
// BeforeClose is called right before a connection is closed and removed from the pool.
|
||||||
|
BeforeClose func(*pgx.Conn)
|
||||||
|
|
||||||
|
// ShouldPing is called after a connection is acquired from the pool. If it returns true, the connection is pinged to check for liveness.
|
||||||
|
// If this func is not set, the default behavior is to ping connections that have been idle for at least 1 second.
|
||||||
|
ShouldPing func(context.Context, ShouldPingParams) bool
|
||||||
|
|
||||||
|
// MaxConnLifetime is the duration since creation after which a connection will be automatically closed.
|
||||||
|
MaxConnLifetime time.Duration
|
||||||
|
|
||||||
|
// MaxConnLifetimeJitter is the duration after MaxConnLifetime to randomly decide to close a connection.
|
||||||
|
// This helps prevent all connections from being closed at the exact same time, starving the pool.
|
||||||
|
MaxConnLifetimeJitter time.Duration
|
||||||
|
|
||||||
|
// MaxConnIdleTime is the duration after which an idle connection will be automatically closed by the health check.
|
||||||
|
MaxConnIdleTime time.Duration
|
||||||
|
|
||||||
|
// PingTimeout is the maximum amount of time to wait for a connection to pong before considering it as unhealthy and
|
||||||
|
// destroying it. If zero, the default is no timeout.
|
||||||
|
PingTimeout time.Duration
|
||||||
|
|
||||||
|
// MaxConns is the maximum size of the pool. The default is the greater of 4 or runtime.NumCPU().
|
||||||
|
MaxConns int32
|
||||||
|
|
||||||
|
// MinConns is the minimum size of the pool. After connection closes, the pool might dip below MinConns. A low
|
||||||
|
// number of MinConns might mean the pool is empty after MaxConnLifetime until the health check has a chance
|
||||||
|
// to create new connections.
|
||||||
|
MinConns int32
|
||||||
|
|
||||||
|
// MinIdleConns is the minimum number of idle connections in the pool. You can increase this to ensure that
|
||||||
|
// there are always idle connections available. This can help reduce tail latencies during request processing,
|
||||||
|
// as you can avoid the latency of establishing a new connection while handling requests. It is superior
|
||||||
|
// to MinConns for this purpose.
|
||||||
|
// Similar to MinConns, the pool might temporarily dip below MinIdleConns after connection closes.
|
||||||
|
MinIdleConns int32
|
||||||
|
|
||||||
|
// HealthCheckPeriod is the duration between checks of the health of idle connections.
|
||||||
|
HealthCheckPeriod time.Duration
|
||||||
|
|
||||||
|
createdByParseConfig bool // Used to enforce created by ParseConfig rule.
|
||||||
|
}
|
||||||
|
|
||||||
|
// Copy returns a deep copy of the config that is safe to use and modify.
|
||||||
|
// The only exception is the tls.Config:
|
||||||
|
// according to the tls.Config docs it must not be modified after creation.
|
||||||
|
func (c *Config) Copy() *Config {
|
||||||
|
newConfig := new(Config)
|
||||||
|
*newConfig = *c
|
||||||
|
newConfig.ConnConfig = c.ConnConfig.Copy()
|
||||||
|
return newConfig
|
||||||
|
}
|
||||||
|
|
||||||
|
// ConnString returns the connection string as parsed by pgxpool.ParseConfig into pgxpool.Config.
|
||||||
|
func (c *Config) ConnString() string { return c.ConnConfig.ConnString() }
|
||||||
|
|
||||||
|
// New creates a new Pool. See [ParseConfig] for information on connString format.
|
||||||
|
func New(ctx context.Context, connString string) (*Pool, error) {
|
||||||
|
config, err := ParseConfig(connString)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
return NewWithConfig(ctx, config)
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewWithConfig creates a new [Pool]. config must have been created by [ParseConfig].
|
||||||
|
func NewWithConfig(ctx context.Context, config *Config) (*Pool, error) {
|
||||||
|
// Default values are set in ParseConfig. Enforce initial creation by ParseConfig rather than setting defaults from
|
||||||
|
// zero values.
|
||||||
|
if !config.createdByParseConfig {
|
||||||
|
panic("config must be created by ParseConfig")
|
||||||
|
}
|
||||||
|
|
||||||
|
prepareConn := config.PrepareConn
|
||||||
|
if prepareConn == nil && config.BeforeAcquire != nil {
|
||||||
|
prepareConn = func(ctx context.Context, conn *pgx.Conn) (bool, error) {
|
||||||
|
return config.BeforeAcquire(ctx, conn), nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
p := &Pool{
|
||||||
|
config: config,
|
||||||
|
beforeConnect: config.BeforeConnect,
|
||||||
|
afterConnect: config.AfterConnect,
|
||||||
|
prepareConn: prepareConn,
|
||||||
|
afterRelease: config.AfterRelease,
|
||||||
|
beforeClose: config.BeforeClose,
|
||||||
|
minConns: config.MinConns,
|
||||||
|
minIdleConns: config.MinIdleConns,
|
||||||
|
maxConns: config.MaxConns,
|
||||||
|
maxConnLifetime: config.MaxConnLifetime,
|
||||||
|
maxConnLifetimeJitter: config.MaxConnLifetimeJitter,
|
||||||
|
maxConnIdleTime: config.MaxConnIdleTime,
|
||||||
|
pingTimeout: config.PingTimeout,
|
||||||
|
healthCheckPeriod: config.HealthCheckPeriod,
|
||||||
|
healthCheckChan: make(chan struct{}, 1),
|
||||||
|
closeChan: make(chan struct{}),
|
||||||
|
}
|
||||||
|
|
||||||
|
if t, ok := config.ConnConfig.Tracer.(AcquireTracer); ok {
|
||||||
|
p.acquireTracer = t
|
||||||
|
}
|
||||||
|
|
||||||
|
if t, ok := config.ConnConfig.Tracer.(ReleaseTracer); ok {
|
||||||
|
p.releaseTracer = t
|
||||||
|
}
|
||||||
|
|
||||||
|
if config.ShouldPing != nil {
|
||||||
|
p.shouldPing = config.ShouldPing
|
||||||
|
} else {
|
||||||
|
p.shouldPing = func(ctx context.Context, params ShouldPingParams) bool {
|
||||||
|
return params.IdleDuration > time.Second
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
var err error
|
||||||
|
p.p, err = puddle.NewPool(
|
||||||
|
&puddle.Config[*connResource]{
|
||||||
|
Constructor: func(ctx context.Context) (*connResource, error) {
|
||||||
|
atomic.AddInt64(&p.newConnsCount, 1)
|
||||||
|
connConfig := p.config.ConnConfig.Copy()
|
||||||
|
|
||||||
|
// Connection will continue in background even if Acquire is canceled. Ensure that a connect won't hang forever.
|
||||||
|
if connConfig.ConnectTimeout <= 0 {
|
||||||
|
connConfig.ConnectTimeout = 2 * time.Minute
|
||||||
|
}
|
||||||
|
|
||||||
|
if p.beforeConnect != nil {
|
||||||
|
if err := p.beforeConnect(ctx, connConfig); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
conn, err := pgx.ConnectConfig(ctx, connConfig)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
if p.afterConnect != nil {
|
||||||
|
err = p.afterConnect(ctx, conn)
|
||||||
|
if err != nil {
|
||||||
|
conn.Close(ctx)
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
jitterSecs := rand.Float64() * config.MaxConnLifetimeJitter.Seconds()
|
||||||
|
maxAgeTime := time.Now().Add(config.MaxConnLifetime).Add(time.Duration(jitterSecs) * time.Second)
|
||||||
|
|
||||||
|
cr := &connResource{
|
||||||
|
conn: conn,
|
||||||
|
conns: make([]Conn, 64),
|
||||||
|
poolRows: make([]poolRow, 64),
|
||||||
|
poolRowss: make([]poolRows, 64),
|
||||||
|
maxAgeTime: maxAgeTime,
|
||||||
|
}
|
||||||
|
|
||||||
|
return cr, nil
|
||||||
|
},
|
||||||
|
Destructor: func(value *connResource) {
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second)
|
||||||
|
conn := value.conn
|
||||||
|
if p.beforeClose != nil {
|
||||||
|
p.beforeClose(conn)
|
||||||
|
}
|
||||||
|
conn.Close(ctx)
|
||||||
|
select {
|
||||||
|
case <-conn.PgConn().CleanupDone():
|
||||||
|
case <-ctx.Done():
|
||||||
|
}
|
||||||
|
cancel()
|
||||||
|
},
|
||||||
|
MaxSize: config.MaxConns,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
go func() {
|
||||||
|
targetIdleResources := max(int(p.minConns), int(p.minIdleConns))
|
||||||
|
p.createIdleResources(ctx, targetIdleResources)
|
||||||
|
p.backgroundHealthCheck()
|
||||||
|
}()
|
||||||
|
|
||||||
|
return p, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// ParseConfig builds a Config from connString. It parses connString with the same behavior as [pgx.ParseConfig] with the
|
||||||
|
// addition of the following variables:
|
||||||
|
//
|
||||||
|
// - pool_max_conns: integer greater than 0 (default 4)
|
||||||
|
// - pool_min_conns: integer 0 or greater (default 0)
|
||||||
|
// - pool_max_conn_lifetime: duration string (default 1 hour)
|
||||||
|
// - pool_max_conn_idle_time: duration string (default 30 minutes)
|
||||||
|
// - pool_health_check_period: duration string (default 1 minute)
|
||||||
|
// - pool_max_conn_lifetime_jitter: duration string (default 0)
|
||||||
|
//
|
||||||
|
// See Config for definitions of these arguments.
|
||||||
|
//
|
||||||
|
// # Example Keyword/Value
|
||||||
|
// user=jack password=secret host=pg.example.com port=5432 dbname=mydb sslmode=verify-ca pool_max_conns=10 pool_max_conn_lifetime=1h30m
|
||||||
|
//
|
||||||
|
// # Example URL
|
||||||
|
// postgres://jack:secret@pg.example.com:5432/mydb?sslmode=verify-ca&pool_max_conns=10&pool_max_conn_lifetime=1h30m
|
||||||
|
func ParseConfig(connString string) (*Config, error) {
|
||||||
|
connConfig, err := pgx.ParseConfig(connString)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
config := &Config{
|
||||||
|
ConnConfig: connConfig,
|
||||||
|
createdByParseConfig: true,
|
||||||
|
}
|
||||||
|
|
||||||
|
if s, ok := config.ConnConfig.Config.RuntimeParams["pool_max_conns"]; ok {
|
||||||
|
delete(connConfig.Config.RuntimeParams, "pool_max_conns")
|
||||||
|
n, err := strconv.ParseInt(s, 10, 32)
|
||||||
|
if err != nil {
|
||||||
|
return nil, pgconn.NewParseConfigError(connString, "cannot parse pool_max_conns", err)
|
||||||
|
}
|
||||||
|
if n < 1 {
|
||||||
|
return nil, pgconn.NewParseConfigError(connString, "pool_max_conns too small", err)
|
||||||
|
}
|
||||||
|
config.MaxConns = int32(n)
|
||||||
|
} else {
|
||||||
|
config.MaxConns = defaultMaxConns
|
||||||
|
if numCPU := int32(runtime.NumCPU()); numCPU > config.MaxConns {
|
||||||
|
config.MaxConns = numCPU
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if s, ok := config.ConnConfig.Config.RuntimeParams["pool_min_conns"]; ok {
|
||||||
|
delete(connConfig.Config.RuntimeParams, "pool_min_conns")
|
||||||
|
n, err := strconv.ParseInt(s, 10, 32)
|
||||||
|
if err != nil {
|
||||||
|
return nil, pgconn.NewParseConfigError(connString, "cannot parse pool_min_conns", err)
|
||||||
|
}
|
||||||
|
config.MinConns = int32(n)
|
||||||
|
} else {
|
||||||
|
config.MinConns = defaultMinConns
|
||||||
|
}
|
||||||
|
|
||||||
|
if s, ok := config.ConnConfig.Config.RuntimeParams["pool_min_idle_conns"]; ok {
|
||||||
|
delete(connConfig.Config.RuntimeParams, "pool_min_idle_conns")
|
||||||
|
n, err := strconv.ParseInt(s, 10, 32)
|
||||||
|
if err != nil {
|
||||||
|
return nil, pgconn.NewParseConfigError(connString, "cannot parse pool_min_idle_conns", err)
|
||||||
|
}
|
||||||
|
config.MinIdleConns = int32(n)
|
||||||
|
} else {
|
||||||
|
config.MinIdleConns = defaultMinIdleConns
|
||||||
|
}
|
||||||
|
|
||||||
|
if s, ok := config.ConnConfig.Config.RuntimeParams["pool_max_conn_lifetime"]; ok {
|
||||||
|
delete(connConfig.Config.RuntimeParams, "pool_max_conn_lifetime")
|
||||||
|
d, err := time.ParseDuration(s)
|
||||||
|
if err != nil {
|
||||||
|
return nil, pgconn.NewParseConfigError(connString, "cannot parse pool_max_conn_lifetime", err)
|
||||||
|
}
|
||||||
|
config.MaxConnLifetime = d
|
||||||
|
} else {
|
||||||
|
config.MaxConnLifetime = defaultMaxConnLifetime
|
||||||
|
}
|
||||||
|
|
||||||
|
if s, ok := config.ConnConfig.Config.RuntimeParams["pool_max_conn_idle_time"]; ok {
|
||||||
|
delete(connConfig.Config.RuntimeParams, "pool_max_conn_idle_time")
|
||||||
|
d, err := time.ParseDuration(s)
|
||||||
|
if err != nil {
|
||||||
|
return nil, pgconn.NewParseConfigError(connString, "cannot parse pool_max_conn_idle_time", err)
|
||||||
|
}
|
||||||
|
config.MaxConnIdleTime = d
|
||||||
|
} else {
|
||||||
|
config.MaxConnIdleTime = defaultMaxConnIdleTime
|
||||||
|
}
|
||||||
|
|
||||||
|
if s, ok := config.ConnConfig.Config.RuntimeParams["pool_health_check_period"]; ok {
|
||||||
|
delete(connConfig.Config.RuntimeParams, "pool_health_check_period")
|
||||||
|
d, err := time.ParseDuration(s)
|
||||||
|
if err != nil {
|
||||||
|
return nil, pgconn.NewParseConfigError(connString, "cannot parse pool_health_check_period", err)
|
||||||
|
}
|
||||||
|
config.HealthCheckPeriod = d
|
||||||
|
} else {
|
||||||
|
config.HealthCheckPeriod = defaultHealthCheckPeriod
|
||||||
|
}
|
||||||
|
|
||||||
|
if s, ok := config.ConnConfig.Config.RuntimeParams["pool_max_conn_lifetime_jitter"]; ok {
|
||||||
|
delete(connConfig.Config.RuntimeParams, "pool_max_conn_lifetime_jitter")
|
||||||
|
d, err := time.ParseDuration(s)
|
||||||
|
if err != nil {
|
||||||
|
return nil, pgconn.NewParseConfigError(connString, "cannot parse pool_max_conn_lifetime_jitter", err)
|
||||||
|
}
|
||||||
|
config.MaxConnLifetimeJitter = d
|
||||||
|
}
|
||||||
|
|
||||||
|
return config, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Close closes all connections in the pool and rejects future [Pool.Acquire] calls. Blocks until all connections are returned
|
||||||
|
// to pool and closed.
|
||||||
|
func (p *Pool) Close() {
|
||||||
|
p.closeOnce.Do(func() {
|
||||||
|
close(p.closeChan)
|
||||||
|
p.p.Close()
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *Pool) isExpired(res *puddle.Resource[*connResource]) bool {
|
||||||
|
return time.Now().After(res.Value().maxAgeTime)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *Pool) triggerHealthCheck() {
|
||||||
|
const healthCheckDelay = 500 * time.Millisecond
|
||||||
|
|
||||||
|
p.healthCheckMu.Lock()
|
||||||
|
defer p.healthCheckMu.Unlock()
|
||||||
|
|
||||||
|
if p.healthCheckTimer == nil {
|
||||||
|
// Destroy is asynchronous so we give it time to actually remove itself from
|
||||||
|
// the pool otherwise we might try to check the pool size too soon
|
||||||
|
p.healthCheckTimer = time.AfterFunc(healthCheckDelay, func() {
|
||||||
|
select {
|
||||||
|
case <-p.closeChan:
|
||||||
|
case p.healthCheckChan <- struct{}{}:
|
||||||
|
default:
|
||||||
|
}
|
||||||
|
})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
p.healthCheckTimer.Reset(healthCheckDelay)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *Pool) backgroundHealthCheck() {
|
||||||
|
ticker := time.NewTicker(p.healthCheckPeriod)
|
||||||
|
defer ticker.Stop()
|
||||||
|
for {
|
||||||
|
select {
|
||||||
|
case <-p.closeChan:
|
||||||
|
return
|
||||||
|
case <-p.healthCheckChan:
|
||||||
|
p.checkHealth()
|
||||||
|
case <-ticker.C:
|
||||||
|
p.checkHealth()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *Pool) checkHealth() {
|
||||||
|
for {
|
||||||
|
// If checkMinConns failed we don't destroy any connections since we couldn't
|
||||||
|
// even get to minConns
|
||||||
|
if err := p.checkMinConns(); err != nil {
|
||||||
|
// Should we log this error somewhere?
|
||||||
|
break
|
||||||
|
}
|
||||||
|
if !p.checkConnsHealth() {
|
||||||
|
// Since we didn't destroy any connections we can stop looping
|
||||||
|
break
|
||||||
|
}
|
||||||
|
// Technically Destroy is asynchronous but 500ms should be enough for it to
|
||||||
|
// remove it from the underlying pool
|
||||||
|
select {
|
||||||
|
case <-p.closeChan:
|
||||||
|
return
|
||||||
|
case <-time.After(500 * time.Millisecond):
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// checkConnsHealth will check all idle connections, destroy a connection if
|
||||||
|
// it's idle or too old, and returns true if any were destroyed
|
||||||
|
func (p *Pool) checkConnsHealth() bool {
|
||||||
|
var destroyed bool
|
||||||
|
totalConns := p.Stat().TotalConns()
|
||||||
|
resources := p.p.AcquireAllIdle()
|
||||||
|
for _, res := range resources {
|
||||||
|
// We're okay going under minConns if the lifetime is up
|
||||||
|
if p.isExpired(res) && totalConns >= p.minConns {
|
||||||
|
atomic.AddInt64(&p.lifetimeDestroyCount, 1)
|
||||||
|
res.Destroy()
|
||||||
|
destroyed = true
|
||||||
|
// Since Destroy is async we manually decrement totalConns.
|
||||||
|
totalConns--
|
||||||
|
} else if res.IdleDuration() > p.maxConnIdleTime && totalConns > p.minConns {
|
||||||
|
atomic.AddInt64(&p.idleDestroyCount, 1)
|
||||||
|
res.Destroy()
|
||||||
|
destroyed = true
|
||||||
|
// Since Destroy is async we manually decrement totalConns.
|
||||||
|
totalConns--
|
||||||
|
} else {
|
||||||
|
res.ReleaseUnused()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return destroyed
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *Pool) checkMinConns() error {
|
||||||
|
// TotalConns can include ones that are being destroyed but we should have
|
||||||
|
// sleep(500ms) around all of the destroys to help prevent that from throwing
|
||||||
|
// off this check
|
||||||
|
|
||||||
|
// Create the number of connections needed to get to both minConns and minIdleConns
|
||||||
|
stat := p.Stat()
|
||||||
|
toCreate := max(p.minConns-stat.TotalConns(), p.minIdleConns-stat.IdleConns())
|
||||||
|
if toCreate > 0 {
|
||||||
|
return p.createIdleResources(context.Background(), int(toCreate))
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *Pool) createIdleResources(parentCtx context.Context, targetResources int) error {
|
||||||
|
ctx, cancel := context.WithCancel(parentCtx)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
errs := make(chan error, targetResources)
|
||||||
|
|
||||||
|
for range targetResources {
|
||||||
|
go func() {
|
||||||
|
err := p.p.CreateResource(ctx)
|
||||||
|
// Ignore ErrNotAvailable since it means that the pool has become full since we started creating resource.
|
||||||
|
if err == puddle.ErrNotAvailable {
|
||||||
|
err = nil
|
||||||
|
}
|
||||||
|
errs <- err
|
||||||
|
}()
|
||||||
|
}
|
||||||
|
|
||||||
|
var firstError error
|
||||||
|
for range targetResources {
|
||||||
|
err := <-errs
|
||||||
|
if err != nil && firstError == nil {
|
||||||
|
cancel()
|
||||||
|
firstError = err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return firstError
|
||||||
|
}
|
||||||
|
|
||||||
|
// Acquire returns a connection ([Conn]) from the [Pool].
|
||||||
|
func (p *Pool) Acquire(ctx context.Context) (c *Conn, err error) {
|
||||||
|
if p.acquireTracer != nil {
|
||||||
|
ctx = p.acquireTracer.TraceAcquireStart(ctx, p, TraceAcquireStartData{})
|
||||||
|
defer func() {
|
||||||
|
var conn *pgx.Conn
|
||||||
|
if c != nil {
|
||||||
|
conn = c.Conn()
|
||||||
|
}
|
||||||
|
p.acquireTracer.TraceAcquireEnd(ctx, p, TraceAcquireEndData{Conn: conn, Err: err})
|
||||||
|
}()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Try to acquire from the connection pool up to maxConns + 1 times, so that
|
||||||
|
// any that fatal errors would empty the pool and still at least try 1 fresh
|
||||||
|
// connection.
|
||||||
|
for range int(p.maxConns) + 1 {
|
||||||
|
res, err := p.p.Acquire(ctx)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
cr := res.Value()
|
||||||
|
|
||||||
|
shouldPingParams := ShouldPingParams{Conn: cr.conn, IdleDuration: res.IdleDuration()}
|
||||||
|
if p.shouldPing(ctx, shouldPingParams) {
|
||||||
|
err := func() error {
|
||||||
|
pingCtx := ctx
|
||||||
|
if p.pingTimeout > 0 {
|
||||||
|
var cancel context.CancelFunc
|
||||||
|
pingCtx, cancel = context.WithTimeout(ctx, p.pingTimeout)
|
||||||
|
defer cancel()
|
||||||
|
}
|
||||||
|
return cr.conn.Ping(pingCtx)
|
||||||
|
}()
|
||||||
|
if err != nil {
|
||||||
|
res.Destroy()
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if p.prepareConn != nil {
|
||||||
|
ok, err := p.prepareConn(ctx, cr.conn)
|
||||||
|
if !ok {
|
||||||
|
res.Destroy()
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
if ok {
|
||||||
|
res.Release()
|
||||||
|
}
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if !ok {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return cr.getConn(p, res), nil
|
||||||
|
}
|
||||||
|
return nil, errors.New("pgxpool: too many failed attempts acquiring connection; likely bug in PrepareConn, BeforeAcquire, or ShouldPing hook")
|
||||||
|
}
|
||||||
|
|
||||||
|
// AcquireFunc acquires a [Conn] and calls f with that [Conn]. ctx will only affect the [Pool.Acquire]. It has no effect on the
|
||||||
|
// call of f. The return value is either an error acquiring the [Conn] or the return value of f. The [Conn] is
|
||||||
|
// automatically released after the call of f.
|
||||||
|
func (p *Pool) AcquireFunc(ctx context.Context, f func(*Conn) error) error {
|
||||||
|
conn, err := p.Acquire(ctx)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
defer conn.Release()
|
||||||
|
|
||||||
|
return f(conn)
|
||||||
|
}
|
||||||
|
|
||||||
|
// AcquireAllIdle atomically acquires all currently idle connections. Its intended use is for health check and
|
||||||
|
// keep-alive functionality. It does not update pool statistics.
|
||||||
|
func (p *Pool) AcquireAllIdle(ctx context.Context) []*Conn {
|
||||||
|
resources := p.p.AcquireAllIdle()
|
||||||
|
conns := make([]*Conn, 0, len(resources))
|
||||||
|
for _, res := range resources {
|
||||||
|
cr := res.Value()
|
||||||
|
if p.prepareConn != nil {
|
||||||
|
ok, err := p.prepareConn(ctx, cr.conn)
|
||||||
|
if !ok || err != nil {
|
||||||
|
res.Destroy()
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
}
|
||||||
|
conns = append(conns, cr.getConn(p, res))
|
||||||
|
}
|
||||||
|
|
||||||
|
return conns
|
||||||
|
}
|
||||||
|
|
||||||
|
// Reset closes all connections, but leaves the pool open. It is intended for use when an error is detected that would
|
||||||
|
// disrupt all connections (such as a network interruption or a server state change).
|
||||||
|
//
|
||||||
|
// It is safe to reset a pool while connections are checked out. Those connections will be closed when they are returned
|
||||||
|
// to the pool.
|
||||||
|
func (p *Pool) Reset() {
|
||||||
|
p.p.Reset()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Config returns a copy of config that was used to initialize this [Pool].
|
||||||
|
func (p *Pool) Config() *Config { return p.config.Copy() }
|
||||||
|
|
||||||
|
// Stat returns a pgxpool.Stat struct with a snapshot of Pool statistics.
|
||||||
|
func (p *Pool) Stat() *Stat {
|
||||||
|
return &Stat{
|
||||||
|
s: p.p.Stat(),
|
||||||
|
newConnsCount: atomic.LoadInt64(&p.newConnsCount),
|
||||||
|
lifetimeDestroyCount: atomic.LoadInt64(&p.lifetimeDestroyCount),
|
||||||
|
idleDestroyCount: atomic.LoadInt64(&p.idleDestroyCount),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Exec acquires a connection from the [Pool] and executes the given SQL.
|
||||||
|
// SQL can be either a prepared statement name or an SQL string.
|
||||||
|
// Arguments should be referenced positionally from the SQL string as $1, $2, etc.
|
||||||
|
// The acquired connection is returned to the pool when the [Pool.Exec] function returns.
|
||||||
|
func (p *Pool) Exec(ctx context.Context, sql string, arguments ...any) (pgconn.CommandTag, error) {
|
||||||
|
c, err := p.Acquire(ctx)
|
||||||
|
if err != nil {
|
||||||
|
return pgconn.CommandTag{}, err
|
||||||
|
}
|
||||||
|
defer c.Release()
|
||||||
|
|
||||||
|
return c.Exec(ctx, sql, arguments...)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Query acquires a connection and executes a query that returns [pgx.Rows].
|
||||||
|
// Arguments should be referenced positionally from the SQL string as $1, $2, etc.
|
||||||
|
// See [pgx.Rows] documentation to close the returned [pgx.Rows] and return the acquired connection to the [Pool].
|
||||||
|
//
|
||||||
|
// If there is an error, the returned [pgx.Rows] will be returned in an error state.
|
||||||
|
// If preferred, ignore the error returned from [Pool.Query] and handle errors using the returned [pgx.Rows].
|
||||||
|
//
|
||||||
|
// For extra control over how the query is executed, the types [pgx.QueryExecMode], [pgx.QueryResultFormats], and
|
||||||
|
// [pgx.QueryResultFormatsByOID] may be used as the first args to control exactly how the query is executed. This is rarely
|
||||||
|
// needed. See the documentation for those types for details.
|
||||||
|
func (p *Pool) Query(ctx context.Context, sql string, args ...any) (pgx.Rows, error) {
|
||||||
|
c, err := p.Acquire(ctx)
|
||||||
|
if err != nil {
|
||||||
|
return errRows{err: err}, err
|
||||||
|
}
|
||||||
|
|
||||||
|
rows, err := c.Query(ctx, sql, args...)
|
||||||
|
if err != nil {
|
||||||
|
c.Release()
|
||||||
|
return errRows{err: err}, err
|
||||||
|
}
|
||||||
|
|
||||||
|
return c.getPoolRows(rows), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// QueryRow acquires a connection and executes a query that is expected
|
||||||
|
// to return at most one row ([pgx.Row]). Errors are deferred until [pgx.Row]'s
|
||||||
|
// Scan method is called. If the query selects no rows, [pgx.Row]'s Scan will
|
||||||
|
// return [pgx.ErrNoRows]. Otherwise, [pgx.Row]'s Scan scans the first selected row
|
||||||
|
// and discards the rest. The acquired connection is returned to the [Pool] when
|
||||||
|
// [pgx.Row]'s Scan method is called.
|
||||||
|
//
|
||||||
|
// Arguments should be referenced positionally from the SQL string as $1, $2, etc.
|
||||||
|
//
|
||||||
|
// For extra control over how the query is executed, the types [pgx.QueryExecMode], [pgx.QueryResultFormats], and
|
||||||
|
// [pgx.QueryResultFormatsByOID] may be used as the first args to control exactly how the query is executed. This is rarely
|
||||||
|
// needed. See the documentation for those types for details.
|
||||||
|
func (p *Pool) QueryRow(ctx context.Context, sql string, args ...any) pgx.Row {
|
||||||
|
c, err := p.Acquire(ctx)
|
||||||
|
if err != nil {
|
||||||
|
return errRow{err: err}
|
||||||
|
}
|
||||||
|
|
||||||
|
row := c.QueryRow(ctx, sql, args...)
|
||||||
|
return c.getPoolRow(row)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *Pool) SendBatch(ctx context.Context, b *pgx.Batch) pgx.BatchResults {
|
||||||
|
c, err := p.Acquire(ctx)
|
||||||
|
if err != nil {
|
||||||
|
return errBatchResults{err: err}
|
||||||
|
}
|
||||||
|
|
||||||
|
br := c.SendBatch(ctx, b)
|
||||||
|
return &poolBatchResults{br: br, c: c}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Begin acquires a connection from the [Pool] and starts a transaction. Unlike [database/sql], the context only affects the begin command. i.e. there is no
|
||||||
|
// auto-rollback on context cancellation. Begin initiates a transaction block without explicitly setting a transaction mode for the block (see [Pool.BeginTx] with [pgx.TxOptions] if transaction mode is required).
|
||||||
|
// [*Tx] is returned, which implements the [pgx.Tx] interface.
|
||||||
|
// [Tx.Commit] or [Tx.Rollback] must be called on the returned transaction to finalize the transaction block.
|
||||||
|
func (p *Pool) Begin(ctx context.Context) (pgx.Tx, error) {
|
||||||
|
return p.BeginTx(ctx, pgx.TxOptions{})
|
||||||
|
}
|
||||||
|
|
||||||
|
// BeginTx acquires a connection from the [Pool] and starts a transaction with [pgx.TxOptions] determining the transaction mode.
|
||||||
|
// Unlike [database/sql], the context only affects the begin command. i.e. there is no auto-rollback on context cancellation.
|
||||||
|
// [*Tx] is returned, which implements the [pgx.Tx] interface.
|
||||||
|
// [Tx.Commit] or [Tx.Rollback] must be called on the returned transaction to finalize the transaction block.
|
||||||
|
func (p *Pool) BeginTx(ctx context.Context, txOptions pgx.TxOptions) (pgx.Tx, error) {
|
||||||
|
c, err := p.Acquire(ctx)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
t, err := c.BeginTx(ctx, txOptions)
|
||||||
|
if err != nil {
|
||||||
|
c.Release()
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
return &Tx{t: t, c: c}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *Pool) CopyFrom(ctx context.Context, tableName pgx.Identifier, columnNames []string, rowSrc pgx.CopyFromSource) (int64, error) {
|
||||||
|
c, err := p.Acquire(ctx)
|
||||||
|
if err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
defer c.Release()
|
||||||
|
|
||||||
|
return c.Conn().CopyFrom(ctx, tableName, columnNames, rowSrc)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Ping acquires a connection from the [Pool] and executes an empty sql statement against it.
|
||||||
|
// If the sql returns without error, the database [Pool.Ping] is considered successful, otherwise, the error is returned.
|
||||||
|
func (p *Pool) Ping(ctx context.Context) error {
|
||||||
|
c, err := p.Acquire(ctx)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
defer c.Release()
|
||||||
|
return c.Ping(ctx)
|
||||||
|
}
|
||||||
+116
@@ -0,0 +1,116 @@
|
|||||||
|
package pgxpool
|
||||||
|
|
||||||
|
import (
|
||||||
|
"github.com/jackc/pgx/v5"
|
||||||
|
"github.com/jackc/pgx/v5/pgconn"
|
||||||
|
)
|
||||||
|
|
||||||
|
type errRows struct {
|
||||||
|
err error
|
||||||
|
}
|
||||||
|
|
||||||
|
func (errRows) Close() {}
|
||||||
|
func (e errRows) Err() error { return e.err }
|
||||||
|
func (errRows) CommandTag() pgconn.CommandTag { return pgconn.CommandTag{} }
|
||||||
|
func (errRows) FieldDescriptions() []pgconn.FieldDescription { return nil }
|
||||||
|
func (errRows) Next() bool { return false }
|
||||||
|
func (e errRows) Scan(dest ...any) error { return e.err }
|
||||||
|
func (e errRows) Values() ([]any, error) { return nil, e.err }
|
||||||
|
func (e errRows) RawValues() [][]byte { return nil }
|
||||||
|
func (e errRows) Conn() *pgx.Conn { return nil }
|
||||||
|
|
||||||
|
type errRow struct {
|
||||||
|
err error
|
||||||
|
}
|
||||||
|
|
||||||
|
func (e errRow) Scan(dest ...any) error { return e.err }
|
||||||
|
|
||||||
|
type poolRows struct {
|
||||||
|
r pgx.Rows
|
||||||
|
c *Conn
|
||||||
|
err error
|
||||||
|
}
|
||||||
|
|
||||||
|
func (rows *poolRows) Close() {
|
||||||
|
rows.r.Close()
|
||||||
|
if rows.c != nil {
|
||||||
|
rows.c.Release()
|
||||||
|
rows.c = nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (rows *poolRows) Err() error {
|
||||||
|
if rows.err != nil {
|
||||||
|
return rows.err
|
||||||
|
}
|
||||||
|
return rows.r.Err()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (rows *poolRows) CommandTag() pgconn.CommandTag {
|
||||||
|
return rows.r.CommandTag()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (rows *poolRows) FieldDescriptions() []pgconn.FieldDescription {
|
||||||
|
return rows.r.FieldDescriptions()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (rows *poolRows) Next() bool {
|
||||||
|
if rows.err != nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
n := rows.r.Next()
|
||||||
|
if !n {
|
||||||
|
rows.Close()
|
||||||
|
}
|
||||||
|
return n
|
||||||
|
}
|
||||||
|
|
||||||
|
func (rows *poolRows) Scan(dest ...any) error {
|
||||||
|
err := rows.r.Scan(dest...)
|
||||||
|
if err != nil {
|
||||||
|
rows.Close()
|
||||||
|
}
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
func (rows *poolRows) Values() ([]any, error) {
|
||||||
|
values, err := rows.r.Values()
|
||||||
|
if err != nil {
|
||||||
|
rows.Close()
|
||||||
|
}
|
||||||
|
return values, err
|
||||||
|
}
|
||||||
|
|
||||||
|
func (rows *poolRows) RawValues() [][]byte {
|
||||||
|
return rows.r.RawValues()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (rows *poolRows) Conn() *pgx.Conn {
|
||||||
|
return rows.r.Conn()
|
||||||
|
}
|
||||||
|
|
||||||
|
type poolRow struct {
|
||||||
|
r pgx.Row
|
||||||
|
c *Conn
|
||||||
|
err error
|
||||||
|
}
|
||||||
|
|
||||||
|
func (row *poolRow) Scan(dest ...any) error {
|
||||||
|
if row.err != nil {
|
||||||
|
return row.err
|
||||||
|
}
|
||||||
|
|
||||||
|
panicked := true
|
||||||
|
defer func() {
|
||||||
|
if panicked && row.c != nil {
|
||||||
|
row.c.Release()
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
err := row.r.Scan(dest...)
|
||||||
|
panicked = false
|
||||||
|
if row.c != nil {
|
||||||
|
row.c.Release()
|
||||||
|
}
|
||||||
|
return err
|
||||||
|
}
|
||||||
+91
@@ -0,0 +1,91 @@
|
|||||||
|
package pgxpool
|
||||||
|
|
||||||
|
import (
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/jackc/puddle/v2"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Stat is a snapshot of Pool statistics.
|
||||||
|
type Stat struct {
|
||||||
|
s *puddle.Stat
|
||||||
|
newConnsCount int64
|
||||||
|
lifetimeDestroyCount int64
|
||||||
|
idleDestroyCount int64
|
||||||
|
}
|
||||||
|
|
||||||
|
// AcquireCount returns the cumulative count of successful acquires from the pool.
|
||||||
|
func (s *Stat) AcquireCount() int64 {
|
||||||
|
return s.s.AcquireCount()
|
||||||
|
}
|
||||||
|
|
||||||
|
// AcquireDuration returns the total duration of all successful acquires from
|
||||||
|
// the pool.
|
||||||
|
func (s *Stat) AcquireDuration() time.Duration {
|
||||||
|
return s.s.AcquireDuration()
|
||||||
|
}
|
||||||
|
|
||||||
|
// AcquiredConns returns the number of currently acquired connections in the pool.
|
||||||
|
func (s *Stat) AcquiredConns() int32 {
|
||||||
|
return s.s.AcquiredResources()
|
||||||
|
}
|
||||||
|
|
||||||
|
// CanceledAcquireCount returns the cumulative count of acquires from the pool
|
||||||
|
// that were canceled by a context.
|
||||||
|
func (s *Stat) CanceledAcquireCount() int64 {
|
||||||
|
return s.s.CanceledAcquireCount()
|
||||||
|
}
|
||||||
|
|
||||||
|
// ConstructingConns returns the number of conns with construction in progress in
|
||||||
|
// the pool.
|
||||||
|
func (s *Stat) ConstructingConns() int32 {
|
||||||
|
return s.s.ConstructingResources()
|
||||||
|
}
|
||||||
|
|
||||||
|
// EmptyAcquireCount returns the cumulative count of successful acquires from the pool
|
||||||
|
// that waited for a resource to be released or constructed because the pool was
|
||||||
|
// empty.
|
||||||
|
func (s *Stat) EmptyAcquireCount() int64 {
|
||||||
|
return s.s.EmptyAcquireCount()
|
||||||
|
}
|
||||||
|
|
||||||
|
// IdleConns returns the number of currently idle conns in the pool.
|
||||||
|
func (s *Stat) IdleConns() int32 {
|
||||||
|
return s.s.IdleResources()
|
||||||
|
}
|
||||||
|
|
||||||
|
// MaxConns returns the maximum size of the pool.
|
||||||
|
func (s *Stat) MaxConns() int32 {
|
||||||
|
return s.s.MaxResources()
|
||||||
|
}
|
||||||
|
|
||||||
|
// TotalConns returns the total number of resources currently in the pool.
|
||||||
|
// The value is the sum of ConstructingConns, AcquiredConns, and
|
||||||
|
// IdleConns.
|
||||||
|
func (s *Stat) TotalConns() int32 {
|
||||||
|
return s.s.TotalResources()
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewConnsCount returns the cumulative count of new connections opened.
|
||||||
|
func (s *Stat) NewConnsCount() int64 {
|
||||||
|
return s.newConnsCount
|
||||||
|
}
|
||||||
|
|
||||||
|
// MaxLifetimeDestroyCount returns the cumulative count of connections destroyed
|
||||||
|
// because they exceeded MaxConnLifetime.
|
||||||
|
func (s *Stat) MaxLifetimeDestroyCount() int64 {
|
||||||
|
return s.lifetimeDestroyCount
|
||||||
|
}
|
||||||
|
|
||||||
|
// MaxIdleDestroyCount returns the cumulative count of connections destroyed because
|
||||||
|
// they exceeded MaxConnIdleTime.
|
||||||
|
func (s *Stat) MaxIdleDestroyCount() int64 {
|
||||||
|
return s.idleDestroyCount
|
||||||
|
}
|
||||||
|
|
||||||
|
// EmptyAcquireWaitTime returns the cumulative time waited for successful acquires
|
||||||
|
// from the pool for a resource to be released or constructed because the pool was
|
||||||
|
// empty.
|
||||||
|
func (s *Stat) EmptyAcquireWaitTime() time.Duration {
|
||||||
|
return s.s.EmptyAcquireWaitTime()
|
||||||
|
}
|
||||||
+33
@@ -0,0 +1,33 @@
|
|||||||
|
package pgxpool
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
|
||||||
|
"github.com/jackc/pgx/v5"
|
||||||
|
)
|
||||||
|
|
||||||
|
// AcquireTracer traces Acquire.
|
||||||
|
type AcquireTracer interface {
|
||||||
|
// TraceAcquireStart is called at the beginning of Acquire.
|
||||||
|
// The returned context is used for the rest of the call and will be passed to the TraceAcquireEnd.
|
||||||
|
TraceAcquireStart(ctx context.Context, pool *Pool, data TraceAcquireStartData) context.Context
|
||||||
|
// TraceAcquireEnd is called when a connection has been acquired.
|
||||||
|
TraceAcquireEnd(ctx context.Context, pool *Pool, data TraceAcquireEndData)
|
||||||
|
}
|
||||||
|
|
||||||
|
type TraceAcquireStartData struct{}
|
||||||
|
|
||||||
|
type TraceAcquireEndData struct {
|
||||||
|
Conn *pgx.Conn
|
||||||
|
Err error
|
||||||
|
}
|
||||||
|
|
||||||
|
// ReleaseTracer traces Release.
|
||||||
|
type ReleaseTracer interface {
|
||||||
|
// TraceRelease is called at the beginning of Release.
|
||||||
|
TraceRelease(pool *Pool, data TraceReleaseData)
|
||||||
|
}
|
||||||
|
|
||||||
|
type TraceReleaseData struct {
|
||||||
|
Conn *pgx.Conn
|
||||||
|
}
|
||||||
+83
@@ -0,0 +1,83 @@
|
|||||||
|
package pgxpool
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
|
||||||
|
"github.com/jackc/pgx/v5"
|
||||||
|
"github.com/jackc/pgx/v5/pgconn"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Tx represents a database transaction acquired from a Pool.
|
||||||
|
type Tx struct {
|
||||||
|
t pgx.Tx
|
||||||
|
c *Conn
|
||||||
|
}
|
||||||
|
|
||||||
|
// Begin starts a pseudo nested transaction implemented with a savepoint.
|
||||||
|
func (tx *Tx) Begin(ctx context.Context) (pgx.Tx, error) {
|
||||||
|
return tx.t.Begin(ctx)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Commit commits the transaction and returns the associated connection back to the Pool. Commit will return an error
|
||||||
|
// where errors.Is(ErrTxClosed) is true if the Tx is already closed, but is otherwise safe to call multiple times. If
|
||||||
|
// the commit fails with a rollback status (e.g. the transaction was already in a broken state) then ErrTxCommitRollback
|
||||||
|
// will be returned.
|
||||||
|
func (tx *Tx) Commit(ctx context.Context) error {
|
||||||
|
err := tx.t.Commit(ctx)
|
||||||
|
if tx.c != nil {
|
||||||
|
tx.c.Release()
|
||||||
|
tx.c = nil
|
||||||
|
}
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Rollback rolls back the transaction and returns the associated connection back to the Pool. Rollback will return
|
||||||
|
// where an error where errors.Is(ErrTxClosed) is true if the Tx is already closed, but is otherwise safe to call
|
||||||
|
// multiple times. Hence, defer tx.Rollback() is safe even if tx.Commit() will be called first in a non-error condition.
|
||||||
|
func (tx *Tx) Rollback(ctx context.Context) error {
|
||||||
|
err := tx.t.Rollback(ctx)
|
||||||
|
if tx.c != nil {
|
||||||
|
tx.c.Release()
|
||||||
|
tx.c = nil
|
||||||
|
}
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
func (tx *Tx) CopyFrom(ctx context.Context, tableName pgx.Identifier, columnNames []string, rowSrc pgx.CopyFromSource) (int64, error) {
|
||||||
|
return tx.t.CopyFrom(ctx, tableName, columnNames, rowSrc)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (tx *Tx) SendBatch(ctx context.Context, b *pgx.Batch) pgx.BatchResults {
|
||||||
|
return tx.t.SendBatch(ctx, b)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (tx *Tx) LargeObjects() pgx.LargeObjects {
|
||||||
|
return tx.t.LargeObjects()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Prepare creates a prepared statement with name and sql. If the name is empty,
|
||||||
|
// an anonymous prepared statement will be used. sql can contain placeholders
|
||||||
|
// for bound parameters. These placeholders are referenced positionally as $1, $2, etc.
|
||||||
|
//
|
||||||
|
// Prepare is idempotent; i.e. it is safe to call Prepare multiple times with the same
|
||||||
|
// name and sql arguments. This allows a code path to Prepare and Query/Exec without
|
||||||
|
// needing to first check whether the statement has already been prepared.
|
||||||
|
func (tx *Tx) Prepare(ctx context.Context, name, sql string) (*pgconn.StatementDescription, error) {
|
||||||
|
return tx.t.Prepare(ctx, name, sql)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (tx *Tx) Exec(ctx context.Context, sql string, arguments ...any) (pgconn.CommandTag, error) {
|
||||||
|
return tx.t.Exec(ctx, sql, arguments...)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (tx *Tx) Query(ctx context.Context, sql string, args ...any) (pgx.Rows, error) {
|
||||||
|
return tx.t.Query(ctx, sql, args...)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (tx *Tx) QueryRow(ctx context.Context, sql string, args ...any) pgx.Row {
|
||||||
|
return tx.t.QueryRow(ctx, sql, args...)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (tx *Tx) Conn() *pgx.Conn {
|
||||||
|
return tx.t.Conn()
|
||||||
|
}
|
||||||
+909
@@ -0,0 +1,909 @@
|
|||||||
|
// Package stdlib is the compatibility layer from pgx to database/sql.
|
||||||
|
//
|
||||||
|
// A database/sql connection can be established through sql.Open.
|
||||||
|
//
|
||||||
|
// db, err := sql.Open("pgx", "postgres://pgx_md5:secret@localhost:5432/pgx_test?sslmode=disable")
|
||||||
|
// if err != nil {
|
||||||
|
// return err
|
||||||
|
// }
|
||||||
|
//
|
||||||
|
// Or from a keyword/value string.
|
||||||
|
//
|
||||||
|
// db, err := sql.Open("pgx", "user=postgres password=secret host=localhost port=5432 database=pgx_test sslmode=disable")
|
||||||
|
// if err != nil {
|
||||||
|
// return err
|
||||||
|
// }
|
||||||
|
//
|
||||||
|
// Or from a *pgxpool.Pool.
|
||||||
|
//
|
||||||
|
// pool, err := pgxpool.New(context.Background(), os.Getenv("DATABASE_URL"))
|
||||||
|
// if err != nil {
|
||||||
|
// return err
|
||||||
|
// }
|
||||||
|
//
|
||||||
|
// db := stdlib.OpenDBFromPool(pool)
|
||||||
|
//
|
||||||
|
// Or a pgx.ConnConfig can be used to set configuration not accessible via connection string. In this case the
|
||||||
|
// pgx.ConnConfig must first be registered with the driver. This registration returns a connection string which is used
|
||||||
|
// with sql.Open.
|
||||||
|
//
|
||||||
|
// connConfig, _ := pgx.ParseConfig(os.Getenv("DATABASE_URL"))
|
||||||
|
// connConfig.Tracer = &tracelog.TraceLog{Logger: myLogger, LogLevel: tracelog.LogLevelInfo}
|
||||||
|
// connStr := stdlib.RegisterConnConfig(connConfig)
|
||||||
|
// db, _ := sql.Open("pgx", connStr)
|
||||||
|
//
|
||||||
|
// pgx uses standard PostgreSQL positional parameters in queries. e.g. $1, $2. It does not support named parameters.
|
||||||
|
//
|
||||||
|
// db.QueryRow("select * from users where id=$1", userID)
|
||||||
|
//
|
||||||
|
// (*sql.Conn) Raw() can be used to get a *pgx.Conn from the standard database/sql.DB connection pool. This allows
|
||||||
|
// operations that use pgx specific functionality.
|
||||||
|
//
|
||||||
|
// // Given db is a *sql.DB
|
||||||
|
// conn, err := db.Conn(context.Background())
|
||||||
|
// if err != nil {
|
||||||
|
// // handle error from acquiring connection from DB pool
|
||||||
|
// }
|
||||||
|
//
|
||||||
|
// err = conn.Raw(func(driverConn any) error {
|
||||||
|
// conn := driverConn.(*stdlib.Conn).Conn() // conn is a *pgx.Conn
|
||||||
|
// // Do pgx specific stuff with conn
|
||||||
|
// conn.CopyFrom(...)
|
||||||
|
// return nil
|
||||||
|
// })
|
||||||
|
// if err != nil {
|
||||||
|
// // handle error that occurred while using *pgx.Conn
|
||||||
|
// }
|
||||||
|
//
|
||||||
|
// # PostgreSQL Specific Data Types
|
||||||
|
//
|
||||||
|
// The pgtype package provides support for PostgreSQL specific types. *pgtype.Map.SQLScanner is an adapter that makes
|
||||||
|
// these types usable as a sql.Scanner.
|
||||||
|
//
|
||||||
|
// m := pgtype.NewMap()
|
||||||
|
// var a []int64
|
||||||
|
// err := db.QueryRow("select '{1,2,3}'::bigint[]").Scan(m.SQLScanner(&a))
|
||||||
|
package stdlib
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"database/sql"
|
||||||
|
"database/sql/driver"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"math"
|
||||||
|
"math/rand/v2"
|
||||||
|
"reflect"
|
||||||
|
"slices"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
"sync"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/jackc/pgx/v5"
|
||||||
|
"github.com/jackc/pgx/v5/pgconn"
|
||||||
|
"github.com/jackc/pgx/v5/pgtype"
|
||||||
|
"github.com/jackc/pgx/v5/pgxpool"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Only intrinsic types should be binary format with database/sql.
|
||||||
|
var databaseSQLResultFormats pgx.QueryResultFormatsByOID
|
||||||
|
|
||||||
|
var pgxDriver *Driver
|
||||||
|
|
||||||
|
func init() {
|
||||||
|
pgxDriver = &Driver{
|
||||||
|
configs: make(map[string]*pgx.ConnConfig),
|
||||||
|
}
|
||||||
|
|
||||||
|
// if pgx driver was already registered by different pgx major version then we
|
||||||
|
// skip registration under the default name.
|
||||||
|
if !slices.Contains(sql.Drivers(), "pgx") {
|
||||||
|
sql.Register("pgx", pgxDriver)
|
||||||
|
}
|
||||||
|
sql.Register("pgx/v5", pgxDriver)
|
||||||
|
|
||||||
|
databaseSQLResultFormats = pgx.QueryResultFormatsByOID{
|
||||||
|
pgtype.BoolOID: 1,
|
||||||
|
pgtype.ByteaOID: 1,
|
||||||
|
pgtype.CIDOID: 1,
|
||||||
|
pgtype.DateOID: 1,
|
||||||
|
pgtype.Float4OID: 1,
|
||||||
|
pgtype.Float8OID: 1,
|
||||||
|
pgtype.Int2OID: 1,
|
||||||
|
pgtype.Int4OID: 1,
|
||||||
|
pgtype.Int8OID: 1,
|
||||||
|
pgtype.OIDOID: 1,
|
||||||
|
pgtype.TimestampOID: 1,
|
||||||
|
pgtype.TimestamptzOID: 1,
|
||||||
|
pgtype.XIDOID: 1,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// OptionOpenDB options for configuring the driver when opening a new db pool.
|
||||||
|
type OptionOpenDB func(*connector)
|
||||||
|
|
||||||
|
// ShouldPingParams are passed to OptionShouldPing to decide whether to ping before reusing a connection.
|
||||||
|
type ShouldPingParams struct {
|
||||||
|
// Conn is the underlying pgx connection.
|
||||||
|
Conn *pgx.Conn
|
||||||
|
// IdleDuration is how long it has been since ResetSession last ran.
|
||||||
|
IdleDuration time.Duration
|
||||||
|
}
|
||||||
|
|
||||||
|
// OptionShouldPing controls whether stdlib should issue a liveness ping before reusing a connection.
|
||||||
|
// If the function returns true, stdlib will ping.
|
||||||
|
// If it returns false, stdlib will skip the ping.
|
||||||
|
// If not provided, default is ping only when IdleDuration > 1s.
|
||||||
|
func OptionShouldPing(f func(context.Context, ShouldPingParams) bool) OptionOpenDB {
|
||||||
|
return func(dc *connector) { dc.ShouldPing = f }
|
||||||
|
}
|
||||||
|
|
||||||
|
// OptionBeforeConnect provides a callback for before connect. It is passed a shallow copy of the ConnConfig that will
|
||||||
|
// be used to connect, so only its immediate members should be modified. Used only if db is opened with *pgx.ConnConfig.
|
||||||
|
func OptionBeforeConnect(bc func(context.Context, *pgx.ConnConfig) error) OptionOpenDB {
|
||||||
|
return func(dc *connector) {
|
||||||
|
dc.BeforeConnect = bc
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// OptionAfterConnect provides a callback for after connect. Used only if db is opened with *pgx.ConnConfig.
|
||||||
|
func OptionAfterConnect(ac func(context.Context, *pgx.Conn) error) OptionOpenDB {
|
||||||
|
return func(dc *connector) {
|
||||||
|
dc.AfterConnect = ac
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// OptionResetSession provides a callback that can be used to add custom logic prior to executing a query on the
|
||||||
|
// connection if the connection has been used before.
|
||||||
|
// If ResetSessionFunc returns ErrBadConn error the connection will be discarded.
|
||||||
|
func OptionResetSession(rs func(context.Context, *pgx.Conn) error) OptionOpenDB {
|
||||||
|
return func(dc *connector) {
|
||||||
|
dc.ResetSession = rs
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// RandomizeHostOrderFunc is a BeforeConnect hook that randomizes the host order in the provided connConfig, so that a
|
||||||
|
// new host becomes primary each time. This is useful to distribute connections for multi-master databases like
|
||||||
|
// CockroachDB. If you use this you likely should set https://golang.org/pkg/database/sql/#DB.SetConnMaxLifetime as well
|
||||||
|
// to ensure that connections are periodically rebalanced across your nodes.
|
||||||
|
func RandomizeHostOrderFunc(ctx context.Context, connConfig *pgx.ConnConfig) error {
|
||||||
|
if len(connConfig.Fallbacks) == 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
newFallbacks := append([]*pgconn.FallbackConfig{{
|
||||||
|
Host: connConfig.Host,
|
||||||
|
Port: connConfig.Port,
|
||||||
|
TLSConfig: connConfig.TLSConfig,
|
||||||
|
}}, connConfig.Fallbacks...)
|
||||||
|
|
||||||
|
rand.Shuffle(len(newFallbacks), func(i, j int) {
|
||||||
|
newFallbacks[i], newFallbacks[j] = newFallbacks[j], newFallbacks[i]
|
||||||
|
})
|
||||||
|
|
||||||
|
// Use the one that sorted last as the primary and keep the rest as the fallbacks
|
||||||
|
newPrimary := newFallbacks[len(newFallbacks)-1]
|
||||||
|
connConfig.Host = newPrimary.Host
|
||||||
|
connConfig.Port = newPrimary.Port
|
||||||
|
connConfig.TLSConfig = newPrimary.TLSConfig
|
||||||
|
connConfig.Fallbacks = newFallbacks[:len(newFallbacks)-1]
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func GetConnector(config pgx.ConnConfig, opts ...OptionOpenDB) driver.Connector {
|
||||||
|
c := connector{
|
||||||
|
ConnConfig: config,
|
||||||
|
BeforeConnect: func(context.Context, *pgx.ConnConfig) error { return nil }, // noop before connect by default
|
||||||
|
AfterConnect: func(context.Context, *pgx.Conn) error { return nil }, // noop after connect by default
|
||||||
|
ResetSession: func(context.Context, *pgx.Conn) error { return nil }, // noop reset session by default
|
||||||
|
driver: pgxDriver,
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, opt := range opts {
|
||||||
|
opt(&c)
|
||||||
|
}
|
||||||
|
return c
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetPoolConnector creates a new driver.Connector from the given *pgxpool.Pool. By using this be sure to set the
|
||||||
|
// maximum idle connections of the *sql.DB created with this connector to zero since they must be managed from the
|
||||||
|
// *pgxpool.Pool. This is required to avoid acquiring all the connections from the pgxpool and starving any direct
|
||||||
|
// users of the pgxpool.
|
||||||
|
func GetPoolConnector(pool *pgxpool.Pool, opts ...OptionOpenDB) driver.Connector {
|
||||||
|
c := connector{
|
||||||
|
pool: pool,
|
||||||
|
ResetSession: func(context.Context, *pgx.Conn) error { return nil }, // noop reset session by default
|
||||||
|
driver: pgxDriver,
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, opt := range opts {
|
||||||
|
opt(&c)
|
||||||
|
}
|
||||||
|
|
||||||
|
return c
|
||||||
|
}
|
||||||
|
|
||||||
|
func OpenDB(config pgx.ConnConfig, opts ...OptionOpenDB) *sql.DB {
|
||||||
|
c := GetConnector(config, opts...)
|
||||||
|
return sql.OpenDB(c)
|
||||||
|
}
|
||||||
|
|
||||||
|
// OpenDBFromPool creates a new *sql.DB from the given *pgxpool.Pool. Note that this method automatically sets the
|
||||||
|
// maximum number of idle connections in *sql.DB to zero, since they must be managed from the *pgxpool.Pool. This is
|
||||||
|
// required to avoid acquiring all the connections from the pgxpool and starving any direct users of the pgxpool. Note
|
||||||
|
// that closing the returned *sql.DB will not close the *pgxpool.Pool.
|
||||||
|
func OpenDBFromPool(pool *pgxpool.Pool, opts ...OptionOpenDB) *sql.DB {
|
||||||
|
c := GetPoolConnector(pool, opts...)
|
||||||
|
db := sql.OpenDB(c)
|
||||||
|
db.SetMaxIdleConns(0)
|
||||||
|
return db
|
||||||
|
}
|
||||||
|
|
||||||
|
type connector struct {
|
||||||
|
pgx.ConnConfig
|
||||||
|
pool *pgxpool.Pool
|
||||||
|
BeforeConnect func(context.Context, *pgx.ConnConfig) error // function to call before creation of every new connection
|
||||||
|
AfterConnect func(context.Context, *pgx.Conn) error // function to call after creation of every new connection
|
||||||
|
ResetSession func(context.Context, *pgx.Conn) error // function is called before a connection is reused
|
||||||
|
ShouldPing func(context.Context, ShouldPingParams) bool // function to decide if stdlib should ping before reusing a connection
|
||||||
|
driver *Driver
|
||||||
|
}
|
||||||
|
|
||||||
|
// Connect implement driver.Connector interface
|
||||||
|
func (c connector) Connect(ctx context.Context) (driver.Conn, error) {
|
||||||
|
var (
|
||||||
|
connConfig pgx.ConnConfig
|
||||||
|
conn *pgx.Conn
|
||||||
|
close func(context.Context) error
|
||||||
|
err error
|
||||||
|
)
|
||||||
|
|
||||||
|
if c.pool == nil {
|
||||||
|
// Create a shallow copy of the config, so that BeforeConnect can safely modify it
|
||||||
|
connConfig = c.ConnConfig
|
||||||
|
|
||||||
|
if err = c.BeforeConnect(ctx, &connConfig); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
if conn, err = pgx.ConnectConfig(ctx, &connConfig); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
if err = c.AfterConnect(ctx, conn); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
close = conn.Close
|
||||||
|
} else {
|
||||||
|
var pconn *pgxpool.Conn
|
||||||
|
|
||||||
|
pconn, err = c.pool.Acquire(ctx)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
conn = pconn.Conn()
|
||||||
|
|
||||||
|
close = func(_ context.Context) error {
|
||||||
|
pconn.Release()
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return &Conn{
|
||||||
|
conn: conn,
|
||||||
|
close: close,
|
||||||
|
driver: c.driver,
|
||||||
|
connConfig: connConfig,
|
||||||
|
resetSessionFunc: c.ResetSession,
|
||||||
|
shouldPing: c.ShouldPing,
|
||||||
|
psRefCounts: make(map[*pgconn.StatementDescription]int),
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Driver implement driver.Connector interface
|
||||||
|
func (c connector) Driver() driver.Driver {
|
||||||
|
return c.driver
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetDefaultDriver returns the driver initialized in the init function
|
||||||
|
// and used when the pgx driver is registered.
|
||||||
|
func GetDefaultDriver() driver.Driver {
|
||||||
|
return pgxDriver
|
||||||
|
}
|
||||||
|
|
||||||
|
type Driver struct {
|
||||||
|
configMutex sync.Mutex
|
||||||
|
configs map[string]*pgx.ConnConfig
|
||||||
|
sequence int
|
||||||
|
}
|
||||||
|
|
||||||
|
func (d *Driver) Open(name string) (driver.Conn, error) {
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), 60*time.Second) // Ensure eventual timeout
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
connector, err := d.OpenConnector(name)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return connector.Connect(ctx)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (d *Driver) OpenConnector(name string) (driver.Connector, error) {
|
||||||
|
return &driverConnector{driver: d, name: name}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (d *Driver) registerConnConfig(c *pgx.ConnConfig) string {
|
||||||
|
d.configMutex.Lock()
|
||||||
|
connStr := fmt.Sprintf("registeredConnConfig%d", d.sequence)
|
||||||
|
d.sequence++
|
||||||
|
d.configs[connStr] = c
|
||||||
|
d.configMutex.Unlock()
|
||||||
|
return connStr
|
||||||
|
}
|
||||||
|
|
||||||
|
func (d *Driver) unregisterConnConfig(connStr string) {
|
||||||
|
d.configMutex.Lock()
|
||||||
|
delete(d.configs, connStr)
|
||||||
|
d.configMutex.Unlock()
|
||||||
|
}
|
||||||
|
|
||||||
|
type driverConnector struct {
|
||||||
|
driver *Driver
|
||||||
|
name string
|
||||||
|
}
|
||||||
|
|
||||||
|
func (dc *driverConnector) Connect(ctx context.Context) (driver.Conn, error) {
|
||||||
|
var connConfig *pgx.ConnConfig
|
||||||
|
|
||||||
|
dc.driver.configMutex.Lock()
|
||||||
|
connConfig = dc.driver.configs[dc.name]
|
||||||
|
dc.driver.configMutex.Unlock()
|
||||||
|
|
||||||
|
if connConfig == nil {
|
||||||
|
var err error
|
||||||
|
connConfig, err = pgx.ParseConfig(dc.name)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
conn, err := pgx.ConnectConfig(ctx, connConfig)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
c := &Conn{
|
||||||
|
conn: conn,
|
||||||
|
close: conn.Close,
|
||||||
|
driver: dc.driver,
|
||||||
|
connConfig: *connConfig,
|
||||||
|
resetSessionFunc: func(context.Context, *pgx.Conn) error { return nil },
|
||||||
|
psRefCounts: make(map[*pgconn.StatementDescription]int),
|
||||||
|
}
|
||||||
|
|
||||||
|
return c, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (dc *driverConnector) Driver() driver.Driver {
|
||||||
|
return dc.driver
|
||||||
|
}
|
||||||
|
|
||||||
|
// RegisterConnConfig registers a ConnConfig and returns the connection string to use with Open.
|
||||||
|
func RegisterConnConfig(c *pgx.ConnConfig) string {
|
||||||
|
return pgxDriver.registerConnConfig(c)
|
||||||
|
}
|
||||||
|
|
||||||
|
// UnregisterConnConfig removes the ConnConfig registration for connStr.
|
||||||
|
func UnregisterConnConfig(connStr string) {
|
||||||
|
pgxDriver.unregisterConnConfig(connStr)
|
||||||
|
}
|
||||||
|
|
||||||
|
type Conn struct {
|
||||||
|
conn *pgx.Conn
|
||||||
|
close func(context.Context) error
|
||||||
|
driver *Driver
|
||||||
|
connConfig pgx.ConnConfig
|
||||||
|
resetSessionFunc func(context.Context, *pgx.Conn) error // Function is called before a connection is reused
|
||||||
|
shouldPing func(context.Context, ShouldPingParams) bool // Function to decide if stdlib should ping before reusing a connection
|
||||||
|
lastResetSessionTime time.Time
|
||||||
|
|
||||||
|
// psRefCounts contains reference counts for prepared statements. Prepare uses the underlying pgx logic to generate
|
||||||
|
// deterministic statement names from the statement text. If this query has already been prepared then the existing
|
||||||
|
// *pgconn.StatementDescription will be returned. However, this means that if Close is called on the returned Stmt
|
||||||
|
// then the underlying prepared statement will be closed even when the underlying prepared statement is still in use
|
||||||
|
// by another database/sql Stmt. To prevent this psRefCounts keeps track of how many database/sql statements are using
|
||||||
|
// the same underlying statement and only closes the underlying statement when the reference count reaches 0.
|
||||||
|
psRefCounts map[*pgconn.StatementDescription]int
|
||||||
|
}
|
||||||
|
|
||||||
|
// Conn returns the underlying *pgx.Conn
|
||||||
|
func (c *Conn) Conn() *pgx.Conn {
|
||||||
|
return c.conn
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *Conn) Prepare(query string) (driver.Stmt, error) {
|
||||||
|
return c.PrepareContext(context.Background(), query)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *Conn) PrepareContext(ctx context.Context, query string) (driver.Stmt, error) {
|
||||||
|
if c.conn.IsClosed() {
|
||||||
|
return nil, driver.ErrBadConn
|
||||||
|
}
|
||||||
|
|
||||||
|
sd, err := c.conn.Prepare(ctx, query, query)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
c.psRefCounts[sd]++
|
||||||
|
|
||||||
|
return &Stmt{sd: sd, conn: c}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *Conn) Close() error {
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), time.Second*5)
|
||||||
|
defer cancel()
|
||||||
|
return c.close(ctx)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *Conn) Begin() (driver.Tx, error) {
|
||||||
|
return c.BeginTx(context.Background(), driver.TxOptions{})
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *Conn) BeginTx(ctx context.Context, opts driver.TxOptions) (driver.Tx, error) {
|
||||||
|
if c.conn.IsClosed() {
|
||||||
|
return nil, driver.ErrBadConn
|
||||||
|
}
|
||||||
|
|
||||||
|
var pgxOpts pgx.TxOptions
|
||||||
|
switch sql.IsolationLevel(opts.Isolation) {
|
||||||
|
case sql.LevelDefault:
|
||||||
|
case sql.LevelReadUncommitted:
|
||||||
|
pgxOpts.IsoLevel = pgx.ReadUncommitted
|
||||||
|
case sql.LevelReadCommitted:
|
||||||
|
pgxOpts.IsoLevel = pgx.ReadCommitted
|
||||||
|
case sql.LevelRepeatableRead, sql.LevelSnapshot:
|
||||||
|
pgxOpts.IsoLevel = pgx.RepeatableRead
|
||||||
|
case sql.LevelSerializable:
|
||||||
|
pgxOpts.IsoLevel = pgx.Serializable
|
||||||
|
default:
|
||||||
|
return nil, fmt.Errorf("unsupported isolation: %v", opts.Isolation)
|
||||||
|
}
|
||||||
|
|
||||||
|
if opts.ReadOnly {
|
||||||
|
pgxOpts.AccessMode = pgx.ReadOnly
|
||||||
|
}
|
||||||
|
|
||||||
|
tx, err := c.conn.BeginTx(ctx, pgxOpts)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
return wrapTx{ctx: ctx, tx: tx}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *Conn) ExecContext(ctx context.Context, query string, argsV []driver.NamedValue) (driver.Result, error) {
|
||||||
|
if c.conn.IsClosed() {
|
||||||
|
return nil, driver.ErrBadConn
|
||||||
|
}
|
||||||
|
|
||||||
|
args := make([]any, len(argsV))
|
||||||
|
convertNamedArguments(args, argsV)
|
||||||
|
|
||||||
|
commandTag, err := c.conn.Exec(ctx, query, args...)
|
||||||
|
// if we got a network error before we had a chance to send the query, retry
|
||||||
|
if err != nil {
|
||||||
|
if pgconn.SafeToRetry(err) {
|
||||||
|
return nil, driver.ErrBadConn
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return driver.RowsAffected(commandTag.RowsAffected()), err
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *Conn) QueryContext(ctx context.Context, query string, argsV []driver.NamedValue) (driver.Rows, error) {
|
||||||
|
if c.conn.IsClosed() {
|
||||||
|
return nil, driver.ErrBadConn
|
||||||
|
}
|
||||||
|
|
||||||
|
args := make([]any, 1+len(argsV))
|
||||||
|
args[0] = databaseSQLResultFormats
|
||||||
|
convertNamedArguments(args[1:], argsV)
|
||||||
|
|
||||||
|
rows, err := c.conn.Query(ctx, query, args...)
|
||||||
|
if err != nil {
|
||||||
|
if pgconn.SafeToRetry(err) {
|
||||||
|
return nil, driver.ErrBadConn
|
||||||
|
}
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Preload first row because otherwise we won't know what columns are available when database/sql asks.
|
||||||
|
more := rows.Next()
|
||||||
|
if err = rows.Err(); err != nil {
|
||||||
|
rows.Close()
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return &Rows{conn: c, rows: rows, skipNext: true, skipNextMore: more}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *Conn) Ping(ctx context.Context) error {
|
||||||
|
if c.conn.IsClosed() {
|
||||||
|
return driver.ErrBadConn
|
||||||
|
}
|
||||||
|
|
||||||
|
err := c.conn.Ping(ctx)
|
||||||
|
if err != nil {
|
||||||
|
// A Ping failure implies some sort of fatal state. The connection is almost certainly already closed by the
|
||||||
|
// failure, but manually close it just to be sure.
|
||||||
|
c.Close()
|
||||||
|
return driver.ErrBadConn
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *Conn) CheckNamedValue(*driver.NamedValue) error {
|
||||||
|
// Underlying pgx supports sql.Scanner and driver.Valuer interfaces natively. So everything can be passed through directly.
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *Conn) ResetSession(ctx context.Context) error {
|
||||||
|
if c.conn.IsClosed() {
|
||||||
|
return driver.ErrBadConn
|
||||||
|
}
|
||||||
|
|
||||||
|
// Discard connection if it has an open transaction. This can happen if the
|
||||||
|
// application did not properly commit or rollback a transaction.
|
||||||
|
if c.conn.PgConn().TxStatus() != 'I' {
|
||||||
|
return driver.ErrBadConn
|
||||||
|
}
|
||||||
|
|
||||||
|
now := time.Now()
|
||||||
|
idle := now.Sub(c.lastResetSessionTime)
|
||||||
|
|
||||||
|
doPing := idle > time.Second // default behavior: ping only if idle > 1s
|
||||||
|
|
||||||
|
if c.shouldPing != nil {
|
||||||
|
doPing = c.shouldPing(ctx, ShouldPingParams{
|
||||||
|
Conn: c.conn,
|
||||||
|
IdleDuration: idle,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
if doPing {
|
||||||
|
if err := c.conn.PgConn().Ping(ctx); err != nil {
|
||||||
|
return driver.ErrBadConn
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
c.lastResetSessionTime = now
|
||||||
|
|
||||||
|
return c.resetSessionFunc(ctx, c.conn)
|
||||||
|
}
|
||||||
|
|
||||||
|
type Stmt struct {
|
||||||
|
sd *pgconn.StatementDescription
|
||||||
|
conn *Conn
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Stmt) Close() error {
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), time.Second*5)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
refCount := s.conn.psRefCounts[s.sd]
|
||||||
|
if refCount == 1 {
|
||||||
|
delete(s.conn.psRefCounts, s.sd)
|
||||||
|
} else {
|
||||||
|
s.conn.psRefCounts[s.sd]--
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
return s.conn.conn.Deallocate(ctx, s.sd.SQL)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Stmt) NumInput() int {
|
||||||
|
return len(s.sd.ParamOIDs)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Stmt) Exec(argsV []driver.Value) (driver.Result, error) {
|
||||||
|
return nil, errors.New("Stmt.Exec deprecated and not implemented")
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Stmt) ExecContext(ctx context.Context, argsV []driver.NamedValue) (driver.Result, error) {
|
||||||
|
return s.conn.ExecContext(ctx, s.sd.SQL, argsV)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Stmt) Query(argsV []driver.Value) (driver.Rows, error) {
|
||||||
|
return nil, errors.New("Stmt.Query deprecated and not implemented")
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Stmt) QueryContext(ctx context.Context, argsV []driver.NamedValue) (driver.Rows, error) {
|
||||||
|
return s.conn.QueryContext(ctx, s.sd.SQL, argsV)
|
||||||
|
}
|
||||||
|
|
||||||
|
type rowValueFunc func(src []byte) (driver.Value, error)
|
||||||
|
|
||||||
|
type Rows struct {
|
||||||
|
conn *Conn
|
||||||
|
rows pgx.Rows
|
||||||
|
valueFuncs []rowValueFunc
|
||||||
|
skipNext bool
|
||||||
|
skipNextMore bool
|
||||||
|
|
||||||
|
columnNames []string
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *Rows) Columns() []string {
|
||||||
|
if r.columnNames == nil {
|
||||||
|
fields := r.rows.FieldDescriptions()
|
||||||
|
r.columnNames = make([]string, len(fields))
|
||||||
|
for i, fd := range fields {
|
||||||
|
r.columnNames[i] = string(fd.Name)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return r.columnNames
|
||||||
|
}
|
||||||
|
|
||||||
|
// ColumnTypeDatabaseTypeName returns the database system type name. If the name is unknown the OID is returned.
|
||||||
|
func (r *Rows) ColumnTypeDatabaseTypeName(index int) string {
|
||||||
|
if dt, ok := r.conn.conn.TypeMap().TypeForOID(r.rows.FieldDescriptions()[index].DataTypeOID); ok {
|
||||||
|
return strings.ToUpper(dt.Name)
|
||||||
|
}
|
||||||
|
|
||||||
|
return strconv.FormatInt(int64(r.rows.FieldDescriptions()[index].DataTypeOID), 10)
|
||||||
|
}
|
||||||
|
|
||||||
|
const varHeaderSize = 4
|
||||||
|
|
||||||
|
// ColumnTypeLength returns the length of the column type if the column is a
|
||||||
|
// variable length type. If the column is not a variable length type ok
|
||||||
|
// should return false.
|
||||||
|
func (r *Rows) ColumnTypeLength(index int) (int64, bool) {
|
||||||
|
fd := r.rows.FieldDescriptions()[index]
|
||||||
|
|
||||||
|
switch fd.DataTypeOID {
|
||||||
|
case pgtype.TextOID, pgtype.ByteaOID:
|
||||||
|
return math.MaxInt64, true
|
||||||
|
case pgtype.VarcharOID, pgtype.BPCharOID:
|
||||||
|
return int64(fd.TypeModifier - varHeaderSize), true
|
||||||
|
case pgtype.VarbitOID:
|
||||||
|
return int64(fd.TypeModifier), true
|
||||||
|
default:
|
||||||
|
return 0, false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ColumnTypePrecisionScale should return the precision and scale for decimal
|
||||||
|
// types. If not applicable, ok should be false.
|
||||||
|
func (r *Rows) ColumnTypePrecisionScale(index int) (precision, scale int64, ok bool) {
|
||||||
|
fd := r.rows.FieldDescriptions()[index]
|
||||||
|
|
||||||
|
switch fd.DataTypeOID {
|
||||||
|
case pgtype.NumericOID:
|
||||||
|
mod := fd.TypeModifier - varHeaderSize
|
||||||
|
precision = int64((mod >> 16) & 0xffff)
|
||||||
|
scale = int64(mod & 0xffff)
|
||||||
|
return precision, scale, true
|
||||||
|
default:
|
||||||
|
return 0, 0, false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ColumnTypeScanType returns the value type that can be used to scan types into.
|
||||||
|
func (r *Rows) ColumnTypeScanType(index int) reflect.Type {
|
||||||
|
fd := r.rows.FieldDescriptions()[index]
|
||||||
|
|
||||||
|
switch fd.DataTypeOID {
|
||||||
|
case pgtype.Float8OID:
|
||||||
|
return reflect.TypeFor[float64]()
|
||||||
|
case pgtype.Float4OID:
|
||||||
|
return reflect.TypeFor[float32]()
|
||||||
|
case pgtype.Int8OID:
|
||||||
|
return reflect.TypeFor[int64]()
|
||||||
|
case pgtype.Int4OID:
|
||||||
|
return reflect.TypeFor[int32]()
|
||||||
|
case pgtype.Int2OID:
|
||||||
|
return reflect.TypeFor[int16]()
|
||||||
|
case pgtype.BoolOID:
|
||||||
|
return reflect.TypeFor[bool]()
|
||||||
|
case pgtype.NumericOID:
|
||||||
|
return reflect.TypeFor[float64]()
|
||||||
|
case pgtype.DateOID, pgtype.TimestampOID, pgtype.TimestamptzOID:
|
||||||
|
return reflect.TypeFor[time.Time]()
|
||||||
|
case pgtype.ByteaOID:
|
||||||
|
return reflect.TypeFor[[]byte]()
|
||||||
|
default:
|
||||||
|
return reflect.TypeFor[string]()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *Rows) Close() error {
|
||||||
|
r.rows.Close()
|
||||||
|
return r.rows.Err()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *Rows) Next(dest []driver.Value) error {
|
||||||
|
m := r.conn.conn.TypeMap()
|
||||||
|
fieldDescriptions := r.rows.FieldDescriptions()
|
||||||
|
|
||||||
|
if r.valueFuncs == nil {
|
||||||
|
r.valueFuncs = make([]rowValueFunc, len(fieldDescriptions))
|
||||||
|
|
||||||
|
for i, fd := range fieldDescriptions {
|
||||||
|
dataTypeOID := fd.DataTypeOID
|
||||||
|
format := fd.Format
|
||||||
|
|
||||||
|
switch fd.DataTypeOID {
|
||||||
|
case pgtype.BoolOID:
|
||||||
|
var d bool
|
||||||
|
scanPlan := m.PlanScan(dataTypeOID, format, &d)
|
||||||
|
r.valueFuncs[i] = func(src []byte) (driver.Value, error) {
|
||||||
|
err := scanPlan.Scan(src, &d)
|
||||||
|
return d, err
|
||||||
|
}
|
||||||
|
case pgtype.ByteaOID:
|
||||||
|
var d []byte
|
||||||
|
scanPlan := m.PlanScan(dataTypeOID, format, &d)
|
||||||
|
r.valueFuncs[i] = func(src []byte) (driver.Value, error) {
|
||||||
|
err := scanPlan.Scan(src, &d)
|
||||||
|
return d, err
|
||||||
|
}
|
||||||
|
case pgtype.CIDOID, pgtype.OIDOID, pgtype.XIDOID:
|
||||||
|
var d pgtype.Uint32
|
||||||
|
scanPlan := m.PlanScan(dataTypeOID, format, &d)
|
||||||
|
r.valueFuncs[i] = func(src []byte) (driver.Value, error) {
|
||||||
|
err := scanPlan.Scan(src, &d)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return d.Value()
|
||||||
|
}
|
||||||
|
case pgtype.DateOID:
|
||||||
|
var d pgtype.Date
|
||||||
|
scanPlan := m.PlanScan(dataTypeOID, format, &d)
|
||||||
|
r.valueFuncs[i] = func(src []byte) (driver.Value, error) {
|
||||||
|
err := scanPlan.Scan(src, &d)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return d.Value()
|
||||||
|
}
|
||||||
|
case pgtype.Float4OID:
|
||||||
|
var d float32
|
||||||
|
scanPlan := m.PlanScan(dataTypeOID, format, &d)
|
||||||
|
r.valueFuncs[i] = func(src []byte) (driver.Value, error) {
|
||||||
|
err := scanPlan.Scan(src, &d)
|
||||||
|
return float64(d), err
|
||||||
|
}
|
||||||
|
case pgtype.Float8OID:
|
||||||
|
var d float64
|
||||||
|
scanPlan := m.PlanScan(dataTypeOID, format, &d)
|
||||||
|
r.valueFuncs[i] = func(src []byte) (driver.Value, error) {
|
||||||
|
err := scanPlan.Scan(src, &d)
|
||||||
|
return d, err
|
||||||
|
}
|
||||||
|
case pgtype.Int2OID:
|
||||||
|
var d int16
|
||||||
|
scanPlan := m.PlanScan(dataTypeOID, format, &d)
|
||||||
|
r.valueFuncs[i] = func(src []byte) (driver.Value, error) {
|
||||||
|
err := scanPlan.Scan(src, &d)
|
||||||
|
return int64(d), err
|
||||||
|
}
|
||||||
|
case pgtype.Int4OID:
|
||||||
|
var d int32
|
||||||
|
scanPlan := m.PlanScan(dataTypeOID, format, &d)
|
||||||
|
r.valueFuncs[i] = func(src []byte) (driver.Value, error) {
|
||||||
|
err := scanPlan.Scan(src, &d)
|
||||||
|
return int64(d), err
|
||||||
|
}
|
||||||
|
case pgtype.Int8OID:
|
||||||
|
var d int64
|
||||||
|
scanPlan := m.PlanScan(dataTypeOID, format, &d)
|
||||||
|
r.valueFuncs[i] = func(src []byte) (driver.Value, error) {
|
||||||
|
err := scanPlan.Scan(src, &d)
|
||||||
|
return d, err
|
||||||
|
}
|
||||||
|
case pgtype.JSONOID, pgtype.JSONBOID:
|
||||||
|
var d []byte
|
||||||
|
scanPlan := m.PlanScan(dataTypeOID, format, &d)
|
||||||
|
r.valueFuncs[i] = func(src []byte) (driver.Value, error) {
|
||||||
|
err := scanPlan.Scan(src, &d)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return d, nil
|
||||||
|
}
|
||||||
|
case pgtype.TimestampOID:
|
||||||
|
var d pgtype.Timestamp
|
||||||
|
scanPlan := m.PlanScan(dataTypeOID, format, &d)
|
||||||
|
r.valueFuncs[i] = func(src []byte) (driver.Value, error) {
|
||||||
|
err := scanPlan.Scan(src, &d)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return d.Value()
|
||||||
|
}
|
||||||
|
case pgtype.TimestamptzOID:
|
||||||
|
var d pgtype.Timestamptz
|
||||||
|
scanPlan := m.PlanScan(dataTypeOID, format, &d)
|
||||||
|
r.valueFuncs[i] = func(src []byte) (driver.Value, error) {
|
||||||
|
err := scanPlan.Scan(src, &d)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return d.Value()
|
||||||
|
}
|
||||||
|
case pgtype.XMLOID:
|
||||||
|
var d []byte
|
||||||
|
scanPlan := m.PlanScan(dataTypeOID, format, &d)
|
||||||
|
r.valueFuncs[i] = func(src []byte) (driver.Value, error) {
|
||||||
|
err := scanPlan.Scan(src, &d)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return d, nil
|
||||||
|
}
|
||||||
|
default:
|
||||||
|
var d string
|
||||||
|
scanPlan := m.PlanScan(dataTypeOID, format, &d)
|
||||||
|
r.valueFuncs[i] = func(src []byte) (driver.Value, error) {
|
||||||
|
err := scanPlan.Scan(src, &d)
|
||||||
|
return d, err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
var more bool
|
||||||
|
if r.skipNext {
|
||||||
|
more = r.skipNextMore
|
||||||
|
r.skipNext = false
|
||||||
|
} else {
|
||||||
|
more = r.rows.Next()
|
||||||
|
}
|
||||||
|
|
||||||
|
if !more {
|
||||||
|
if r.rows.Err() == nil {
|
||||||
|
return io.EOF
|
||||||
|
} else {
|
||||||
|
return r.rows.Err()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
for i, rv := range r.rows.RawValues() {
|
||||||
|
if rv != nil {
|
||||||
|
var err error
|
||||||
|
dest[i], err = r.valueFuncs[i](rv)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("convert field %d failed: %w", i, err)
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
dest[i] = nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func convertNamedArguments(args []any, argsV []driver.NamedValue) {
|
||||||
|
for i, v := range argsV {
|
||||||
|
if v.Value != nil {
|
||||||
|
args[i] = v.Value.(any)
|
||||||
|
} else {
|
||||||
|
args[i] = nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
type wrapTx struct {
|
||||||
|
ctx context.Context
|
||||||
|
tx pgx.Tx
|
||||||
|
}
|
||||||
|
|
||||||
|
func (wtx wrapTx) Commit() error { return wtx.tx.Commit(wtx.ctx) }
|
||||||
|
|
||||||
|
func (wtx wrapTx) Rollback() error { return wtx.tx.Rollback(wtx.ctx) }
|
||||||
+79
@@ -0,0 +1,79 @@
|
|||||||
|
# 2.2.2 (September 10, 2024)
|
||||||
|
|
||||||
|
* Add empty acquire time to stats (Maxim Ivanov)
|
||||||
|
* Stop importing nanotime from runtime via linkname (maypok86)
|
||||||
|
|
||||||
|
# 2.2.1 (July 15, 2023)
|
||||||
|
|
||||||
|
* Fix: CreateResource cannot overflow pool. This changes documented behavior of CreateResource. Previously,
|
||||||
|
CreateResource could create a resource even if the pool was full. This could cause the pool to overflow. While this
|
||||||
|
was documented, it was documenting incorrect behavior. CreateResource now returns an error if the pool is full.
|
||||||
|
|
||||||
|
# 2.2.0 (February 11, 2023)
|
||||||
|
|
||||||
|
* Use Go 1.19 atomics and drop go.uber.org/atomic dependency
|
||||||
|
|
||||||
|
# 2.1.2 (November 12, 2022)
|
||||||
|
|
||||||
|
* Restore support to Go 1.18 via go.uber.org/atomic
|
||||||
|
|
||||||
|
# 2.1.1 (November 11, 2022)
|
||||||
|
|
||||||
|
* Fix create resource concurrently with Stat call race
|
||||||
|
|
||||||
|
# 2.1.0 (October 28, 2022)
|
||||||
|
|
||||||
|
* Concurrency control is now implemented with a semaphore. This simplifies some internal logic, resolves a few error conditions (including a deadlock), and improves performance. (Jan Dubsky)
|
||||||
|
* Go 1.19 is now required for the improved atomic support.
|
||||||
|
|
||||||
|
# 2.0.1 (October 28, 2022)
|
||||||
|
|
||||||
|
* Fix race condition when Close is called concurrently with multiple constructors
|
||||||
|
|
||||||
|
# 2.0.0 (September 17, 2022)
|
||||||
|
|
||||||
|
* Use generics instead of interface{} (Столяров Владимир Алексеевич)
|
||||||
|
* Add Reset
|
||||||
|
* Do not cancel resource construction when Acquire is canceled
|
||||||
|
* NewPool takes Config
|
||||||
|
|
||||||
|
# 1.3.0 (August 27, 2022)
|
||||||
|
|
||||||
|
* Acquire creates resources in background to allow creation to continue after Acquire is canceled (James Hartig)
|
||||||
|
|
||||||
|
# 1.2.1 (December 2, 2021)
|
||||||
|
|
||||||
|
* TryAcquire now does not block when background constructing resource
|
||||||
|
|
||||||
|
# 1.2.0 (November 20, 2021)
|
||||||
|
|
||||||
|
* Add TryAcquire (A. Jensen)
|
||||||
|
* Fix: remove memory leak / unintentionally pinned memory when shrinking slices (Alexander Staubo)
|
||||||
|
* Fix: Do not leave pool locked after panic from nil context
|
||||||
|
|
||||||
|
# 1.1.4 (September 11, 2021)
|
||||||
|
|
||||||
|
* Fix: Deadlock in CreateResource if pool was closed during resource acquisition (Dmitriy Matrenichev)
|
||||||
|
|
||||||
|
# 1.1.3 (December 3, 2020)
|
||||||
|
|
||||||
|
* Fix: Failed resource creation could cause concurrent Acquire to hang. (Evgeny Vanslov)
|
||||||
|
|
||||||
|
# 1.1.2 (September 26, 2020)
|
||||||
|
|
||||||
|
* Fix: Resource.Destroy no longer removes itself from the pool before its destructor has completed.
|
||||||
|
* Fix: Prevent crash when pool is closed while resource is being created.
|
||||||
|
|
||||||
|
# 1.1.1 (April 2, 2020)
|
||||||
|
|
||||||
|
* Pool.Close can be safely called multiple times
|
||||||
|
* AcquireAllIDle immediately returns nil if pool is closed
|
||||||
|
* CreateResource checks if pool is closed before taking any action
|
||||||
|
* Fix potential race condition when CreateResource and Close are called concurrently. CreateResource now checks if pool is closed before adding newly created resource to pool.
|
||||||
|
|
||||||
|
# 1.1.0 (February 5, 2020)
|
||||||
|
|
||||||
|
* Use runtime.nanotime for faster tracking of acquire time and last usage time.
|
||||||
|
* Track resource idle time to enable client health check logic. (Patrick Ellul)
|
||||||
|
* Add CreateResource to construct a new resource without acquiring it. (Patrick Ellul)
|
||||||
|
* Fix deadlock race when acquire is cancelled. (Michael Tharp)
|
||||||
+22
@@ -0,0 +1,22 @@
|
|||||||
|
Copyright (c) 2018 Jack Christensen
|
||||||
|
|
||||||
|
MIT License
|
||||||
|
|
||||||
|
Permission is hereby granted, free of charge, to any person obtaining
|
||||||
|
a copy of this software and associated documentation files (the
|
||||||
|
"Software"), to deal in the Software without restriction, including
|
||||||
|
without limitation the rights to use, copy, modify, merge, publish,
|
||||||
|
distribute, sublicense, and/or sell copies of the Software, and to
|
||||||
|
permit persons to whom the Software is furnished to do so, subject to
|
||||||
|
the following conditions:
|
||||||
|
|
||||||
|
The above copyright notice and this permission notice shall be
|
||||||
|
included in all copies or substantial portions of the Software.
|
||||||
|
|
||||||
|
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND,
|
||||||
|
EXPRESS OR IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF
|
||||||
|
MERCHANTABILITY, FITNESS FOR A PARTICULAR PURPOSE AND
|
||||||
|
NONINFRINGEMENT. IN NO EVENT SHALL THE AUTHORS OR COPYRIGHT HOLDERS BE
|
||||||
|
LIABLE FOR ANY CLAIM, DAMAGES OR OTHER LIABILITY, WHETHER IN AN ACTION
|
||||||
|
OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, OUT OF OR IN CONNECTION
|
||||||
|
WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE.
|
||||||
+80
@@ -0,0 +1,80 @@
|
|||||||
|
[](https://pkg.go.dev/github.com/jackc/puddle/v2)
|
||||||
|

|
||||||
|
|
||||||
|
# Puddle
|
||||||
|
|
||||||
|
Puddle is a tiny generic resource pool library for Go that uses the standard
|
||||||
|
context library to signal cancellation of acquires. It is designed to contain
|
||||||
|
the minimum functionality required for a resource pool. It can be used directly
|
||||||
|
or it can be used as the base for a domain specific resource pool. For example,
|
||||||
|
a database connection pool may use puddle internally and implement health checks
|
||||||
|
and keep-alive behavior without needing to implement any concurrent code of its
|
||||||
|
own.
|
||||||
|
|
||||||
|
## Features
|
||||||
|
|
||||||
|
* Acquire cancellation via context standard library
|
||||||
|
* Statistics API for monitoring pool pressure
|
||||||
|
* No dependencies outside of standard library and golang.org/x/sync
|
||||||
|
* High performance
|
||||||
|
* 100% test coverage of reachable code
|
||||||
|
|
||||||
|
## Example Usage
|
||||||
|
|
||||||
|
```go
|
||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"log"
|
||||||
|
"net"
|
||||||
|
|
||||||
|
"github.com/jackc/puddle/v2"
|
||||||
|
)
|
||||||
|
|
||||||
|
func main() {
|
||||||
|
constructor := func(context.Context) (net.Conn, error) {
|
||||||
|
return net.Dial("tcp", "127.0.0.1:8080")
|
||||||
|
}
|
||||||
|
destructor := func(value net.Conn) {
|
||||||
|
value.Close()
|
||||||
|
}
|
||||||
|
maxPoolSize := int32(10)
|
||||||
|
|
||||||
|
pool, err := puddle.NewPool(&puddle.Config[net.Conn]{Constructor: constructor, Destructor: destructor, MaxSize: maxPoolSize})
|
||||||
|
if err != nil {
|
||||||
|
log.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Acquire resource from the pool.
|
||||||
|
res, err := pool.Acquire(context.Background())
|
||||||
|
if err != nil {
|
||||||
|
log.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Use resource.
|
||||||
|
_, err = res.Value().Write([]byte{1})
|
||||||
|
if err != nil {
|
||||||
|
log.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Release when done.
|
||||||
|
res.Release()
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
## Status
|
||||||
|
|
||||||
|
Puddle is stable and feature complete.
|
||||||
|
|
||||||
|
* Bug reports and fixes are welcome.
|
||||||
|
* New features will usually not be accepted if they can be feasibly implemented in a wrapper.
|
||||||
|
* Performance optimizations will usually not be accepted unless the performance issue rises to the level of a bug.
|
||||||
|
|
||||||
|
## Supported Go Versions
|
||||||
|
|
||||||
|
puddle supports the same versions of Go that are supported by the Go project. For [Go](https://golang.org/doc/devel/release.html#policy) that is the two most recent major releases. This means puddle supports Go 1.19 and higher.
|
||||||
|
|
||||||
|
## License
|
||||||
|
|
||||||
|
MIT
|
||||||
+24
@@ -0,0 +1,24 @@
|
|||||||
|
package puddle
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
// valueCancelCtx combines two contexts into one. One context is used for values and the other is used for cancellation.
|
||||||
|
type valueCancelCtx struct {
|
||||||
|
valueCtx context.Context
|
||||||
|
cancelCtx context.Context
|
||||||
|
}
|
||||||
|
|
||||||
|
func (ctx *valueCancelCtx) Deadline() (time.Time, bool) { return ctx.cancelCtx.Deadline() }
|
||||||
|
func (ctx *valueCancelCtx) Done() <-chan struct{} { return ctx.cancelCtx.Done() }
|
||||||
|
func (ctx *valueCancelCtx) Err() error { return ctx.cancelCtx.Err() }
|
||||||
|
func (ctx *valueCancelCtx) Value(key any) any { return ctx.valueCtx.Value(key) }
|
||||||
|
|
||||||
|
func newValueCancelCtx(valueCtx, cancelContext context.Context) context.Context {
|
||||||
|
return &valueCancelCtx{
|
||||||
|
valueCtx: valueCtx,
|
||||||
|
cancelCtx: cancelContext,
|
||||||
|
}
|
||||||
|
}
|
||||||
+11
@@ -0,0 +1,11 @@
|
|||||||
|
// Package puddle is a generic resource pool with type-parametrized api.
|
||||||
|
/*
|
||||||
|
|
||||||
|
Puddle is a tiny generic resource pool library for Go that uses the standard
|
||||||
|
context library to signal cancellation of acquires. It is designed to contain
|
||||||
|
the minimum functionality a resource pool needs that cannot be implemented
|
||||||
|
without concurrency concerns. For example, a database connection pool may use
|
||||||
|
puddle internally and implement health checks and keep-alive behavior without
|
||||||
|
needing to implement any concurrent code of its own.
|
||||||
|
*/
|
||||||
|
package puddle
|
||||||
+85
@@ -0,0 +1,85 @@
|
|||||||
|
package genstack
|
||||||
|
|
||||||
|
// GenStack implements a generational stack.
|
||||||
|
//
|
||||||
|
// GenStack works as common stack except for the fact that all elements in the
|
||||||
|
// older generation are guaranteed to be popped before any element in the newer
|
||||||
|
// generation. New elements are always pushed to the current (newest)
|
||||||
|
// generation.
|
||||||
|
//
|
||||||
|
// We could also say that GenStack behaves as a stack in case of a single
|
||||||
|
// generation, but it behaves as a queue of individual generation stacks.
|
||||||
|
type GenStack[T any] struct {
|
||||||
|
// We can represent arbitrary number of generations using 2 stacks. The
|
||||||
|
// new stack stores all new pushes and the old stack serves all reads.
|
||||||
|
// Old stack can represent multiple generations. If old == new, then all
|
||||||
|
// elements pushed in previous (not current) generations have already
|
||||||
|
// been popped.
|
||||||
|
|
||||||
|
old *stack[T]
|
||||||
|
new *stack[T]
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewGenStack creates a new empty GenStack.
|
||||||
|
func NewGenStack[T any]() *GenStack[T] {
|
||||||
|
s := &stack[T]{}
|
||||||
|
return &GenStack[T]{
|
||||||
|
old: s,
|
||||||
|
new: s,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *GenStack[T]) Pop() (T, bool) {
|
||||||
|
// Pushes always append to the new stack, so if the old once becomes
|
||||||
|
// empty, it will remail empty forever.
|
||||||
|
if s.old.len() == 0 && s.old != s.new {
|
||||||
|
s.old = s.new
|
||||||
|
}
|
||||||
|
|
||||||
|
if s.old.len() == 0 {
|
||||||
|
var zero T
|
||||||
|
return zero, false
|
||||||
|
}
|
||||||
|
|
||||||
|
return s.old.pop(), true
|
||||||
|
}
|
||||||
|
|
||||||
|
// Push pushes a new element at the top of the stack.
|
||||||
|
func (s *GenStack[T]) Push(v T) { s.new.push(v) }
|
||||||
|
|
||||||
|
// NextGen starts a new stack generation.
|
||||||
|
func (s *GenStack[T]) NextGen() {
|
||||||
|
if s.old == s.new {
|
||||||
|
s.new = &stack[T]{}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// We need to pop from the old stack to the top of the new stack. Let's
|
||||||
|
// have an example:
|
||||||
|
//
|
||||||
|
// Old: <bottom> 4 3 2 1
|
||||||
|
// New: <bottom> 8 7 6 5
|
||||||
|
// PopOrder: 1 2 3 4 5 6 7 8
|
||||||
|
//
|
||||||
|
//
|
||||||
|
// To preserve pop order, we have to take all elements from the old
|
||||||
|
// stack and push them to the top of new stack:
|
||||||
|
//
|
||||||
|
// New: 8 7 6 5 4 3 2 1
|
||||||
|
//
|
||||||
|
s.new.push(s.old.takeAll()...)
|
||||||
|
|
||||||
|
// We have the old stack allocated and empty, so why not to reuse it as
|
||||||
|
// new new stack.
|
||||||
|
s.old, s.new = s.new, s.old
|
||||||
|
}
|
||||||
|
|
||||||
|
// Len returns number of elements in the stack.
|
||||||
|
func (s *GenStack[T]) Len() int {
|
||||||
|
l := s.old.len()
|
||||||
|
if s.old != s.new {
|
||||||
|
l += s.new.len()
|
||||||
|
}
|
||||||
|
|
||||||
|
return l
|
||||||
|
}
|
||||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user