Compare commits

..

24 Commits

Author SHA1 Message Date
Hein b95b74f0a3 chore(release): update package version to 1.0.65
Release / test (push) Successful in 1m40s
Release / release (push) Successful in 2m4s
Release / pkg-aur (push) Failing after 20s
Release / pkg-deb (push) Successful in 3m17s
Release / pkg-rpm (push) Successful in 3m41s
2026-07-21 12:45:28 +02:00
Hein 5d9ff5df03 feat(bun): generate native Go array slices for PostgreSQL array columns
Bun's pgdialect scans/appends native slices directly, so array columns
(text[], integer[], uuid[], ...) always generate as plain []string,
[]int32, etc. with an explicit "array" bun tag, regardless of --types
(sqltypes/stdlib/baselib). The SqlXxxArray wrapper types are no longer
used for Bun array columns (gorm is unaffected and keeps using them).

Adds --array-nullable pointer_slice to represent nullable array columns
as *[]T instead of []T, so callers can distinguish SQL NULL (nil) from
'{}' (pointer to an empty slice). Verified end-to-end against a live
PostgreSQL instance for NULL/{}/populated arrays in every --types mode.

Closes #13
2026-07-21 12:41:38 +02:00
Hein 2cecb4c11c feat(cli): add report command for filing bugs/features from the CLI
Adds `relspec report bug|feature <title>` which files an issue directly
against the RelSpec tracker, tagging the title with the system's OS
machine-id (falling back to a persisted UUID) and rate-limiting
submissions to one per minute via a local state file.

Closes #14
2026-07-21 10:31:20 +02:00
Hein 316d9b0e7f chore(release): update package version to 1.0.64
Release / test (push) Successful in 35s
Release / release (push) Successful in 40s
Release / pkg-deb (push) Successful in 54s
Release / pkg-aur (push) Successful in 1m1s
Release / pkg-rpm (push) Successful in 2m59s
2026-07-20 13:59:44 +02:00
Hein 17ae8e050a fix(assetloader): name embedDirectiveLiteral return values to satisfy gocritic 2026-07-20 13:59:19 +02:00
Hein f0410221d8 fix(bun): use PostgreSQL internal array type name for sqltypes array columns
bun's pgdialect overrides Field.Scan/Append with its own slice-only array
handling whenever the tag's type: value ends in "[]", clobbering the
sql.Scanner/driver.Valuer implemented on SqlXxxArray wrapper types and
causing "bun: Scan(unsupported sqltypes.SqlStringArray)" at query time.
Emit the underscore-prefixed internal type name (e.g. _text) instead,
which is DDL-valid but doesn't end in "[]" so bun leaves our scanner alone.
2026-07-20 13:58:24 +02:00
warkanum 1c217b546c Merge pull request 'feat(scripts): support external file embedding' (#12) from issue-6-external-file-embedding into master
Reviewed-on: #12
Reviewed-by: Warky <2+warkanum@noreply@warky.dev>
2026-07-20 11:09:39 +00:00
SG Command 1bcdf29206 feat(scripts): support external file embedding 2026-07-20 00:13:05 +02:00
sgcommand 5c31deb630 Merge pull request #11: fix deterministic template table index ordering 2026-07-19 14:11:11 +00:00
SG Command c2def00bcf fix(template): make map helper ordering deterministic 2026-07-19 15:19:33 +02:00
warkanum 784dc1f0da chore(release): update package version to 1.0.63
Release / test (push) Successful in 52s
Release / release (push) Successful in 1m45s
Release / pkg-aur (push) Successful in 1m1s
Release / pkg-deb (push) Successful in 2m48s
Release / pkg-rpm (push) Successful in 2m49s
2026-07-18 22:41:30 +02:00
warkanum 7d93bee4bd chore: Fixed linitng issues 2026-07-18 22:41:23 +02:00
warkanum 2aecd1312e Merge pull request 'feat(assets): add native Go asset/file loader for migrate-apply (#7)' (#8) from issue-7-native-asset-loader into master
Reviewed-on: #8
2026-07-18 20:35:45 +00:00
warkanum 60c5cc40b2 feat(assets): add native Go asset/file loader for migrate-apply
Implements a new `relspec assets` command (list/execute subcommands) that
loads local binary and text asset files into PostgreSQL by binding file bytes
as native pgx query parameters — never as SQL text literals — so binary data
stays byte-exact with no escaping overhead.

Key design points:
- YAML manifest (assets.yaml) colocated with files describes each entry:
  file path, SQL call with :bytes/:filename/:param named placeholders, and
  optional static params map.
- Placeholder substitution converts :name to positional $N params; PostgreSQL
  ::cast syntax is protected before substitution to avoid false matches.
- Directory scan follows the existing {priority}_{sequence}_{name} naming
  convention, enabling asset-loading steps to be correctly interleaved with
  relspec scripts execute in a migrate-apply pipeline.
- Symlink components and path traversal (../) are silently skipped to prevent
  directory escape attacks.
- 14 unit tests cover manifest loading, directory scanning, ordering, symlink
  skipping, path traversal rejection, placeholder substitution edge cases
  (repeated, cast protection, binary byte-exact, unknown).

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
2026-07-17 01:14:11 +02:00
warkanum 5edb004799 Merge pull request 'fix(bun): support extra generated model fields' (#5) from fix/bun-extra-fields-issue-4 into master
Reviewed-on: #5
Reviewed-by: Warky <warkanum@warky.dev>
2026-07-09 04:53:33 +00:00
Hermes Agent 764d00c249 fix(bun): read extra fields from file 2026-07-09 06:02:01 +02:00
Hermes Agent 47b77763cb fix(bun): support extra generated model fields 2026-07-09 05:50:33 +02:00
Hein 7805d9b6f0 chore(release): update package version to 1.0.62
Release / test (push) Successful in 2m34s
Release / release (push) Successful in 1m43s
Release / pkg-aur (push) Successful in 1m2s
Release / pkg-deb (push) Successful in 2m57s
Release / pkg-rpm (push) Successful in 3m2s
2026-07-07 15:29:55 +02:00
Hein 40a0e6a0aa refactor: ♻️ change resolvspec types to build in sqltypes 2026-07-07 15:28:01 +02:00
Hein 99d63aa5f4 test(uuid): add integration tests for SqlUUID with database
* Implement tests for inserting, updating, and retrieving UUIDs in a SQLite database.
* Verify that Value() returns a string representation of the UUID.
* Ensure compatibility with custom types implementing fmt.Stringer.
* Update type mapper imports to reflect new package path.
2026-07-02 17:06:41 +02:00
warkanum ee94ddc133 chore(release): update package version to 1.0.61
Release / test (push) Successful in 1m12s
Release / release (push) Successful in 1m40s
Release / pkg-aur (push) Successful in 1m2s
Release / pkg-deb (push) Successful in 45s
Release / pkg-rpm (push) Successful in 1m18s
2026-06-26 00:27:33 +02:00
warkanum 651c7aa3f4 fix(reflectutil): correct pointer type check in dereference logic 2026-06-26 00:27:25 +02:00
warkanum 1cd9cd8803 chore(release): update package version to 1.0.60
Release / test (push) Failing after 0s
Release / release (push) Has been skipped
Release / pkg-aur (push) Has been skipped
Release / pkg-deb (push) Has been skipped
Release / pkg-rpm (push) Has been skipped
2026-06-25 23:26:08 +02:00
warkanum ab735d1f3a fix(writer): update nullable type handling to use baselib
* change default nullable type from resolvespec to baselib
* update type mappings in tests to reflect new nullable types
* adjust comments for clarity on nullable type options
2026-06-25 23:25:36 +02:00
111 changed files with 16745 additions and 281 deletions
+18
View File
@@ -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)
+214
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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,
+2 -1
View File
@@ -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,
} }
+243
View File
@@ -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
}
+2
View File
@@ -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")
} }
+5 -1
View File
@@ -23,6 +23,7 @@ var (
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)
+17
View File
@@ -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
+11 -11
View File
@@ -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"`
+20 -20
View File
@@ -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"`
+12 -12
View File
@@ -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"`
+8 -8
View File
@@ -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"`
+3
View File
@@ -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
+2
View File
@@ -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
View File
@@ -1,6 +1,6 @@
# Maintainer: Hein (Warky Devs) <hein@warky.dev> # Maintainer: Hein (Warky Devs) <hein@warky.dev>
pkgname=relspec pkgname=relspec
pkgver=1.0.59 pkgver=1.0.65
pkgrel=1 pkgrel=1
pkgdesc="RelSpec is a comprehensive database relations management tool that reads, transforms, and writes database table specifications across multiple formats and ORMs." pkgdesc="RelSpec is a comprehensive database relations management tool that reads, transforms, and writes database table specifications across multiple formats and ORMs."
arch=('x86_64' 'aarch64') arch=('x86_64' 'aarch64')
+1 -1
View File
@@ -1,5 +1,5 @@
Name: relspec Name: relspec
Version: 1.0.59 Version: 1.0.65
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.
+165
View File
@@ -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_]|$)`)
}
+143
View File
@@ -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)
}
}
+106
View File
@@ -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
}
+200
View File
@@ -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
}
+341
View File
@@ -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)
}
}
+2
View File
@@ -360,6 +360,7 @@ type Script struct {
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 +469,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(),
} }
} }
+2 -2
View File
@@ -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")
} }
} }
+15
View File
@@ -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 -1
View File
@@ -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)
+60
View File
@@ -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"
@@ -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)
}
}
+33 -6
View File
@@ -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{} {
+136
View File
@@ -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
+485
View File
@@ -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")
}
})
}
+981
View File
@@ -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}}
}
+134
View File
@@ -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)
}
})
}
+958
View File
@@ -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)
}
}
+678
View File
@@ -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)
}
}
+164
View File
@@ -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)
}
}
+180
View File
@@ -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)
}
+116 -19
View File
@@ -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,34 +48,40 @@ 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"
) )
@@ -82,9 +90,9 @@ type User struct {
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
+141
View File
@@ -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 := &notNullArrayRow{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 := &notNullArrayRow{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)
}
})
}
+89 -3
View File
@@ -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
@@ -141,10 +227,10 @@ func NewModelData(table *models.Table, schema string, typeMapper *TypeMapper, fl
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 {
@@ -203,7 +289,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
+7
View File
@@ -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}}
+73 -42
View File
@@ -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")
} }
} }
@@ -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())
} }
+46 -3
View File
@@ -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()
+241 -32
View File
@@ -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,245 @@ 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_ArraysAreNativeInEveryMode verifies that array
mapper := NewTypeMapper(writers.NullableTypeStdlib) // 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 { cases := []struct {
name string name string
column *models.Column column *models.Column
wantSubstr string
}{ }{
{name: "text array", column: &models.Column{Name: "tags", Type: "text[]"}}, {name: "text array", column: &models.Column{Name: "tags", Type: "text[]"}, wantSubstr: "type:text[],array,"},
{name: "integer array", column: &models.Column{Name: "scores", Type: "integer[]", NotNull: true}}, {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 { for _, tt := range cases {
t.Run(tt.name, func(t *testing.T) { t.Run(tt.name, func(t *testing.T) {
result := mapper.BuildBunTag(tt.column, nil) result := mapper.BuildBunTag(tt.column, nil)
if !strings.Contains(result, "array") { if !strings.Contains(result, tt.wantSubstr) {
t.Errorf("BuildBunTag() = %q, expected 'array' in stdlib mode", result) 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",
+11 -9
View File
@@ -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` |
+1 -1
View File
@@ -202,7 +202,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
+51 -10
View File
@@ -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{
@@ -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())
} }
+2 -2
View File
@@ -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
+9 -9
View File
@@ -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 {
+31 -1
View File
@@ -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)
}
+37
View File
@@ -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)
}
}
+12 -7
View File
@@ -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:
+87
View File
@@ -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
}
+36 -4
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -0,0 +1,80 @@
[![Go Reference](https://pkg.go.dev/badge/github.com/jackc/puddle/v2.svg)](https://pkg.go.dev/github.com/jackc/puddle/v2)
![Build Status](https://github.com/jackc/puddle/actions/workflows/ci.yml/badge.svg)
# 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
View File
@@ -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
View File
@@ -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
View File
@@ -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
}
+39
View File
@@ -0,0 +1,39 @@
package genstack
// stack is a wrapper around an array implementing a stack.
//
// We cannot use slice to represent the stack because append might change the
// pointer value of the slice. That would be an issue in GenStack
// implementation.
type stack[T any] struct {
arr []T
}
// push pushes a new element at the top of a stack.
func (s *stack[T]) push(vs ...T) { s.arr = append(s.arr, vs...) }
// pop pops the stack top-most element.
//
// If stack length is zero, this method panics.
func (s *stack[T]) pop() T {
idx := s.len() - 1
val := s.arr[idx]
// Avoid memory leak
var zero T
s.arr[idx] = zero
s.arr = s.arr[:idx]
return val
}
// takeAll returns all elements in the stack in order as they are stored - i.e.
// the top-most stack element is the last one.
func (s *stack[T]) takeAll() []T {
arr := s.arr
s.arr = nil
return arr
}
// len returns number of elements in the stack.
func (s *stack[T]) len() int { return len(s.arr) }
+32
View File
@@ -0,0 +1,32 @@
package puddle
import "unsafe"
type ints interface {
int | int8 | int16 | int32 | int64 | uint | uint8 | uint16 | uint32 | uint64
}
// log2Int returns log2 of an integer. This function panics if val < 0. For val
// == 0, returns 0.
func log2Int[T ints](val T) uint8 {
if val <= 0 {
panic("log2 of non-positive number does not exist")
}
return log2IntRange(val, 0, uint8(8*unsafe.Sizeof(val)))
}
func log2IntRange[T ints](val T, begin, end uint8) uint8 {
length := end - begin
if length == 1 {
return begin
}
delim := begin + length/2
mask := T(1) << delim
if mask > val {
return log2IntRange(val, begin, delim)
} else {
return log2IntRange(val, delim, end)
}
}
+16
View File
@@ -0,0 +1,16 @@
package puddle
import "time"
// nanotime returns the time in nanoseconds since process start.
//
// This approach, described at
// https://github.com/golang/go/issues/61765#issuecomment-1672090302,
// is fast, monotonic, and portable, and avoids the previous
// dependence on runtime.nanotime using the (unsafe) linkname hack.
// In particular, time.Since does less work than time.Now.
func nanotime() int64 {
return time.Since(globalStart).Nanoseconds()
}
var globalStart = time.Now()
+710
View File
@@ -0,0 +1,710 @@
package puddle
import (
"context"
"errors"
"sync"
"sync/atomic"
"time"
"github.com/jackc/puddle/v2/internal/genstack"
"golang.org/x/sync/semaphore"
)
const (
resourceStatusConstructing = 0
resourceStatusIdle = iota
resourceStatusAcquired = iota
resourceStatusHijacked = iota
)
// ErrClosedPool occurs on an attempt to acquire a connection from a closed pool
// or a pool that is closed while the acquire is waiting.
var ErrClosedPool = errors.New("closed pool")
// ErrNotAvailable occurs on an attempt to acquire a resource from a pool
// that is at maximum capacity and has no available resources.
var ErrNotAvailable = errors.New("resource not available")
// Constructor is a function called by the pool to construct a resource.
type Constructor[T any] func(ctx context.Context) (res T, err error)
// Destructor is a function called by the pool to destroy a resource.
type Destructor[T any] func(res T)
// Resource is the resource handle returned by acquiring from the pool.
type Resource[T any] struct {
value T
pool *Pool[T]
creationTime time.Time
lastUsedNano int64
poolResetCount int
status byte
}
// Value returns the resource value.
func (res *Resource[T]) Value() T {
if !(res.status == resourceStatusAcquired || res.status == resourceStatusHijacked) {
panic("tried to access resource that is not acquired or hijacked")
}
return res.value
}
// Release returns the resource to the pool. res must not be subsequently used.
func (res *Resource[T]) Release() {
if res.status != resourceStatusAcquired {
panic("tried to release resource that is not acquired")
}
res.pool.releaseAcquiredResource(res, nanotime())
}
// ReleaseUnused returns the resource to the pool without updating when it was last used used. i.e. LastUsedNanotime
// will not change. res must not be subsequently used.
func (res *Resource[T]) ReleaseUnused() {
if res.status != resourceStatusAcquired {
panic("tried to release resource that is not acquired")
}
res.pool.releaseAcquiredResource(res, res.lastUsedNano)
}
// Destroy returns the resource to the pool for destruction. res must not be
// subsequently used.
func (res *Resource[T]) Destroy() {
if res.status != resourceStatusAcquired {
panic("tried to destroy resource that is not acquired")
}
go res.pool.destroyAcquiredResource(res)
}
// Hijack assumes ownership of the resource from the pool. Caller is responsible
// for cleanup of resource value.
func (res *Resource[T]) Hijack() {
if res.status != resourceStatusAcquired {
panic("tried to hijack resource that is not acquired")
}
res.pool.hijackAcquiredResource(res)
}
// CreationTime returns when the resource was created by the pool.
func (res *Resource[T]) CreationTime() time.Time {
if !(res.status == resourceStatusAcquired || res.status == resourceStatusHijacked) {
panic("tried to access resource that is not acquired or hijacked")
}
return res.creationTime
}
// LastUsedNanotime returns when Release was last called on the resource measured in nanoseconds from an arbitrary time
// (a monotonic time). Returns creation time if Release has never been called. This is only useful to compare with
// other calls to LastUsedNanotime. In almost all cases, IdleDuration should be used instead.
func (res *Resource[T]) LastUsedNanotime() int64 {
if !(res.status == resourceStatusAcquired || res.status == resourceStatusHijacked) {
panic("tried to access resource that is not acquired or hijacked")
}
return res.lastUsedNano
}
// IdleDuration returns the duration since Release was last called on the resource. This is equivalent to subtracting
// LastUsedNanotime to the current nanotime.
func (res *Resource[T]) IdleDuration() time.Duration {
if !(res.status == resourceStatusAcquired || res.status == resourceStatusHijacked) {
panic("tried to access resource that is not acquired or hijacked")
}
return time.Duration(nanotime() - res.lastUsedNano)
}
// Pool is a concurrency-safe resource pool.
type Pool[T any] struct {
// mux is the pool internal lock. Any modification of shared state of
// the pool (but Acquires of acquireSem) must be performed only by
// holder of the lock. Long running operations are not allowed when mux
// is held.
mux sync.Mutex
// acquireSem provides an allowance to acquire a resource.
//
// Releases are allowed only when caller holds mux. Acquires have to
// happen before mux is locked (doesn't apply to semaphore.TryAcquire in
// AcquireAllIdle).
acquireSem *semaphore.Weighted
destructWG sync.WaitGroup
allResources resList[T]
idleResources *genstack.GenStack[*Resource[T]]
constructor Constructor[T]
destructor Destructor[T]
maxSize int32
acquireCount int64
acquireDuration time.Duration
emptyAcquireCount int64
emptyAcquireWaitTime time.Duration
canceledAcquireCount atomic.Int64
resetCount int
baseAcquireCtx context.Context
cancelBaseAcquireCtx context.CancelFunc
closed bool
}
type Config[T any] struct {
Constructor Constructor[T]
Destructor Destructor[T]
MaxSize int32
}
// NewPool creates a new pool. Returns an error iff MaxSize is less than 1.
func NewPool[T any](config *Config[T]) (*Pool[T], error) {
if config.MaxSize < 1 {
return nil, errors.New("MaxSize must be >= 1")
}
baseAcquireCtx, cancelBaseAcquireCtx := context.WithCancel(context.Background())
return &Pool[T]{
acquireSem: semaphore.NewWeighted(int64(config.MaxSize)),
idleResources: genstack.NewGenStack[*Resource[T]](),
maxSize: config.MaxSize,
constructor: config.Constructor,
destructor: config.Destructor,
baseAcquireCtx: baseAcquireCtx,
cancelBaseAcquireCtx: cancelBaseAcquireCtx,
}, nil
}
// Close destroys all resources in the pool and rejects future Acquire calls.
// Blocks until all resources are returned to pool and destroyed.
func (p *Pool[T]) Close() {
defer p.destructWG.Wait()
p.mux.Lock()
defer p.mux.Unlock()
if p.closed {
return
}
p.closed = true
p.cancelBaseAcquireCtx()
for res, ok := p.idleResources.Pop(); ok; res, ok = p.idleResources.Pop() {
p.allResources.remove(res)
go p.destructResourceValue(res.value)
}
}
// Stat is a snapshot of Pool statistics.
type Stat struct {
constructingResources int32
acquiredResources int32
idleResources int32
maxResources int32
acquireCount int64
acquireDuration time.Duration
emptyAcquireCount int64
emptyAcquireWaitTime time.Duration
canceledAcquireCount int64
}
// TotalResources returns the total number of resources currently in the pool.
// The value is the sum of ConstructingResources, AcquiredResources, and
// IdleResources.
func (s *Stat) TotalResources() int32 {
return s.constructingResources + s.acquiredResources + s.idleResources
}
// ConstructingResources returns the number of resources with construction in progress in
// the pool.
func (s *Stat) ConstructingResources() int32 {
return s.constructingResources
}
// AcquiredResources returns the number of currently acquired resources in the pool.
func (s *Stat) AcquiredResources() int32 {
return s.acquiredResources
}
// IdleResources returns the number of currently idle resources in the pool.
func (s *Stat) IdleResources() int32 {
return s.idleResources
}
// MaxResources returns the maximum size of the pool.
func (s *Stat) MaxResources() int32 {
return s.maxResources
}
// AcquireCount returns the cumulative count of successful acquires from the pool.
func (s *Stat) AcquireCount() int64 {
return s.acquireCount
}
// AcquireDuration returns the total duration of all successful acquires from
// the pool.
func (s *Stat) AcquireDuration() time.Duration {
return s.acquireDuration
}
// 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.emptyAcquireCount
}
// 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.emptyAcquireWaitTime
}
// CanceledAcquireCount returns the cumulative count of acquires from the pool
// that were canceled by a context.
func (s *Stat) CanceledAcquireCount() int64 {
return s.canceledAcquireCount
}
// Stat returns the current pool statistics.
func (p *Pool[T]) Stat() *Stat {
p.mux.Lock()
defer p.mux.Unlock()
s := &Stat{
maxResources: p.maxSize,
acquireCount: p.acquireCount,
emptyAcquireCount: p.emptyAcquireCount,
emptyAcquireWaitTime: p.emptyAcquireWaitTime,
canceledAcquireCount: p.canceledAcquireCount.Load(),
acquireDuration: p.acquireDuration,
}
for _, res := range p.allResources {
switch res.status {
case resourceStatusConstructing:
s.constructingResources += 1
case resourceStatusIdle:
s.idleResources += 1
case resourceStatusAcquired:
s.acquiredResources += 1
}
}
return s
}
// tryAcquireIdleResource checks if there is any idle resource. If there is
// some, this method removes it from idle list and returns it. If the idle pool
// is empty, this method returns nil and doesn't modify the idleResources slice.
//
// WARNING: Caller of this method must hold the pool mutex!
func (p *Pool[T]) tryAcquireIdleResource() *Resource[T] {
res, ok := p.idleResources.Pop()
if !ok {
return nil
}
res.status = resourceStatusAcquired
return res
}
// createNewResource creates a new resource and inserts it into list of pool
// resources.
//
// WARNING: Caller of this method must hold the pool mutex!
func (p *Pool[T]) createNewResource() *Resource[T] {
res := &Resource[T]{
pool: p,
creationTime: time.Now(),
lastUsedNano: nanotime(),
poolResetCount: p.resetCount,
status: resourceStatusConstructing,
}
p.allResources.append(res)
p.destructWG.Add(1)
return res
}
// Acquire gets a resource from the pool. If no resources are available and the pool is not at maximum capacity it will
// create a new resource. If the pool is at maximum capacity it will block until a resource is available. ctx can be
// used to cancel the Acquire.
//
// If Acquire creates a new resource the resource constructor function will receive a context that delegates Value() to
// ctx. Canceling ctx will cause Acquire to return immediately but it will not cancel the resource creation. This avoids
// the problem of it being impossible to create resources when the time to create a resource is greater than any one
// caller of Acquire is willing to wait.
func (p *Pool[T]) Acquire(ctx context.Context) (_ *Resource[T], err error) {
select {
case <-ctx.Done():
p.canceledAcquireCount.Add(1)
return nil, ctx.Err()
default:
}
return p.acquire(ctx)
}
// acquire is a continuation of Acquire function that doesn't check context
// validity.
//
// This function exists solely only for benchmarking purposes.
func (p *Pool[T]) acquire(ctx context.Context) (*Resource[T], error) {
startNano := nanotime()
var waitedForLock bool
if !p.acquireSem.TryAcquire(1) {
waitedForLock = true
err := p.acquireSem.Acquire(ctx, 1)
if err != nil {
p.canceledAcquireCount.Add(1)
return nil, err
}
}
p.mux.Lock()
if p.closed {
p.acquireSem.Release(1)
p.mux.Unlock()
return nil, ErrClosedPool
}
// If a resource is available in the pool.
if res := p.tryAcquireIdleResource(); res != nil {
waitTime := time.Duration(nanotime() - startNano)
if waitedForLock {
p.emptyAcquireCount += 1
p.emptyAcquireWaitTime += waitTime
}
p.acquireCount += 1
p.acquireDuration += waitTime
p.mux.Unlock()
return res, nil
}
if len(p.allResources) >= int(p.maxSize) {
// Unreachable code.
panic("bug: semaphore allowed more acquires than pool allows")
}
// The resource is not idle, but there is enough space to create one.
res := p.createNewResource()
p.mux.Unlock()
res, err := p.initResourceValue(ctx, res)
if err != nil {
return nil, err
}
p.mux.Lock()
defer p.mux.Unlock()
p.emptyAcquireCount += 1
p.acquireCount += 1
waitTime := time.Duration(nanotime() - startNano)
p.acquireDuration += waitTime
p.emptyAcquireWaitTime += waitTime
return res, nil
}
func (p *Pool[T]) initResourceValue(ctx context.Context, res *Resource[T]) (*Resource[T], error) {
// Create the resource in a goroutine to immediately return from Acquire
// if ctx is canceled without also canceling the constructor.
//
// See:
// - https://github.com/jackc/pgx/issues/1287
// - https://github.com/jackc/pgx/issues/1259
constructErrChan := make(chan error)
go func() {
constructorCtx := newValueCancelCtx(ctx, p.baseAcquireCtx)
value, err := p.constructor(constructorCtx)
if err != nil {
p.mux.Lock()
p.allResources.remove(res)
p.destructWG.Done()
// The resource won't be acquired because its
// construction failed. We have to allow someone else to
// take that resouce.
p.acquireSem.Release(1)
p.mux.Unlock()
select {
case constructErrChan <- err:
case <-ctx.Done():
// The caller is cancelled, so no-one awaits the
// error. This branch avoid goroutine leak.
}
return
}
// The resource is already in p.allResources where it might be read. So we need to acquire the lock to update its
// status.
p.mux.Lock()
res.value = value
res.status = resourceStatusAcquired
p.mux.Unlock()
// This select works because the channel is unbuffered.
select {
case constructErrChan <- nil:
case <-ctx.Done():
p.releaseAcquiredResource(res, res.lastUsedNano)
}
}()
select {
case <-ctx.Done():
p.canceledAcquireCount.Add(1)
return nil, ctx.Err()
case err := <-constructErrChan:
if err != nil {
return nil, err
}
return res, nil
}
}
// TryAcquire gets a resource from the pool if one is immediately available. If not, it returns ErrNotAvailable. If no
// resources are available but the pool has room to grow, a resource will be created in the background. ctx is only
// used to cancel the background creation.
func (p *Pool[T]) TryAcquire(ctx context.Context) (*Resource[T], error) {
if !p.acquireSem.TryAcquire(1) {
return nil, ErrNotAvailable
}
p.mux.Lock()
defer p.mux.Unlock()
if p.closed {
p.acquireSem.Release(1)
return nil, ErrClosedPool
}
// If a resource is available now
if res := p.tryAcquireIdleResource(); res != nil {
p.acquireCount += 1
return res, nil
}
if len(p.allResources) >= int(p.maxSize) {
// Unreachable code.
panic("bug: semaphore allowed more acquires than pool allows")
}
res := p.createNewResource()
go func() {
value, err := p.constructor(ctx)
p.mux.Lock()
defer p.mux.Unlock()
// We have to create the resource and only then release the
// semaphore - For the time being there is no resource that
// someone could acquire.
defer p.acquireSem.Release(1)
if err != nil {
p.allResources.remove(res)
p.destructWG.Done()
return
}
res.value = value
res.status = resourceStatusIdle
p.idleResources.Push(res)
}()
return nil, ErrNotAvailable
}
// acquireSemAll tries to acquire num free tokens from sem. This function is
// guaranteed to acquire at least the lowest number of tokens that has been
// available in the semaphore during runtime of this function.
//
// For the time being, semaphore doesn't allow to acquire all tokens atomically
// (see https://github.com/golang/sync/pull/19). We simulate this by trying all
// powers of 2 that are less or equal to num.
//
// For example, let's immagine we have 19 free tokens in the semaphore which in
// total has 24 tokens (i.e. the maxSize of the pool is 24 resources). Then if
// num is 24, the log2Uint(24) is 4 and we try to acquire 16, 8, 4, 2 and 1
// tokens. Out of those, the acquire of 16, 2 and 1 tokens will succeed.
//
// Naturally, Acquires and Releases of the semaphore might take place
// concurrently. For this reason, it's not guaranteed that absolutely all free
// tokens in the semaphore will be acquired. But it's guaranteed that at least
// the minimal number of tokens that has been present over the whole process
// will be acquired. This is sufficient for the use-case we have in this
// package.
//
// TODO: Replace this with acquireSem.TryAcquireAll() if it gets to
// upstream. https://github.com/golang/sync/pull/19
func acquireSemAll(sem *semaphore.Weighted, num int) int {
if sem.TryAcquire(int64(num)) {
return num
}
var acquired int
for i := int(log2Int(num)); i >= 0; i-- {
val := 1 << i
if sem.TryAcquire(int64(val)) {
acquired += val
}
}
return acquired
}
// AcquireAllIdle acquires all currently idle resources. Its intended use is for
// health check and keep-alive functionality. It does not update pool
// statistics.
func (p *Pool[T]) AcquireAllIdle() []*Resource[T] {
p.mux.Lock()
defer p.mux.Unlock()
if p.closed {
return nil
}
numIdle := p.idleResources.Len()
if numIdle == 0 {
return nil
}
// In acquireSemAll we use only TryAcquire and not Acquire. Because
// TryAcquire cannot block, the fact that we hold mutex locked and try
// to acquire semaphore cannot result in dead-lock.
//
// Because the mutex is locked, no parallel Release can run. This
// implies that the number of tokens can only decrease because some
// Acquire/TryAcquire call can consume the semaphore token. Consequently
// acquired is always less or equal to numIdle. Moreover if acquired <
// numIdle, then there are some parallel Acquire/TryAcquire calls that
// will take the remaining idle connections.
acquired := acquireSemAll(p.acquireSem, numIdle)
idle := make([]*Resource[T], acquired)
for i := range idle {
res, _ := p.idleResources.Pop()
res.status = resourceStatusAcquired
idle[i] = res
}
// We have to bump the generation to ensure that Acquire/TryAcquire
// calls running in parallel (those which caused acquired < numIdle)
// will consume old connections and not freshly released connections
// instead.
p.idleResources.NextGen()
return idle
}
// CreateResource constructs a new resource without acquiring it. It goes straight in the IdlePool. If the pool is full
// it returns an error. It can be useful to maintain warm resources under little load.
func (p *Pool[T]) CreateResource(ctx context.Context) error {
if !p.acquireSem.TryAcquire(1) {
return ErrNotAvailable
}
p.mux.Lock()
if p.closed {
p.acquireSem.Release(1)
p.mux.Unlock()
return ErrClosedPool
}
if len(p.allResources) >= int(p.maxSize) {
p.acquireSem.Release(1)
p.mux.Unlock()
return ErrNotAvailable
}
res := p.createNewResource()
p.mux.Unlock()
value, err := p.constructor(ctx)
p.mux.Lock()
defer p.mux.Unlock()
defer p.acquireSem.Release(1)
if err != nil {
p.allResources.remove(res)
p.destructWG.Done()
return err
}
res.value = value
res.status = resourceStatusIdle
// If closed while constructing resource then destroy it and return an error
if p.closed {
go p.destructResourceValue(res.value)
return ErrClosedPool
}
p.idleResources.Push(res)
return nil
}
// Reset destroys all resources, but leaves the pool open. It is intended for use when an error is detected that would
// disrupt all resources (such as a network interruption or a server state change).
//
// It is safe to reset a pool while resources are checked out. Those resources will be destroyed when they are returned
// to the pool.
func (p *Pool[T]) Reset() {
p.mux.Lock()
defer p.mux.Unlock()
p.resetCount++
for res, ok := p.idleResources.Pop(); ok; res, ok = p.idleResources.Pop() {
p.allResources.remove(res)
go p.destructResourceValue(res.value)
}
}
// releaseAcquiredResource returns res to the the pool.
func (p *Pool[T]) releaseAcquiredResource(res *Resource[T], lastUsedNano int64) {
p.mux.Lock()
defer p.mux.Unlock()
defer p.acquireSem.Release(1)
if p.closed || res.poolResetCount != p.resetCount {
p.allResources.remove(res)
go p.destructResourceValue(res.value)
} else {
res.lastUsedNano = lastUsedNano
res.status = resourceStatusIdle
p.idleResources.Push(res)
}
}
// Remove removes res from the pool and closes it. If res is not part of the
// pool Remove will panic.
func (p *Pool[T]) destroyAcquiredResource(res *Resource[T]) {
p.destructResourceValue(res.value)
p.mux.Lock()
defer p.mux.Unlock()
defer p.acquireSem.Release(1)
p.allResources.remove(res)
}
func (p *Pool[T]) hijackAcquiredResource(res *Resource[T]) {
p.mux.Lock()
defer p.mux.Unlock()
defer p.acquireSem.Release(1)
p.allResources.remove(res)
res.status = resourceStatusHijacked
p.destructWG.Done() // not responsible for destructing hijacked resources
}
func (p *Pool[T]) destructResourceValue(value T) {
p.destructor(value)
p.destructWG.Done()
}
+28
View File
@@ -0,0 +1,28 @@
package puddle
type resList[T any] []*Resource[T]
func (l *resList[T]) append(val *Resource[T]) { *l = append(*l, val) }
func (l *resList[T]) popBack() *Resource[T] {
idx := len(*l) - 1
val := (*l)[idx]
(*l)[idx] = nil // Avoid memory leak
*l = (*l)[:idx]
return val
}
func (l *resList[T]) remove(val *Resource[T]) {
for i, elem := range *l {
if elem == val {
lastIdx := len(*l) - 1
(*l)[i] = (*l)[lastIdx]
(*l)[lastIdx] = nil // Avoid memory leak
(*l) = (*l)[:lastIdx]
return
}
}
panic("BUG: removeResource could not find res in slice")
}
+24
View File
@@ -0,0 +1,24 @@
Copyright (c) 2021 Vladimir Mihailenco. All rights reserved.
Redistribution and use in source and binary forms, with or without
modification, are permitted provided that the following conditions are
met:
* Redistributions of source code must retain the above copyright
notice, this list of conditions and the following disclaimer.
* Redistributions in binary form must reproduce the above
copyright notice, this list of conditions and the following disclaimer
in the documentation and/or other materials provided with the
distribution.
THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS
"AS IS" AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT
LIMITED TO, THE IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR
A PARTICULAR PURPOSE ARE DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT
OWNER OR CONTRIBUTORS BE LIABLE FOR ANY DIRECT, INDIRECT, INCIDENTAL,
SPECIAL, EXEMPLARY, OR CONSEQUENTIAL DAMAGES (INCLUDING, BUT NOT
LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES; LOSS OF USE,
DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER CAUSED AND ON ANY
THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, OR TORT
(INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
+284
View File
@@ -0,0 +1,284 @@
package pgdialect
import (
"context"
"fmt"
"strings"
"github.com/uptrace/bun"
"github.com/uptrace/bun/migrate"
"github.com/uptrace/bun/migrate/sqlschema"
"github.com/uptrace/bun/schema"
)
func (d *Dialect) NewMigrator(db *bun.DB, schemaName string) sqlschema.Migrator {
return &migrator{db: db, schemaName: schemaName, BaseMigrator: sqlschema.NewBaseMigrator(db)}
}
type migrator struct {
*sqlschema.BaseMigrator
db *bun.DB
schemaName string
}
var _ sqlschema.Migrator = (*migrator)(nil)
func (m *migrator) AppendSQL(b []byte, operation any) (_ []byte, err error) {
gen := m.db.QueryGen()
// Append ALTER TABLE statement to the enclosed query bytes []byte.
appendAlterTable := func(query []byte, tableName string) []byte {
query = append(query, "ALTER TABLE "...)
query = m.appendFQN(gen, query, tableName)
return append(query, " "...)
}
switch change := operation.(type) {
case *migrate.CreateTableOp:
return m.AppendCreateTable(b, change.Model)
case *migrate.DropTableOp:
return m.AppendDropTable(b, m.schemaName, change.TableName)
case *migrate.RenameTableOp:
b, err = m.renameTable(gen, appendAlterTable(b, change.TableName), change)
case *migrate.RenameColumnOp:
b, err = m.renameColumn(gen, appendAlterTable(b, change.TableName), change)
case *migrate.AddColumnOp:
b, err = m.addColumn(gen, appendAlterTable(b, change.TableName), change)
case *migrate.DropColumnOp:
b, err = m.dropColumn(gen, appendAlterTable(b, change.TableName), change)
case *migrate.AddPrimaryKeyOp:
b, err = m.addPrimaryKey(gen, appendAlterTable(b, change.TableName), change.PrimaryKey)
case *migrate.ChangePrimaryKeyOp:
b, err = m.changePrimaryKey(gen, appendAlterTable(b, change.TableName), change)
case *migrate.DropPrimaryKeyOp:
b, err = m.dropConstraint(gen, appendAlterTable(b, change.TableName), change.PrimaryKey.Name)
case *migrate.AddUniqueConstraintOp:
b, err = m.addUnique(gen, appendAlterTable(b, change.TableName), change)
case *migrate.DropUniqueConstraintOp:
b, err = m.dropConstraint(gen, appendAlterTable(b, change.TableName), change.Unique.Name)
case *migrate.ChangeColumnTypeOp:
// If column changes to SERIAL, create sequence first.
// https://gist.github.com/oleglomako/185df689706c5499612a0d54d3ffe856
if !change.From.GetIsAutoIncrement() && change.To.GetIsAutoIncrement() {
change.To, b, err = m.createDefaultSequence(gen, b, change)
}
b, err = m.changeColumnType(gen, appendAlterTable(b, change.TableName), change)
case *migrate.AddForeignKeyOp:
b, err = m.addForeignKey(gen, appendAlterTable(b, change.TableName()), change)
case *migrate.DropForeignKeyOp:
b, err = m.dropConstraint(gen, appendAlterTable(b, change.TableName()), change.ConstraintName)
default:
return nil, fmt.Errorf("append sql: unknown operation %T", change)
}
if err != nil {
return nil, fmt.Errorf("append sql: %w", err)
}
return b, nil
}
func (m *migrator) appendFQN(gen schema.QueryGen, b []byte, tableName string) []byte {
return gen.AppendQuery(b, "?.?", bun.Ident(m.schemaName), bun.Ident(tableName))
}
func (m *migrator) renameTable(gen schema.QueryGen, b []byte, rename *migrate.RenameTableOp) (_ []byte, err error) {
b = append(b, "RENAME TO "...)
b = gen.AppendName(b, rename.NewName)
return b, nil
}
func (m *migrator) renameColumn(gen schema.QueryGen, b []byte, rename *migrate.RenameColumnOp) (_ []byte, err error) {
b = append(b, "RENAME COLUMN "...)
b = gen.AppendName(b, rename.OldName)
b = append(b, " TO "...)
b = gen.AppendName(b, rename.NewName)
return b, nil
}
func (m *migrator) addColumn(gen schema.QueryGen, b []byte, add *migrate.AddColumnOp) (_ []byte, err error) {
b = append(b, "ADD COLUMN "...)
b = gen.AppendName(b, add.ColumnName)
b = append(b, " "...)
b, err = add.Column.AppendQuery(gen, b)
if err != nil {
return nil, err
}
if add.Column.GetDefaultValue() != "" {
b = append(b, " DEFAULT "...)
b = append(b, add.Column.GetDefaultValue()...)
b = append(b, " "...)
}
if add.Column.GetIsIdentity() {
b = appendGeneratedAsIdentity(b)
}
return b, nil
}
func (m *migrator) dropColumn(gen schema.QueryGen, b []byte, drop *migrate.DropColumnOp) (_ []byte, err error) {
b = append(b, "DROP COLUMN "...)
b = gen.AppendName(b, drop.ColumnName)
return b, nil
}
func (m *migrator) addPrimaryKey(gen schema.QueryGen, b []byte, pk sqlschema.PrimaryKey) (_ []byte, err error) {
b = append(b, "ADD PRIMARY KEY ("...)
b, _ = pk.Columns.AppendQuery(gen, b)
b = append(b, ")"...)
return b, nil
}
func (m *migrator) changePrimaryKey(gen schema.QueryGen, b []byte, change *migrate.ChangePrimaryKeyOp) (_ []byte, err error) {
b, _ = m.dropConstraint(gen, b, change.Old.Name)
b = append(b, ", "...)
b, _ = m.addPrimaryKey(gen, b, change.New)
return b, nil
}
func (m *migrator) addUnique(gen schema.QueryGen, b []byte, change *migrate.AddUniqueConstraintOp) (_ []byte, err error) {
b = append(b, "ADD CONSTRAINT "...)
if change.Unique.Name != "" {
b = gen.AppendName(b, change.Unique.Name)
} else {
// Default naming scheme for unique constraints in Postgres is <table>_<column>_key
b = gen.AppendName(b, fmt.Sprintf("%s_%s_key", change.TableName, change.Unique.Columns))
}
b = append(b, " UNIQUE ("...)
b, _ = change.Unique.Columns.AppendQuery(gen, b)
b = append(b, ")"...)
return b, nil
}
func (m *migrator) dropConstraint(gen schema.QueryGen, b []byte, name string) (_ []byte, err error) {
b = append(b, "DROP CONSTRAINT "...)
b = gen.AppendName(b, name)
return b, nil
}
func (m *migrator) addForeignKey(gen schema.QueryGen, b []byte, add *migrate.AddForeignKeyOp) (_ []byte, err error) {
b = append(b, "ADD CONSTRAINT "...)
name := add.ConstraintName
if name == "" {
colRef := add.ForeignKey.From
columns := strings.Join(colRef.Column.Split(), "_")
name = fmt.Sprintf("%s_%s_fkey", colRef.TableName, columns)
}
b = gen.AppendName(b, name)
b = append(b, " FOREIGN KEY ("...)
if b, err = add.ForeignKey.From.Column.AppendQuery(gen, b); err != nil {
return b, err
}
b = append(b, ")"...)
b = append(b, " REFERENCES "...)
b = m.appendFQN(gen, b, add.ForeignKey.To.TableName)
b = append(b, " ("...)
if b, err = add.ForeignKey.To.Column.AppendQuery(gen, b); err != nil {
return b, err
}
b = append(b, ")"...)
return b, nil
}
// createDefaultSequence creates a SEQUENCE to back a serial column.
// Having a backing sequence is necessary to change column type to SERIAL.
// The updated Column's default is set to "nextval" of the new sequence.
func (m *migrator) createDefaultSequence(_ schema.QueryGen, b []byte, op *migrate.ChangeColumnTypeOp) (_ sqlschema.Column, _ []byte, err error) {
var last int
if err = m.db.NewSelect().Table(op.TableName).
ColumnExpr("MAX(?)", op.Column).Scan(context.TODO(), &last); err != nil {
return nil, b, err
}
seq := op.TableName + "_" + op.Column + "_seq"
fqn := op.TableName + "." + op.Column
// A sequence that is OWNED BY a table will be dropped
// if the table is dropped with CASCADE action.
b = append(b, "CREATE SEQUENCE "...)
b = append(b, seq...)
b = append(b, " START WITH "...)
b = append(b, fmt.Sprint(last+1)...) // start with next value
b = append(b, " OWNED BY "...)
b = append(b, fqn...)
b = append(b, ";\n"...)
return &Column{
Name: op.To.GetName(),
SQLType: op.To.GetSQLType(),
VarcharLen: op.To.GetVarcharLen(),
DefaultValue: fmt.Sprintf("nextval('%s'::regclass)", seq),
IsNullable: op.To.GetIsNullable(),
IsAutoIncrement: op.To.GetIsAutoIncrement(),
IsIdentity: op.To.GetIsIdentity(),
}, b, nil
}
func (m *migrator) changeColumnType(gen schema.QueryGen, b []byte, colDef *migrate.ChangeColumnTypeOp) (_ []byte, err error) {
// alterColumn never re-assigns err, so there is no need to check for err != nil after calling it
var i int
appendAlterColumn := func() {
if i > 0 {
b = append(b, ", "...)
}
b = append(b, "ALTER COLUMN "...)
b = gen.AppendName(b, colDef.Column)
i++
}
got, want := colDef.From, colDef.To
inspector := m.db.Dialect().(sqlschema.InspectorDialect)
if !inspector.CompareType(want, got) {
appendAlterColumn()
b = append(b, " SET DATA TYPE "...)
if b, err = want.AppendQuery(gen, b); err != nil {
return b, err
}
}
// Column must be declared NOT NULL before identity can be added.
// Although PG can resolve the order of operations itself, we make this explicit in the query.
if want.GetIsNullable() != got.GetIsNullable() {
appendAlterColumn()
if !want.GetIsNullable() {
b = append(b, " SET NOT NULL"...)
} else {
b = append(b, " DROP NOT NULL"...)
}
}
if want.GetIsIdentity() != got.GetIsIdentity() {
appendAlterColumn()
if !want.GetIsIdentity() {
b = append(b, " DROP IDENTITY"...)
} else {
b = append(b, " ADD"...)
b = appendGeneratedAsIdentity(b)
}
}
if want.GetDefaultValue() != got.GetDefaultValue() {
appendAlterColumn()
if want.GetDefaultValue() == "" {
b = append(b, " DROP DEFAULT"...)
} else {
b = append(b, " SET DEFAULT "...)
b = append(b, want.GetDefaultValue()...)
}
}
return b, nil
}
+87
View File
@@ -0,0 +1,87 @@
package pgdialect
import (
"database/sql/driver"
"fmt"
"reflect"
"time"
"github.com/uptrace/bun/dialect"
"github.com/uptrace/bun/schema"
)
var (
driverValuerType = reflect.TypeFor[driver.Valuer]()
stringType = reflect.TypeFor[string]()
sliceStringType = reflect.TypeFor[[]string]()
intType = reflect.TypeFor[int]()
sliceIntType = reflect.TypeFor[[]int]()
int64Type = reflect.TypeFor[int64]()
sliceInt64Type = reflect.TypeFor[[]int64]()
float64Type = reflect.TypeFor[float64]()
sliceFloat64Type = reflect.TypeFor[[]float64]()
timeType = reflect.TypeFor[time.Time]()
sliceTimeType = reflect.TypeFor[[]time.Time]()
)
func appendTime(buf []byte, tm time.Time) []byte {
return tm.UTC().AppendFormat(buf, "2006-01-02 15:04:05.999999-07:00")
}
var mapStringStringType = reflect.TypeOf(map[string]string(nil))
func (d *Dialect) hstoreAppender(typ reflect.Type) schema.AppenderFunc {
kind := typ.Kind()
switch kind {
case reflect.Ptr:
if fn := d.hstoreAppender(typ.Elem()); fn != nil {
return schema.PtrAppender(fn)
}
case reflect.Map:
// ok:
default:
return nil
}
if typ.Key() == stringType && typ.Elem() == stringType {
return appendMapStringStringValue
}
return func(gen schema.QueryGen, b []byte, v reflect.Value) []byte {
err := fmt.Errorf("bun: Hstore(unsupported %s)", v.Type())
return dialect.AppendError(b, err)
}
}
func appendMapStringString(b []byte, m map[string]string) []byte {
if m == nil {
return dialect.AppendNull(b)
}
b = append(b, '\'')
for key, value := range m {
b = appendStringElem(b, key)
b = append(b, '=', '>')
b = appendStringElem(b, value)
b = append(b, ',')
}
if len(m) > 0 {
b = b[:len(b)-1] // Strip trailing comma.
}
b = append(b, '\'')
return b
}
func appendMapStringStringValue(gen schema.QueryGen, b []byte, v reflect.Value) []byte {
m := v.Convert(mapStringStringType).Interface().(map[string]string)
return appendMapStringString(b, m)
}
+608
View File
@@ -0,0 +1,608 @@
package pgdialect
import (
"database/sql"
"database/sql/driver"
"fmt"
"math"
"reflect"
"strconv"
"time"
"github.com/uptrace/bun/dialect"
"github.com/uptrace/bun/internal"
"github.com/uptrace/bun/schema"
)
type ArrayValue struct {
v reflect.Value
append schema.AppenderFunc
scan schema.ScannerFunc
}
// Array accepts a slice and returns a wrapper for working with PostgreSQL
// array data type.
//
// For struct fields you can use array tag:
//
// Emails []string `bun:",array"`
func Array(vi any) *ArrayValue {
v := reflect.ValueOf(vi)
if !v.IsValid() {
panic(fmt.Errorf("bun: Array(nil)"))
}
return &ArrayValue{
v: v,
append: pgDialect.arrayAppender(v.Type()),
scan: arrayScanner(v.Type()),
}
}
var (
_ schema.QueryAppender = (*ArrayValue)(nil)
_ sql.Scanner = (*ArrayValue)(nil)
)
func (a *ArrayValue) AppendQuery(gen schema.QueryGen, b []byte) ([]byte, error) {
if a.append == nil {
panic(fmt.Errorf("bun: Array(unsupported %s)", a.v.Type()))
}
return a.append(gen, b, a.v), nil
}
func (a *ArrayValue) Scan(src any) error {
if a.scan == nil {
return fmt.Errorf("bun: Array(unsupported %s)", a.v.Type())
}
if a.v.Kind() != reflect.Ptr {
return fmt.Errorf("bun: Array(non-pointer %s)", a.v.Type())
}
return a.scan(a.v, src)
}
func (a *ArrayValue) Value() any {
if a.v.IsValid() {
return a.v.Interface()
}
return nil
}
//------------------------------------------------------------------------------
func (d *Dialect) arrayAppender(typ reflect.Type) schema.AppenderFunc {
kind := typ.Kind()
switch kind {
case reflect.Ptr:
if fn := d.arrayAppender(typ.Elem()); fn != nil {
return schema.PtrAppender(fn)
}
case reflect.Slice, reflect.Array:
// continue below
default:
return nil
}
elemType := typ.Elem()
if kind == reflect.Slice {
switch elemType {
case stringType:
return appendStringSliceValue
case intType:
return appendIntSliceValue
case int64Type:
return appendInt64SliceValue
case float64Type:
return appendFloat64SliceValue
case timeType:
return appendTimeSliceValue
}
}
appendElem := d.arrayElemAppender(elemType)
if appendElem == nil {
panic(fmt.Errorf("pgdialect: %s is not supported", typ))
}
return func(gen schema.QueryGen, b []byte, v reflect.Value) []byte {
kind := v.Kind()
switch kind {
case reflect.Ptr, reflect.Slice:
if v.IsNil() {
return dialect.AppendNull(b)
}
}
if kind == reflect.Ptr {
v = v.Elem()
}
b = append(b, "'{"...)
ln := v.Len()
for i := 0; i < ln; i++ {
elem := v.Index(i)
if i > 0 {
b = append(b, ',')
}
b = appendElem(gen, b, elem)
}
b = append(b, "}'"...)
return b
}
}
func (d *Dialect) arrayElemAppender(typ reflect.Type) schema.AppenderFunc {
if typ.Implements(driverValuerType) {
return arrayAppendDriverValue
}
if typ == timeType {
return appendTimeElemValue
}
switch typ.Kind() {
case reflect.String:
return appendStringElemValue
case reflect.Slice:
if typ.Elem().Kind() == reflect.Uint8 {
return appendBytesElemValue
}
case reflect.Ptr:
return schema.PtrAppender(d.arrayElemAppender(typ.Elem()))
}
return schema.Appender(d, typ)
}
func appendTimeElemValue(gen schema.QueryGen, b []byte, v reflect.Value) []byte {
ts := v.Convert(timeType).Interface().(time.Time)
b = append(b, '"')
b = appendTime(b, ts)
return append(b, '"')
}
func appendStringElemValue(gen schema.QueryGen, b []byte, v reflect.Value) []byte {
return appendStringElem(b, v.String())
}
func appendBytesElemValue(gen schema.QueryGen, b []byte, v reflect.Value) []byte {
return appendBytesElem(b, v.Bytes())
}
func arrayAppendDriverValue(gen schema.QueryGen, b []byte, v reflect.Value) []byte {
iface, err := v.Interface().(driver.Valuer).Value()
if err != nil {
return dialect.AppendError(b, err)
}
return appendElem(b, iface)
}
func appendStringSliceValue(gen schema.QueryGen, b []byte, v reflect.Value) []byte {
ss := v.Convert(sliceStringType).Interface().([]string)
return appendStringSlice(b, ss)
}
func appendStringSlice(b []byte, ss []string) []byte {
if ss == nil {
return dialect.AppendNull(b)
}
b = append(b, '\'')
b = append(b, '{')
for _, s := range ss {
b = appendStringElem(b, s)
b = append(b, ',')
}
if len(ss) > 0 {
b[len(b)-1] = '}' // Replace trailing comma.
} else {
b = append(b, '}')
}
b = append(b, '\'')
return b
}
func appendIntSliceValue(gen schema.QueryGen, b []byte, v reflect.Value) []byte {
ints := v.Convert(sliceIntType).Interface().([]int)
return appendIntSlice(b, ints)
}
func appendIntSlice(b []byte, ints []int) []byte {
if ints == nil {
return dialect.AppendNull(b)
}
b = append(b, '\'')
b = append(b, '{')
for _, n := range ints {
b = strconv.AppendInt(b, int64(n), 10)
b = append(b, ',')
}
if len(ints) > 0 {
b[len(b)-1] = '}' // Replace trailing comma.
} else {
b = append(b, '}')
}
b = append(b, '\'')
return b
}
func appendInt64SliceValue(gen schema.QueryGen, b []byte, v reflect.Value) []byte {
ints := v.Convert(sliceInt64Type).Interface().([]int64)
return appendInt64Slice(b, ints)
}
func appendInt64Slice(b []byte, ints []int64) []byte {
if ints == nil {
return dialect.AppendNull(b)
}
b = append(b, '\'')
b = append(b, '{')
for _, n := range ints {
b = strconv.AppendInt(b, n, 10)
b = append(b, ',')
}
if len(ints) > 0 {
b[len(b)-1] = '}' // Replace trailing comma.
} else {
b = append(b, '}')
}
b = append(b, '\'')
return b
}
func appendFloat64SliceValue(gen schema.QueryGen, b []byte, v reflect.Value) []byte {
floats := v.Convert(sliceFloat64Type).Interface().([]float64)
return appendFloat64Slice(b, floats)
}
func appendFloat64Slice(b []byte, floats []float64) []byte {
if floats == nil {
return dialect.AppendNull(b)
}
b = append(b, '\'')
b = append(b, '{')
for _, n := range floats {
b = arrayAppendFloat64(b, n)
b = append(b, ',')
}
if len(floats) > 0 {
b[len(b)-1] = '}' // Replace trailing comma.
} else {
b = append(b, '}')
}
b = append(b, '\'')
return b
}
func arrayAppendFloat64(b []byte, num float64) []byte {
switch {
case math.IsNaN(num):
return append(b, "NaN"...)
case math.IsInf(num, 1):
return append(b, "Infinity"...)
case math.IsInf(num, -1):
return append(b, "-Infinity"...)
default:
return strconv.AppendFloat(b, num, 'f', -1, 64)
}
}
func appendTimeSliceValue(gen schema.QueryGen, b []byte, v reflect.Value) []byte {
ts := v.Convert(sliceTimeType).Interface().([]time.Time)
return appendTimeSlice(gen, b, ts)
}
func appendTimeSlice(gen schema.QueryGen, b []byte, ts []time.Time) []byte {
if ts == nil {
return dialect.AppendNull(b)
}
b = append(b, '\'')
b = append(b, '{')
for _, t := range ts {
b = append(b, '"')
b = appendTime(b, t)
b = append(b, '"')
b = append(b, ',')
}
if len(ts) > 0 {
b[len(b)-1] = '}' // Replace trailing comma.
} else {
b = append(b, '}')
}
b = append(b, '\'')
return b
}
//------------------------------------------------------------------------------
func arrayScanner(typ reflect.Type) schema.ScannerFunc {
kind := typ.Kind()
switch kind {
case reflect.Ptr:
if fn := arrayScanner(typ.Elem()); fn != nil {
return schema.PtrScanner(fn)
}
case reflect.Slice, reflect.Array:
// ok:
default:
return nil
}
elemType := typ.Elem()
if kind == reflect.Slice {
switch elemType {
case stringType:
return scanStringSliceValue
case intType:
return scanIntSliceValue
case int64Type:
return scanInt64SliceValue
case float64Type:
return scanFloat64SliceValue
}
}
scanElem := schema.Scanner(elemType)
return func(dest reflect.Value, src any) error {
dest = reflect.Indirect(dest)
if !dest.CanSet() {
return fmt.Errorf("bun: Scan(non-settable %s)", dest.Type())
}
kind := dest.Kind()
if src == nil {
if kind != reflect.Slice || !dest.IsNil() {
dest.Set(reflect.Zero(dest.Type()))
}
return nil
}
if kind == reflect.Slice {
if dest.IsNil() {
dest.Set(reflect.MakeSlice(dest.Type(), 0, 0))
} else if dest.Len() > 0 {
dest.Set(dest.Slice(0, 0))
}
}
if src == nil {
return nil
}
b, err := toBytes(src)
if err != nil {
return err
}
p := newArrayParser(b)
nextValue := internal.MakeSliceNextElemFunc(dest)
for p.Next() {
elem := p.Elem()
elemValue := nextValue()
if err := scanElem(elemValue, elem); err != nil {
return fmt.Errorf("scanElem failed: %w", err)
}
}
return p.Err()
}
}
func scanStringSliceValue(dest reflect.Value, src any) error {
dest = reflect.Indirect(dest)
if !dest.CanSet() {
return fmt.Errorf("bun: Scan(non-settable %s)", dest.Type())
}
slice, err := decodeStringSlice(src)
if err != nil {
return err
}
dest.Set(reflect.ValueOf(slice))
return nil
}
func decodeStringSlice(src any) ([]string, error) {
if src == nil {
return nil, nil
}
b, err := toBytes(src)
if err != nil {
return nil, err
}
slice := make([]string, 0)
p := newArrayParser(b)
for p.Next() {
elem := p.Elem()
slice = append(slice, string(elem))
}
if err := p.Err(); err != nil {
return nil, err
}
return slice, nil
}
func scanIntSliceValue(dest reflect.Value, src any) error {
dest = reflect.Indirect(dest)
if !dest.CanSet() {
return fmt.Errorf("bun: Scan(non-settable %s)", dest.Type())
}
slice, err := decodeIntSlice(src)
if err != nil {
return err
}
dest.Set(reflect.ValueOf(slice))
return nil
}
func decodeIntSlice(src any) ([]int, error) {
if src == nil {
return nil, nil
}
b, err := toBytes(src)
if err != nil {
return nil, err
}
slice := make([]int, 0)
p := newArrayParser(b)
for p.Next() {
elem := p.Elem()
if elem == nil {
slice = append(slice, 0)
continue
}
n, err := strconv.Atoi(internal.String(elem))
if err != nil {
return nil, err
}
slice = append(slice, n)
}
if err := p.Err(); err != nil {
return nil, err
}
return slice, nil
}
func scanInt64SliceValue(dest reflect.Value, src any) error {
dest = reflect.Indirect(dest)
if !dest.CanSet() {
return fmt.Errorf("bun: Scan(non-settable %s)", dest.Type())
}
slice, err := decodeInt64Slice(src)
if err != nil {
return err
}
dest.Set(reflect.ValueOf(slice))
return nil
}
func decodeInt64Slice(src any) ([]int64, error) {
if src == nil {
return nil, nil
}
b, err := toBytes(src)
if err != nil {
return nil, err
}
slice := make([]int64, 0)
p := newArrayParser(b)
for p.Next() {
elem := p.Elem()
if elem == nil {
slice = append(slice, 0)
continue
}
n, err := strconv.ParseInt(internal.String(elem), 10, 64)
if err != nil {
return nil, err
}
slice = append(slice, n)
}
if err := p.Err(); err != nil {
return nil, err
}
return slice, nil
}
func scanFloat64SliceValue(dest reflect.Value, src any) error {
dest = reflect.Indirect(dest)
if !dest.CanSet() {
return fmt.Errorf("bun: Scan(non-settable %s)", dest.Type())
}
slice, err := scanFloat64Slice(src)
if err != nil {
return err
}
dest.Set(reflect.ValueOf(slice))
return nil
}
func scanFloat64Slice(src any) ([]float64, error) {
if src == nil {
return nil, nil
}
b, err := toBytes(src)
if err != nil {
return nil, err
}
slice := make([]float64, 0)
p := newArrayParser(b)
for p.Next() {
elem := p.Elem()
if elem == nil {
slice = append(slice, 0)
continue
}
n, err := strconv.ParseFloat(internal.String(elem), 64)
if err != nil {
return nil, err
}
slice = append(slice, n)
}
if err := p.Err(); err != nil {
return nil, err
}
return slice, nil
}
func toBytes(src any) ([]byte, error) {
switch src := src.(type) {
case string:
return internal.Bytes(src), nil
case []byte:
return src, nil
default:
return nil, fmt.Errorf("pgdialect: got %T, wanted []byte or string", src)
}
}
+119
View File
@@ -0,0 +1,119 @@
package pgdialect
import (
"bytes"
"fmt"
"io"
)
type arrayParser struct {
p pgparser
elem []byte
err error
isJson bool
}
func newArrayParser(b []byte) *arrayParser {
p := new(arrayParser)
if b[0] == 'n' {
p.p.Reset(nil)
return p
}
if len(b) < 2 || (b[0] != '{' && b[0] != '[') || (b[len(b)-1] != '}' && b[len(b)-1] != ']') {
p.err = fmt.Errorf("pgdialect: can't parse array: %q", b)
return p
}
p.isJson = b[0] == '['
p.p.Reset(b[1 : len(b)-1])
return p
}
func (p *arrayParser) Next() bool {
if p.err != nil {
return false
}
p.err = p.readNext()
return p.err == nil
}
func (p *arrayParser) Err() error {
if p.err != io.EOF {
return p.err
}
return nil
}
func (p *arrayParser) Elem() []byte {
return p.elem
}
func (p *arrayParser) readNext() error {
ch := p.p.Read()
if ch == 0 {
return io.EOF
}
switch ch {
case '}', ']':
return io.EOF
case '"':
b, err := p.p.ReadSubstring(ch)
if err != nil {
return err
}
if p.p.Peek() == ',' {
p.p.Advance()
}
p.elem = b
return nil
case '[', '(':
rng, err := p.p.ReadRange(ch)
if err != nil {
return err
}
if p.p.Peek() == ',' {
p.p.Advance()
}
p.elem = rng
return nil
default:
if ch == '{' && p.isJson {
json, err := p.p.ReadJSON()
if err != nil {
return err
}
for {
if p.p.Peek() == ',' || p.p.Peek() == ' ' {
p.p.Advance()
} else {
break
}
}
p.elem = json
return nil
} else {
lit := p.p.ReadLiteral(ch)
if bytes.Equal(lit, []byte("NULL")) {
lit = nil
}
if p.p.Peek() == ',' {
p.p.Advance()
}
p.elem = lit
return nil
}
}
}
+163
View File
@@ -0,0 +1,163 @@
package pgdialect
import (
"database/sql"
"fmt"
"strconv"
"strings"
"github.com/uptrace/bun"
"github.com/uptrace/bun/dialect"
"github.com/uptrace/bun/dialect/feature"
"github.com/uptrace/bun/dialect/sqltype"
"github.com/uptrace/bun/migrate/sqlschema"
"github.com/uptrace/bun/schema"
)
var pgDialect = New()
func init() {
if Version() != bun.Version() {
panic(fmt.Errorf("pgdialect and Bun must have the same version: v%s != v%s",
Version(), bun.Version()))
}
}
type Dialect struct {
schema.BaseDialect
tables *schema.Tables
features feature.Feature
uintAsInt bool
}
var _ schema.Dialect = (*Dialect)(nil)
var _ sqlschema.InspectorDialect = (*Dialect)(nil)
var _ sqlschema.MigratorDialect = (*Dialect)(nil)
func New(opts ...DialectOption) *Dialect {
d := new(Dialect)
d.tables = schema.NewTables(d)
d.features = feature.CTE |
feature.WithValues |
feature.Returning |
feature.InsertReturning |
feature.DefaultPlaceholder |
feature.DoubleColonCast |
feature.InsertTableAlias |
feature.UpdateTableAlias |
feature.DeleteTableAlias |
feature.TableCascade |
feature.TableIdentity |
feature.TableTruncate |
feature.TableNotExists |
feature.InsertOnConflict |
feature.SelectExists |
feature.GeneratedIdentity |
feature.CompositeIn |
feature.FKDefaultOnAction |
feature.DeleteReturning |
feature.Merge |
feature.MergeReturning |
feature.AlterColumnExists |
feature.CreateIndexIfNotExists
for _, opt := range opts {
opt(d)
}
return d
}
type DialectOption func(d *Dialect)
func WithoutFeature(other feature.Feature) DialectOption {
return func(d *Dialect) {
d.features = d.features.Remove(other)
}
}
func WithAppendUintAsInt(on bool) DialectOption {
return func(d *Dialect) {
d.uintAsInt = on
}
}
func (d *Dialect) Init(*sql.DB) {}
func (d *Dialect) Name() dialect.Name {
return dialect.PG
}
func (d *Dialect) Features() feature.Feature {
return d.features
}
func (d *Dialect) Tables() *schema.Tables {
return d.tables
}
func (d *Dialect) OnTable(table *schema.Table) {
for _, field := range table.FieldMap {
d.onField(field)
}
}
func (d *Dialect) onField(field *schema.Field) {
field.DiscoveredSQLType = fieldSQLType(field)
if field.AutoIncrement && !field.Identity {
switch field.DiscoveredSQLType {
case sqltype.SmallInt:
field.CreateTableSQLType = pgTypeSmallSerial
case sqltype.Integer:
field.CreateTableSQLType = pgTypeSerial
case sqltype.BigInt:
field.CreateTableSQLType = pgTypeBigSerial
}
}
if field.Tag.HasOption("array") || strings.HasSuffix(field.UserSQLType, "[]") {
field.Append = d.arrayAppender(field.StructField.Type)
field.Scan = arrayScanner(field.StructField.Type)
return
}
if field.Tag.HasOption("multirange") || strings.HasSuffix(field.UserSQLType, "multirange") {
field.Scan = arrayScanner(field.StructField.Type)
return
}
switch field.DiscoveredSQLType {
case sqltype.HSTORE:
field.Append = d.hstoreAppender(field.StructField.Type)
field.Scan = hstoreScanner(field.StructField.Type)
}
}
func (d *Dialect) IdentQuote() byte {
return '"'
}
func (d *Dialect) AppendUint32(b []byte, n uint32) []byte {
if d.uintAsInt {
return strconv.AppendInt(b, int64(int32(n)), 10)
}
return strconv.AppendUint(b, uint64(n), 10)
}
func (d *Dialect) AppendUint64(b []byte, n uint64) []byte {
if d.uintAsInt {
return strconv.AppendInt(b, int64(n), 10)
}
return strconv.AppendUint(b, n, 10)
}
func (d *Dialect) AppendSequence(b []byte, _ *schema.Table, _ *schema.Field) []byte {
return appendGeneratedAsIdentity(b)
}
// appendGeneratedAsIdentity appends GENERATED BY DEFAULT AS IDENTITY to the column definition.
func appendGeneratedAsIdentity(b []byte) []byte {
return append(b, " GENERATED BY DEFAULT AS IDENTITY"...)
}
+87
View File
@@ -0,0 +1,87 @@
package pgdialect
import (
"database/sql/driver"
"encoding/hex"
"fmt"
"strconv"
"time"
"unicode/utf8"
"github.com/uptrace/bun/dialect"
)
func appendElem(buf []byte, val any) []byte {
switch val := val.(type) {
case int64:
return strconv.AppendInt(buf, val, 10)
case float64:
return arrayAppendFloat64(buf, val)
case bool:
return dialect.AppendBool(buf, val)
case []byte:
return appendBytesElem(buf, val)
case string:
return appendStringElem(buf, val)
case time.Time:
buf = append(buf, '"')
buf = appendTime(buf, val)
buf = append(buf, '"')
return buf
case driver.Valuer:
val2, err := val.Value()
if err != nil {
err := fmt.Errorf("pgdialect: can't append elem value: %w", err)
return dialect.AppendError(buf, err)
}
return appendElem(buf, val2)
default:
err := fmt.Errorf("pgdialect: can't append elem %T", val)
return dialect.AppendError(buf, err)
}
}
func appendBytesElem(b []byte, bs []byte) []byte {
if bs == nil {
return dialect.AppendNull(b)
}
b = append(b, `"\\x`...)
s := len(b)
b = append(b, make([]byte, hex.EncodedLen(len(bs)))...)
hex.Encode(b[s:], bs)
b = append(b, '"')
return b
}
func appendStringElem(b []byte, s string) []byte {
b = append(b, '"')
for _, r := range s {
switch r {
case 0:
// ignore
case '\'':
b = append(b, "''"...)
case '"':
b = append(b, '\\', '"')
case '\\':
b = append(b, '\\', '\\')
default:
if r < utf8.RuneSelf {
b = append(b, byte(r))
break
}
l := len(b)
if cap(b)-l < utf8.UTFMax {
b = append(b, make([]byte, utf8.UTFMax)...)
}
n := utf8.EncodeRune(b[l:l+utf8.UTFMax], r)
b = b[:l+n]
}
}
b = append(b, '"')
return b
}
+73
View File
@@ -0,0 +1,73 @@
package pgdialect
import (
"database/sql"
"fmt"
"reflect"
"github.com/uptrace/bun/schema"
)
type HStoreValue struct {
v reflect.Value
append schema.AppenderFunc
scan schema.ScannerFunc
}
// HStore accepts a map[string]string and returns a wrapper for working with PostgreSQL
// hstore data type.
//
// For struct fields you can use hstore tag:
//
// Attrs map[string]string `bun:",hstore"`
func HStore(vi any) *HStoreValue {
v := reflect.ValueOf(vi)
if !v.IsValid() {
panic(fmt.Errorf("bun: HStore(nil)"))
}
typ := v.Type()
if typ.Kind() == reflect.Ptr {
typ = typ.Elem()
}
if typ.Kind() != reflect.Map {
panic(fmt.Errorf("bun: Hstore(unsupported %s)", typ))
}
return &HStoreValue{
v: v,
append: pgDialect.hstoreAppender(v.Type()),
scan: hstoreScanner(v.Type()),
}
}
var (
_ schema.QueryAppender = (*HStoreValue)(nil)
_ sql.Scanner = (*HStoreValue)(nil)
)
func (h *HStoreValue) AppendQuery(gen schema.QueryGen, b []byte) ([]byte, error) {
if h.append == nil {
panic(fmt.Errorf("bun: HStore(unsupported %s)", h.v.Type()))
}
return h.append(gen, b, h.v), nil
}
func (h *HStoreValue) Scan(src any) error {
if h.scan == nil {
return fmt.Errorf("bun: HStore(unsupported %s)", h.v.Type())
}
if h.v.Kind() != reflect.Ptr {
return fmt.Errorf("bun: HStore(non-pointer %s)", h.v.Type())
}
return h.scan(h.v.Elem(), src)
}
func (h *HStoreValue) Value() any {
if h.v.IsValid() {
return h.v.Interface()
}
return nil
}
+100
View File
@@ -0,0 +1,100 @@
package pgdialect
import (
"bytes"
"fmt"
"io"
)
type hstoreParser struct {
p pgparser
key string
value string
err error
}
func newHStoreParser(b []byte) *hstoreParser {
p := new(hstoreParser)
if len(b) != 0 && (len(b) < 6 || b[0] != '"') {
p.err = fmt.Errorf("pgdialect: can't parse hstore: %q", b)
return p
}
p.p.Reset(b)
return p
}
func (p *hstoreParser) Next() bool {
if p.err != nil {
return false
}
p.err = p.readNext()
return p.err == nil
}
func (p *hstoreParser) Err() error {
if p.err != io.EOF {
return p.err
}
return nil
}
func (p *hstoreParser) Key() string {
return p.key
}
func (p *hstoreParser) Value() string {
return p.value
}
func (p *hstoreParser) readNext() error {
if !p.p.Valid() {
return io.EOF
}
if err := p.p.Skip('"'); err != nil {
return err
}
key, err := p.p.ReadUnescapedSubstring('"')
if err != nil {
return err
}
p.key = string(key)
if err := p.p.SkipPrefix([]byte("=>")); err != nil {
return err
}
ch, err := p.p.ReadByte()
if err != nil {
return err
}
switch ch {
case '"':
value, err := p.p.ReadUnescapedSubstring(ch)
if err != nil {
return err
}
p.skipComma()
p.value = string(value)
return nil
default:
value := p.p.ReadLiteral(ch)
if bytes.Equal(value, []byte("NULL")) {
p.value = ""
}
p.skipComma()
return nil
}
}
func (p *hstoreParser) skipComma() {
if p.p.Peek() == ',' {
p.p.Advance()
}
if p.p.Peek() == ' ' {
p.p.Advance()
}
}
+67
View File
@@ -0,0 +1,67 @@
package pgdialect
import (
"fmt"
"reflect"
"github.com/uptrace/bun/schema"
)
func hstoreScanner(typ reflect.Type) schema.ScannerFunc {
kind := typ.Kind()
switch kind {
case reflect.Ptr:
if fn := hstoreScanner(typ.Elem()); fn != nil {
return schema.PtrScanner(fn)
}
case reflect.Map:
// ok:
default:
return nil
}
if typ.Key() == stringType && typ.Elem() == stringType {
return scanMapStringStringValue
}
return func(dest reflect.Value, src any) error {
return fmt.Errorf("bun: Hstore(unsupported %s)", dest.Type())
}
}
func scanMapStringStringValue(dest reflect.Value, src any) error {
dest = reflect.Indirect(dest)
if !dest.CanSet() {
return fmt.Errorf("bun: Scan(non-settable %s)", dest.Type())
}
m, err := decodeMapStringString(src)
if err != nil {
return err
}
dest.Set(reflect.ValueOf(m))
return nil
}
func decodeMapStringString(src any) (map[string]string, error) {
if src == nil {
return nil, nil
}
b, err := toBytes(src)
if err != nil {
return nil, err
}
m := make(map[string]string)
p := newHStoreParser(b)
for p.Next() {
m[p.Key()] = p.Value()
}
if err := p.Err(); err != nil {
return nil, err
}
return m, nil
}
+300
View File
@@ -0,0 +1,300 @@
package pgdialect
import (
"context"
"strings"
"github.com/uptrace/bun"
"github.com/uptrace/bun/migrate/sqlschema"
)
type (
Schema = sqlschema.BaseDatabase
Table = sqlschema.BaseTable
Column = sqlschema.BaseColumn
)
func (d *Dialect) NewInspector(db *bun.DB, options ...sqlschema.InspectorOption) sqlschema.Inspector {
return newInspector(db, options...)
}
type Inspector struct {
sqlschema.InspectorConfig
db *bun.DB
}
var _ sqlschema.Inspector = (*Inspector)(nil)
func newInspector(db *bun.DB, options ...sqlschema.InspectorOption) *Inspector {
i := &Inspector{db: db}
sqlschema.ApplyInspectorOptions(&i.InspectorConfig, options...)
return i
}
func (in *Inspector) Inspect(ctx context.Context) (sqlschema.Database, error) {
dbSchema := Schema{
ForeignKeys: make(map[sqlschema.ForeignKey]string),
}
exclude := in.ExcludeTables
if len(exclude) == 0 {
// Avoid getting NOT LIKE ALL (ARRAY[NULL]) if bun.List() is called with an empty slice.
exclude = []string{""}
}
var tables []*InformationSchemaTable
if err := in.db.NewRaw(sqlInspectTables, in.SchemaName, bun.List(exclude)).Scan(ctx, &tables); err != nil {
return dbSchema, err
}
var fks []*ForeignKey
if err := in.db.NewRaw(sqlInspectForeignKeys, in.SchemaName, bun.List(exclude), bun.List(exclude)).Scan(ctx, &fks); err != nil {
return dbSchema, err
}
dbSchema.ForeignKeys = make(map[sqlschema.ForeignKey]string, len(fks))
for _, table := range tables {
var columns []*InformationSchemaColumn
if err := in.db.NewRaw(sqlInspectColumnsQuery, table.Schema, table.Name).Scan(ctx, &columns); err != nil {
return dbSchema, err
}
var colDefs []sqlschema.Column
uniqueGroups := make(map[string][]string)
for _, c := range columns {
def := c.Default
if c.IsSerial || c.IsIdentity {
def = ""
} else if !c.IsDefaultLiteral {
def = strings.ToLower(def)
}
colDefs = append(colDefs, &Column{
Name: c.Name,
SQLType: c.DataType,
VarcharLen: c.VarcharLen,
DefaultValue: def,
IsNullable: c.IsNullable,
IsAutoIncrement: c.IsSerial,
IsIdentity: c.IsIdentity,
})
for _, group := range c.UniqueGroups {
uniqueGroups[group] = append(uniqueGroups[group], c.Name)
}
}
var unique []sqlschema.Unique
for name, columns := range uniqueGroups {
unique = append(unique, sqlschema.Unique{
Name: name,
Columns: sqlschema.NewColumns(columns...),
})
}
var pk *sqlschema.PrimaryKey
if len(table.PrimaryKey.Columns) > 0 {
pk = &sqlschema.PrimaryKey{
Name: table.PrimaryKey.ConstraintName,
Columns: sqlschema.NewColumns(table.PrimaryKey.Columns...),
}
}
dbSchema.Tables = append(dbSchema.Tables, &Table{
Schema: table.Schema,
Name: table.Name,
Columns: colDefs,
PrimaryKey: pk,
UniqueConstraints: unique,
})
}
for _, fk := range fks {
dbFK := sqlschema.ForeignKey{
From: sqlschema.NewColumnReference(fk.SourceTable, fk.SourceColumns...),
To: sqlschema.NewColumnReference(fk.TargetTable, fk.TargetColumns...),
}
if _, exclude := in.ExcludeForeignKeys[dbFK]; exclude {
continue
}
dbSchema.ForeignKeys[dbFK] = fk.ConstraintName
}
return dbSchema, nil
}
type InformationSchemaTable struct {
Schema string `bun:"table_schema,pk"`
Name string `bun:"table_name,pk"`
PrimaryKey PrimaryKey `bun:"embed:primary_key_"`
Columns []*InformationSchemaColumn `bun:"rel:has-many,join:table_schema=table_schema,join:table_name=table_name"`
}
type InformationSchemaColumn struct {
Schema string `bun:"table_schema"`
Table string `bun:"table_name"`
Name string `bun:"column_name"`
DataType string `bun:"data_type"`
VarcharLen int `bun:"varchar_len"`
IsArray bool `bun:"is_array"`
ArrayDims int `bun:"array_dims"`
Default string `bun:"default"`
IsDefaultLiteral bool `bun:"default_is_literal_expr"`
IsIdentity bool `bun:"is_identity"`
IndentityType string `bun:"identity_type"`
IsSerial bool `bun:"is_serial"`
IsNullable bool `bun:"is_nullable"`
UniqueGroups []string `bun:"unique_groups,array"`
}
type ForeignKey struct {
ConstraintName string `bun:"constraint_name"`
SourceSchema string `bun:"schema_name"`
SourceTable string `bun:"table_name"`
SourceColumns []string `bun:"columns,array"`
TargetSchema string `bun:"target_schema"`
TargetTable string `bun:"target_table"`
TargetColumns []string `bun:"target_columns,array"`
}
type PrimaryKey struct {
ConstraintName string `bun:"name"`
Columns []string `bun:"columns,array"`
}
const (
// sqlInspectTables retrieves all user-defined tables in the selected schema.
// Pass bun.List([]string{...}) to exclude tables from this inspection or bun.List([]string{''}) to include all results.
sqlInspectTables = `
SELECT
"t".table_schema,
"t".table_name,
pk.name AS primary_key_name,
pk.columns AS primary_key_columns
FROM information_schema.tables "t"
LEFT JOIN (
SELECT i.indrelid, "idx".relname AS "name", ARRAY_AGG("a".attname) AS "columns"
FROM pg_index i
JOIN pg_attribute "a"
ON "a".attrelid = i.indrelid
AND "a".attnum = ANY("i".indkey)
AND i.indisprimary
JOIN pg_class "idx" ON i.indexrelid = "idx".oid
GROUP BY 1, 2
) pk
ON ("t".table_schema || '.' || "t".table_name)::regclass = pk.indrelid
WHERE table_type = 'BASE TABLE'
AND "t".table_schema = ?
AND "t".table_schema NOT LIKE 'pg_%'
AND "table_name" NOT LIKE ALL (ARRAY[?])
ORDER BY "t".table_schema, "t".table_name
`
// sqlInspectColumnsQuery retrieves column definitions for the specified table.
// Unlike sqlInspectTables and sqlInspectSchema, it should be passed to bun.NewRaw
// with additional args for table_schema and table_name.
sqlInspectColumnsQuery = `
SELECT
"c".table_schema,
"c".table_name,
"c".column_name,
"c".data_type,
"c".character_maximum_length::integer AS varchar_len,
"c".data_type = 'ARRAY' AS is_array,
COALESCE("c".array_dims, 0) AS array_dims,
CASE
WHEN "c".column_default ~ '^''.*''::.*$' THEN substring("c".column_default FROM '^''(.*)''::.*$')
ELSE "c".column_default
END AS "default",
"c".column_default ~ '^''.*''::.*$' OR "c".column_default ~ '^[0-9\.]+$' AS default_is_literal_expr,
"c".is_identity = 'YES' AS is_identity,
"c".column_default = format('nextval(''%s_%s_seq''::regclass)', "c".table_name, "c".column_name) AS is_serial,
COALESCE("c".identity_type, '') AS identity_type,
"c".is_nullable = 'YES' AS is_nullable,
"c"."unique_groups" AS unique_groups
FROM (
SELECT
"table_schema",
"table_name",
"column_name",
"c".data_type,
"c".character_maximum_length,
"c".column_default,
"c".is_identity,
"c".is_nullable,
att.array_dims,
att.identity_type,
att."unique_groups",
att."constraint_type"
FROM information_schema.columns "c"
LEFT JOIN (
SELECT
s.nspname AS "table_schema",
"t".relname AS "table_name",
"c".attname AS "column_name",
"c".attndims AS array_dims,
"c".attidentity AS identity_type,
ARRAY_AGG(con.conname) FILTER (WHERE con.contype = 'u') AS "unique_groups",
ARRAY_AGG(con.contype) AS "constraint_type"
FROM (
SELECT
conname,
contype,
connamespace,
conrelid,
conrelid AS attrelid,
UNNEST(conkey) AS attnum
FROM pg_constraint
) con
LEFT JOIN pg_attribute "c" USING (attrelid, attnum)
LEFT JOIN pg_namespace s ON s.oid = con.connamespace
LEFT JOIN pg_class "t" ON "t".oid = con.conrelid
GROUP BY 1, 2, 3, 4, 5
) att USING ("table_schema", "table_name", "column_name")
) "c"
WHERE "table_schema" = ? AND "table_name" = ?
ORDER BY "table_schema", "table_name", "column_name"
`
// sqlInspectForeignKeys get FK definitions for user-defined tables.
// Pass bun.List([]string{...}) to exclude tables from this inspection or bun.List([]string{''}) to include all results.
sqlInspectForeignKeys = `
WITH
"schemas" AS (
SELECT oid, nspname
FROM pg_namespace
),
"tables" AS (
SELECT oid, relnamespace, relname, relkind
FROM pg_class
),
"columns" AS (
SELECT attrelid, attname, attnum
FROM pg_attribute
WHERE attisdropped = false
)
SELECT DISTINCT
co.conname AS "constraint_name",
ss.nspname AS schema_name,
s.relname AS "table_name",
ARRAY_AGG(sc.attname) AS "columns",
ts.nspname AS target_schema,
"t".relname AS target_table,
ARRAY_AGG(tc.attname) AS target_columns
FROM pg_constraint co
LEFT JOIN "tables" s ON s.oid = co.conrelid
LEFT JOIN "schemas" ss ON ss.oid = s.relnamespace
LEFT JOIN "columns" sc ON sc.attrelid = s.oid AND sc.attnum = ANY(co.conkey)
LEFT JOIN "tables" t ON t.oid = co.confrelid
LEFT JOIN "schemas" ts ON ts.oid = "t".relnamespace
LEFT JOIN "columns" tc ON tc.attrelid = "t".oid AND tc.attnum = ANY(co.confkey)
WHERE co.contype = 'f'
AND co.conrelid IN (SELECT oid FROM pg_class WHERE relkind = 'r')
AND ARRAY_POSITION(co.conkey, sc.attnum) = ARRAY_POSITION(co.confkey, tc.attnum)
AND ss.nspname = ?
AND s.relname NOT LIKE ALL (ARRAY[?])
AND "t".relname NOT LIKE ALL (ARRAY[?])
GROUP BY "constraint_name", "schema_name", "table_name", target_schema, target_table
`
)
+143
View File
@@ -0,0 +1,143 @@
package pgdialect
import (
"bytes"
"encoding/hex"
"github.com/uptrace/bun/internal/parser"
)
type pgparser struct {
parser.Parser
buf []byte
}
func newParser(b []byte) *pgparser {
p := new(pgparser)
p.Reset(b)
return p
}
func (p *pgparser) ReadLiteral(ch byte) []byte {
p.Unread()
lit, _ := p.ReadSep(',')
return lit
}
func (p *pgparser) ReadUnescapedSubstring(ch byte) ([]byte, error) {
return p.readSubstring(ch, false)
}
func (p *pgparser) ReadSubstring(ch byte) ([]byte, error) {
return p.readSubstring(ch, true)
}
func (p *pgparser) readSubstring(ch byte, escaped bool) ([]byte, error) {
ch, err := p.ReadByte()
if err != nil {
return nil, err
}
p.buf = p.buf[:0]
for {
if ch == '"' {
break
}
next, err := p.ReadByte()
if err != nil {
return nil, err
}
if ch == '\\' {
switch next {
case '\\', '"':
p.buf = append(p.buf, next)
ch, err = p.ReadByte()
if err != nil {
return nil, err
}
default:
p.buf = append(p.buf, '\\')
ch = next
}
continue
}
if escaped && ch == '\'' && next == '\'' {
p.buf = append(p.buf, next)
ch, err = p.ReadByte()
if err != nil {
return nil, err
}
continue
}
p.buf = append(p.buf, ch)
ch = next
}
if bytes.HasPrefix(p.buf, []byte("\\x")) && len(p.buf)%2 == 0 {
data := p.buf[2:]
buf := make([]byte, hex.DecodedLen(len(data)))
n, err := hex.Decode(buf, data)
if err != nil {
return nil, err
}
return buf[:n], nil
}
return p.buf, nil
}
func (p *pgparser) ReadRange(ch byte) ([]byte, error) {
p.buf = p.buf[:0]
p.buf = append(p.buf, ch)
for p.Valid() {
ch = p.Read()
p.buf = append(p.buf, ch)
if ch == ']' || ch == ')' {
break
}
}
return p.buf, nil
}
func (p *pgparser) ReadJSON() ([]byte, error) {
p.Unread()
c, err := p.ReadByte()
if err != nil {
return nil, err
}
p.buf = p.buf[:0]
depth := 0
for {
switch c {
case '{':
depth++
case '}':
depth--
}
p.buf = append(p.buf, c)
if depth == 0 {
break
}
next, err := p.ReadByte()
if err != nil {
return nil, err
}
c = next
}
return p.buf, nil
}
+217
View File
@@ -0,0 +1,217 @@
package pgdialect
import (
"bytes"
"database/sql"
"fmt"
"time"
"github.com/uptrace/bun/internal"
"github.com/uptrace/bun/schema"
)
type Range[T any] struct {
Lower, Upper T
LowerBound, UpperBound RangeBound
}
type MultiRange[T any] []Range[T]
type RangeBound byte
const (
// RangeBoundUnset indicates that no bound is set.
// This usually means the range is uninitialized or unspecified.
RangeBoundUnset RangeBound = 0x0
// RangeBoundEmpty is a special marker for an empty range.
// This is NOT a valid PostgreSQL bound character, but is used internally
// to represent a range that contains no values.
RangeBoundEmpty RangeBound = 'E'
RangeBoundInclusiveLeft RangeBound = '['
RangeBoundInclusiveRight RangeBound = ']'
RangeBoundExclusiveLeft RangeBound = '('
RangeBoundExclusiveRight RangeBound = ')'
)
type RangeOption[T any] func(*Range[T])
func NewRange[T any](lower, upper T) Range[T] {
r := Range[T]{
Lower: lower,
Upper: upper,
LowerBound: RangeBoundInclusiveLeft,
UpperBound: RangeBoundExclusiveRight,
}
return r
}
func NewEmptyRange[T any]() Range[T] {
return Range[T]{LowerBound: RangeBoundEmpty, UpperBound: RangeBoundEmpty}
}
func (r *Range[T]) IsZero() bool {
// NOTE: r.LowerBound represent
return r == nil || r.LowerBound == 0
}
func (r Range[T]) IsEmpty() bool {
return r.LowerBound == RangeBoundEmpty
}
var _ sql.Scanner = (*Range[any])(nil)
func (r *Range[T]) Scan(raw any) (err error) {
var src []byte
switch v := raw.(type) {
case []byte:
src = v
case string:
src = []byte(v)
case nil:
return nil
default:
return fmt.Errorf("pgdialect: Range can't scan %T", raw)
}
src = bytes.TrimSpace(src)
if len(src) == 0 {
return nil
}
if string(src) == "empty" {
r.LowerBound, r.UpperBound = RangeBoundEmpty, RangeBoundEmpty
return nil
}
switch src[0] {
case byte(RangeBoundInclusiveLeft), byte(RangeBoundExclusiveLeft):
r.LowerBound = RangeBound(src[0])
default:
return fmt.Errorf("unexpected lower bound: %s", string(src[:1]))
}
switch src[len(src)-1] {
case byte(RangeBoundInclusiveRight), byte(RangeBoundExclusiveRight):
r.UpperBound = RangeBound(src[len(src)-1])
default:
return fmt.Errorf("unexpected upper bound: %s", string(src[len(src)-1:]))
}
src = src[1 : len(src)-1]
ind := bytes.IndexByte(src, ',')
if ind == -1 {
return fmt.Errorf("invalid range: wanted comma, got %s", string(src))
}
left, right := src[:ind], src[ind+1:]
if len(left) > 0 {
_, err := scanElem(&r.Lower, left)
if err != nil {
return err
}
} else {
r.LowerBound = RangeBoundUnset
}
if len(right) > 0 {
_, err = scanElem(&r.Upper, right)
if err != nil {
return err
}
} else {
r.UpperBound = RangeBoundUnset
}
return nil
}
var _ schema.QueryAppender = (*Range[any])(nil)
func (r Range[T]) AppendQuery(_ schema.QueryGen, buf []byte) ([]byte, error) {
buf = append(buf, '\'')
buf = appendRange(buf, r)
buf = append(buf, '\'')
return buf, nil
}
func appendRange[T any](buf []byte, r Range[T]) []byte {
if r.IsEmpty() {
buf = append(buf, []byte("empty")...)
return buf
}
if r.LowerBound == RangeBoundUnset {
// NOTE from pg's document:
// > Specifying a missing bound as inclusive is automatically converted to exclusive, e.g., [,] is converted to (,).
buf = append(buf, byte(RangeBoundExclusiveLeft))
} else {
buf = append(buf, byte(r.LowerBound))
buf = appendElem(buf, r.Lower)
}
buf = append(buf, ',')
if r.UpperBound == RangeBoundUnset {
buf = append(buf, byte(RangeBoundExclusiveRight))
} else {
buf = appendElem(buf, r.Upper)
buf = append(buf, byte(r.UpperBound))
}
return buf
}
func (m *MultiRange[T]) Len() int {
if m == nil {
return 0
}
return len(([]Range[T])(*m))
}
func (m *MultiRange[T]) IsZero() bool {
return m.Len() == 0
}
func (m MultiRange[T]) AppendQuery(_ schema.QueryGen, buf []byte) ([]byte, error) {
if m == nil {
return append(buf, []byte("'{}'")...), nil
}
rs := ([]Range[T])(m)
buf = append(buf, '\'', '{')
for _, r := range rs {
buf = appendRange(buf, r)
buf = append(buf, ',')
}
if len(rs) > 0 {
buf[len(buf)-1] = '}'
} else {
buf = append(buf, '}')
}
buf = append(buf, '\'')
return buf, nil
}
func scanElem(ptr any, src []byte) ([]byte, error) {
// NOTE: for daterange, pg return 2024-12-01, for tzrange, pg return "2024-12-01 12:00:00"
if len(src) >= 2 && src[0] == '"' {
src = src[1 : len(src)-1]
}
switch ptr := ptr.(type) {
case *time.Time:
tm, err := internal.ParseTime(internal.String(src))
if err != nil {
return nil, err
}
*ptr = tm
return src, nil
case sql.Scanner:
if err := ptr.Scan(src); err != nil {
return nil, err
}
return src, nil
default:
panic(fmt.Errorf("unsupported range type: %T", ptr))
}
}
+190
View File
@@ -0,0 +1,190 @@
package pgdialect
import (
"database/sql"
"encoding/json"
"net"
"reflect"
"strings"
"github.com/uptrace/bun/dialect/sqltype"
"github.com/uptrace/bun/migrate/sqlschema"
"github.com/uptrace/bun/schema"
)
const (
// Date / Time
pgTypeTimestamp = "TIMESTAMP" // Timestamp
pgTypeTimestampWithTz = "TIMESTAMP WITH TIME ZONE" // Timestamp with a time zone
pgTypeTimestampTz = "TIMESTAMPTZ" // Timestamp with a time zone (alias)
pgTypeDate = "DATE" // Date
pgTypeTime = "TIME" // Time without a time zone
pgTypeTimeTz = "TIME WITH TIME ZONE" // Time with a time zone
pgTypeInterval = "INTERVAL" // Time interval
// Network Addresses
pgTypeInet = "INET" // IPv4 or IPv6 hosts and networks
pgTypeCidr = "CIDR" // IPv4 or IPv6 networks
pgTypeMacaddr = "MACADDR" // MAC addresses
// Serial Types
pgTypeSmallSerial = "SMALLSERIAL" // 2 byte autoincrementing integer
pgTypeSerial = "SERIAL" // 4 byte autoincrementing integer
pgTypeBigSerial = "BIGSERIAL" // 8 byte autoincrementing integer
// Character Types
pgTypeChar = "CHAR" // fixed length string (blank padded)
pgTypeCharacter = "CHARACTER" // alias for CHAR
pgTypeText = "TEXT" // variable length string without limit
pgTypeVarchar = "VARCHAR" // variable length string with optional limit
pgTypeCharacterVarying = "CHARACTER VARYING" // alias for VARCHAR
// Binary Data Types
pgTypeBytea = "BYTEA" // binary string
)
var (
ipType = reflect.TypeFor[net.IP]()
ipNetType = reflect.TypeFor[net.IPNet]()
jsonRawMessageType = reflect.TypeFor[json.RawMessage]()
nullStringType = reflect.TypeFor[sql.NullString]()
)
func (d *Dialect) DefaultVarcharLen() int {
return 0
}
func (d *Dialect) DefaultSchema() string {
return "public"
}
func fieldSQLType(field *schema.Field) string {
if field.UserSQLType != "" {
return field.UserSQLType
}
if v, ok := field.Tag.Option("composite"); ok {
return v
}
if field.Tag.HasOption("hstore") {
return sqltype.HSTORE
}
if field.Tag.HasOption("array") {
switch field.IndirectType.Kind() {
case reflect.Slice, reflect.Array:
sqlType := sqlType(field.IndirectType.Elem())
return sqlType + "[]"
}
}
if field.DiscoveredSQLType == sqltype.Blob {
return pgTypeBytea
}
return sqlType(field.IndirectType)
}
func sqlType(typ reflect.Type) string {
if typ.Kind() == reflect.Ptr {
typ = typ.Elem()
}
switch typ {
case nullStringType: // typ.Kind() == reflect.Struct, test for exact match
return sqltype.VarChar
case ipType:
return pgTypeInet
case ipNetType:
return pgTypeCidr
case jsonRawMessageType:
return sqltype.JSONB
}
sqlType := schema.DiscoverSQLType(typ)
switch sqlType {
case sqltype.Timestamp:
sqlType = pgTypeTimestampTz
}
switch typ.Kind() {
case reflect.Map, reflect.Struct: // except typ == nullStringType, see above
if sqlType == sqltype.VarChar {
return sqltype.JSONB
}
return sqlType
case reflect.Array, reflect.Slice:
if typ.Elem().Kind() == reflect.Uint8 {
return pgTypeBytea
}
return sqltype.JSONB
}
return sqlType
}
var (
char = newAliases(pgTypeChar, pgTypeCharacter)
varchar = newAliases(pgTypeVarchar, pgTypeCharacterVarying)
timestampTz = newAliases(sqltype.Timestamp, pgTypeTimestampTz, pgTypeTimestampWithTz)
bigint = newAliases(sqltype.BigInt, pgTypeBigSerial)
integer = newAliases(sqltype.Integer, pgTypeSerial)
smallint = newAliases(sqltype.SmallInt, pgTypeSmallSerial)
)
func (d *Dialect) CompareType(col1, col2 sqlschema.Column) bool {
typ1, typ2 := strings.ToUpper(col1.GetSQLType()), strings.ToUpper(col2.GetSQLType())
if typ1 == typ2 {
return checkVarcharLen(col1, col2, d.DefaultVarcharLen())
}
switch {
case char.IsAlias(typ1) && char.IsAlias(typ2):
return checkVarcharLen(col1, col2, d.DefaultVarcharLen())
case varchar.IsAlias(typ1) && varchar.IsAlias(typ2):
return checkVarcharLen(col1, col2, d.DefaultVarcharLen())
case timestampTz.IsAlias(typ1) && timestampTz.IsAlias(typ2):
return true
case bigint.IsAlias(typ1) && bigint.IsAlias(typ2),
integer.IsAlias(typ1) && integer.IsAlias(typ2),
smallint.IsAlias(typ1) && smallint.IsAlias(typ2):
return true
}
return false
}
// checkVarcharLen returns true if columns have the same VarcharLen, or,
// if one specifies no VarcharLen and the other one has the default lenght for pgdialect.
// We assume that the types are otherwise equivalent and that any non-character column
// would have VarcharLen == 0;
func checkVarcharLen(col1, col2 sqlschema.Column, defaultLen int) bool {
vl1, vl2 := col1.GetVarcharLen(), col2.GetVarcharLen()
if vl1 == vl2 {
return true
}
if (vl1 == 0 && vl2 == defaultLen) || (vl1 == defaultLen && vl2 == 0) {
return true
}
return false
}
// typeAlias defines aliases for common data types. It is a lightweight string set implementation.
type typeAlias map[string]struct{}
// IsAlias checks if typ1 and typ2 are aliases of the same data type.
func (t typeAlias) IsAlias(typ string) bool {
_, ok := t[typ]
return ok
}
// newAliases creates a set of aliases.
func newAliases(aliases ...string) typeAlias {
types := make(typeAlias)
for _, a := range aliases {
types[a] = struct{}{}
}
return types
}
+6
View File
@@ -0,0 +1,6 @@
package pgdialect
// Version is the current release version.
func Version() string {
return "1.2.18"
}
+125
View File
@@ -0,0 +1,125 @@
package ordered
// Pair represents a key-value pair in the ordered map.
type Pair[K comparable, V any] struct {
Key K
Value V
next, prev *Pair[K, V] // Pointers to the next and previous pairs in the linked list.
}
// Map represents an ordered map.
type Map[K comparable, V any] struct {
root *Pair[K, V] // Sentinel node for the circular doubly linked list.
zero V // Zero value for the value type.
pairs map[K]*Pair[K, V] // Map from keys to pairs.
}
// NewMap creates a new ordered map with optional initial data.
func NewMap[K comparable, V any](initialData ...Pair[K, V]) *Map[K, V] {
m := &Map[K, V]{}
m.Clear()
for _, pair := range initialData {
m.Store(pair.Key, pair.Value)
}
return m
}
// Clear removes all pairs from the map.
func (m *Map[K, V]) Clear() {
if m.root != nil {
m.root.next, m.root.prev = nil, nil // avoid memory leaks
}
for _, pair := range m.pairs {
pair.next, pair.prev = nil, nil // avoid memory leaks
}
m.root = &Pair[K, V]{}
m.root.next, m.root.prev = m.root, m.root
m.pairs = make(map[K]*Pair[K, V])
}
// Len returns the number of pairs in the map.
func (m *Map[K, V]) Len() int {
return len(m.pairs)
}
// Load returns the value associated with the key, and a boolean indicating if the key was found.
func (m *Map[K, V]) Load(key K) (V, bool) {
if pair, present := m.pairs[key]; present {
return pair.Value, true
}
return m.zero, false
}
// Value returns the value associated with the key, or the zero value if the key is not found.
func (m *Map[K, V]) Value(key K) V {
if pair, present := m.pairs[key]; present {
return pair.Value
}
return m.zero
}
// Store adds or updates a key-value pair in the map.
func (m *Map[K, V]) Store(key K, value V) {
if pair, present := m.pairs[key]; present {
pair.Value = value
return
}
pair := &Pair[K, V]{Key: key, Value: value}
pair.prev = m.root.prev
m.root.prev.next = pair
m.root.prev = pair
pair.next = m.root
m.pairs[key] = pair
}
// Delete removes a key-value pair from the map.
func (m *Map[K, V]) Delete(key K) {
if pair, present := m.pairs[key]; present {
pair.prev.next = pair.next
pair.next.prev = pair.prev
pair.next, pair.prev = nil, nil // avoid memory leaks
delete(m.pairs, key)
}
}
// Range calls the given function for each key-value pair in the map in order.
func (m *Map[K, V]) Range(yield func(key K, value V) bool) {
for pair := m.root.next; pair != m.root; pair = pair.next {
if !yield(pair.Key, pair.Value) {
break
}
}
}
// Keys returns a slice of all keys in the map in order.
func (m *Map[K, V]) Keys() []K {
keys := make([]K, 0, len(m.pairs))
m.Range(func(key K, _ V) bool {
keys = append(keys, key)
return true
})
return keys
}
// Values returns a slice of all values in the map in order.
func (m *Map[K, V]) Values() []V {
values := make([]V, 0, len(m.pairs))
m.Range(func(_ K, value V) bool {
values = append(values, value)
return true
})
return values
}
// Pairs returns a slice of all key-value pairs in the map in order.
func (m *Map[K, V]) Pairs() []Pair[K, V] {
pairs := make([]Pair[K, V], 0, len(m.pairs))
m.Range(func(key K, value V) bool {
pairs = append(pairs, Pair[K, V]{Key: key, Value: value})
return true
})
return pairs
}
+475
View File
@@ -0,0 +1,475 @@
package migrate
import (
"bytes"
"context"
"errors"
"fmt"
"io"
"os"
"path/filepath"
"github.com/uptrace/bun"
"github.com/uptrace/bun/internal"
"github.com/uptrace/bun/migrate/sqlschema"
"github.com/uptrace/bun/schema"
)
// AutoMigratorOption configures an AutoMigrator.
type AutoMigratorOption func(m *AutoMigrator)
// WithModel adds a bun.Model to the migration scope.
func WithModel(models ...any) AutoMigratorOption {
return func(m *AutoMigrator) {
m.includeModels = append(m.includeModels, models...)
}
}
// WithExcludeTable tells AutoMigrator to exclude database tables from the migration scope.
// This prevents AutoMigrator from dropping tables which may exist in the schema
// but which are not used by the application.
//
// Expressions may make use of the wildcards supported by the SQL LIKE operator:
// - % as a wildcard
// - _ as a single character
//
// Do not exclude tables included via WithModel, as BunModelInspector ignores this setting.
func WithExcludeTable(tables ...string) AutoMigratorOption {
return func(m *AutoMigrator) {
m.excludeTables = append(m.excludeTables, tables...)
}
}
// WithExcludeForeignKeys tells AutoMigrator to exclude a foreign key constaint
// from the migration scope. This prevents AutoMigrator from dropping foreign keys
// that are defined manually via CreateTableQuery.ForeignKey().
func WithExcludeForeignKeys(fks ...sqlschema.ForeignKey) AutoMigratorOption {
return func(m *AutoMigrator) {
m.excludeForeignKeys = append(m.excludeForeignKeys, fks...)
}
}
// WithSchemaName sets the database schema to migrate objects in.
// By default, dialects' default schema is used.
func WithSchemaName(schemaName string) AutoMigratorOption {
return func(m *AutoMigrator) {
m.schemaName = schemaName
}
}
// WithTableNameAuto overrides default migrations table name.
func WithTableNameAuto(table string) AutoMigratorOption {
return func(m *AutoMigrator) {
m.table = table
m.migratorOpts = append(m.migratorOpts, WithTableName(table))
}
}
// WithLocksTableNameAuto overrides default migration locks table name.
func WithLocksTableNameAuto(table string) AutoMigratorOption {
return func(m *AutoMigrator) {
m.locksTable = table
m.migratorOpts = append(m.migratorOpts, WithLocksTableName(table))
}
}
// WithMarkAppliedOnSuccessAuto sets the migrator to only mark migrations as applied/unapplied
// when their up/down is successful.
func WithMarkAppliedOnSuccessAuto(enabled bool) AutoMigratorOption {
return func(m *AutoMigrator) {
m.migratorOpts = append(m.migratorOpts, WithMarkAppliedOnSuccess(enabled))
}
}
// WithMigrationsDirectoryAuto overrides the default directory for migration files.
func WithMigrationsDirectoryAuto(directory string) AutoMigratorOption {
return func(m *AutoMigrator) {
m.migrationsOpts = append(m.migrationsOpts, WithMigrationsDirectory(directory))
}
}
// AutoMigrator performs automated schema migrations.
//
// It is designed to be a drop-in replacement for some Migrator functionality and supports all existing
// configuration options.
// Similarly to Migrator, it has methods to create SQL migrations, write them to a file, and apply them.
// Unlike Migrator, it detects the differences between the state defined by bun models and the current
// database schema automatically.
//
// Usage:
// 1. Generate migrations and apply them at once with AutoMigrator.Migrate().
// 2. Create up- and down-SQL migration files and apply migrations using Migrator.Migrate().
//
// While both methods produce complete, reversible migrations (with entries in the database
// and SQL migration files), prefer creating migrations and applying them separately for
// any non-trivial cases to ensure AutoMigrator detects expected changes correctly.
//
// Limitations:
// - AutoMigrator only supports a subset of the possible ALTER TABLE modifications.
// - Some changes are not automatically reversible. For example, you would need to manually
// add a CREATE TABLE query to the .down migration file to revert a DROP TABLE migration.
// - Does not validate most dialect-specific constraints. For example, when changing column
// data type, make sure the data con be auto-casted to the new type.
// - Due to how the schema-state diff is calculated, it is not possible to rename a table and
// modify any of its columns' _data type_ in a single run. This will cause the AutoMigrator
// to drop and re-create the table under a different name; it is better to apply this change in 2 steps.
// Renaming a table and renaming its columns at the same time is possible.
// - Renaming table/column to an existing name, i.e. like this [A->B] [B->C], is not possible due to how
// AutoMigrator distinguishes "rename" and "unchanged" columns.
//
// Dialect must implement both sqlschema.Inspector and sqlschema.Migrator to be used with AutoMigrator.
type AutoMigrator struct {
db *bun.DB
// dbInspector creates the current state for the target database.
dbInspector sqlschema.Inspector
// modelInspector creates the desired state based on the model definitions.
modelInspector sqlschema.Inspector
// dbMigrator executes ALTER TABLE queries.
dbMigrator sqlschema.Migrator
table string // Migrations table (excluded from database inspection)
locksTable string // Migration locks table (excluded from database inspection)
// schemaName is the database schema considered for migration.
schemaName string
// includeModels define the migration scope.
includeModels []any
excludeTables []string // excludeTables are excluded from database inspection.
excludeForeignKeys []sqlschema.ForeignKey // excludeForeignKeys are excluded from database inspection.
// diffOpts are passed to detector constructor.
diffOpts []diffOption
// migratorOpts are passed to Migrator constructor.
migratorOpts []MigratorOption
// migrationsOpts are passed to Migrations constructor.
migrationsOpts []MigrationsOption
}
// NewAutoMigrator creates an AutoMigrator that detects schema differences and generates migrations.
func NewAutoMigrator(db *bun.DB, opts ...AutoMigratorOption) (*AutoMigrator, error) {
am := &AutoMigrator{
db: db,
table: defaultTable,
locksTable: defaultLocksTable,
schemaName: db.Dialect().DefaultSchema(),
}
for _, opt := range opts {
opt(am)
}
am.excludeTables = append(am.excludeTables, am.table, am.locksTable)
dbInspector, err := sqlschema.NewInspector(db,
sqlschema.WithSchemaName(am.schemaName),
sqlschema.WithExcludeTables(am.excludeTables...),
sqlschema.WithExcludeForeignKeys(am.excludeForeignKeys...),
)
if err != nil {
return nil, err
}
am.dbInspector = dbInspector
am.diffOpts = append(am.diffOpts, withCompareTypeFunc(db.Dialect().(sqlschema.InspectorDialect).CompareType))
dbMigrator, err := sqlschema.NewMigrator(db, am.schemaName)
if err != nil {
return nil, err
}
am.dbMigrator = dbMigrator
tables := schema.NewTables(db.Dialect())
tables.Register(am.includeModels...)
am.modelInspector = sqlschema.NewBunModelInspector(tables, sqlschema.WithSchemaName(am.schemaName))
return am, nil
}
func (am *AutoMigrator) plan(ctx context.Context) (*changeset, error) {
var err error
got, err := am.dbInspector.Inspect(ctx)
if err != nil {
return nil, err
}
want, err := am.modelInspector.Inspect(ctx)
if err != nil {
return nil, err
}
changes := diff(got, want, am.diffOpts...)
if err := changes.ResolveDependencies(); err != nil {
return nil, fmt.Errorf("plan migrations: %w", err)
}
return changes, nil
}
// Migrate writes required changes to a new migration file and runs the migration.
// This will create an entry in the migrations table, making it possible to revert
// the changes with Migrator.Rollback(). MigrationOptions are passed on to Migrator.Migrate().
func (am *AutoMigrator) Migrate(ctx context.Context, opts ...MigrationOption) (*MigrationGroup, error) {
migrations, _, err := am.createSQLMigrations(ctx, false)
if err != nil {
if err == errNothingToMigrate {
return new(MigrationGroup), nil
}
return nil, fmt.Errorf("auto migrate: %w", err)
}
migrator := NewMigrator(am.db, migrations, am.migratorOpts...)
if err := migrator.Init(ctx); err != nil {
return nil, fmt.Errorf("auto migrate: %w", err)
}
group, err := migrator.Migrate(ctx, opts...)
if err != nil {
return nil, fmt.Errorf("auto migrate: %w", err)
}
return group, nil
}
// CreateSQLMigration writes required changes to a new migration file.
// Use migrate.Migrator to apply the generated migrations.
func (am *AutoMigrator) CreateSQLMigrations(ctx context.Context) ([]*MigrationFile, error) {
_, files, err := am.createSQLMigrations(ctx, false)
if err == errNothingToMigrate {
return files, nil
}
return files, err
}
// CreateTxSQLMigration writes required changes to a new migration file making sure they will be executed
// in a transaction when applied. Use migrate.Migrator to apply the generated migrations.
func (am *AutoMigrator) CreateTxSQLMigrations(ctx context.Context) ([]*MigrationFile, error) {
_, files, err := am.createSQLMigrations(ctx, true)
if err == errNothingToMigrate {
return files, nil
}
return files, err
}
// errNothingToMigrate is a sentinel error which means the database is already in a desired state.
// Should not be returned to the user -- return a nil-error instead.
var errNothingToMigrate = errors.New("nothing to migrate")
func (am *AutoMigrator) createSQLMigrations(ctx context.Context, transactional bool) (*Migrations, []*MigrationFile, error) {
changes, err := am.plan(ctx)
if err != nil {
return nil, nil, fmt.Errorf("create sql migrations: %w", err)
}
if changes.Len() == 0 {
return nil, nil, errNothingToMigrate
}
name, _ := genMigrationName(am.schemaName + "_auto")
migrations := NewMigrations(am.migrationsOpts...)
migrations.Add(Migration{
Name: name,
Up: wrapGoMigrationFunc(changes.Up(am.dbMigrator)),
Down: wrapGoMigrationFunc(changes.Down(am.dbMigrator)),
Comment: "Changes detected by bun.AutoMigrator",
})
// Append .tx.up.sql or .up.sql to migration name, depending if it should be transactional.
fname := func(direction string) string {
return name + map[bool]string{true: ".tx.", false: "."}[transactional] + direction + ".sql"
}
up, err := am.createSQL(ctx, migrations, fname("up"), changes, transactional)
if err != nil {
return nil, nil, fmt.Errorf("create sql migration up: %w", err)
}
down, err := am.createSQL(ctx, migrations, fname("down"), changes.GetReverse(), transactional)
if err != nil {
return nil, nil, fmt.Errorf("create sql migration down: %w", err)
}
return migrations, []*MigrationFile{up, down}, nil
}
func (am *AutoMigrator) createSQL(_ context.Context, migrations *Migrations, fname string, changes *changeset, transactional bool) (*MigrationFile, error) {
var buf bytes.Buffer
if transactional {
buf.WriteString("SET statement_timeout = 0;")
}
if err := changes.WriteTo(&buf, am.dbMigrator); err != nil {
return nil, err
}
content := buf.Bytes()
fpath := filepath.Join(migrations.getDirectory(), fname)
if err := os.WriteFile(fpath, content, 0o644); err != nil {
return nil, err
}
mf := &MigrationFile{
Name: fname,
Path: fpath,
Content: string(content),
}
return mf, nil
}
func (c *changeset) Len() int {
return len(c.operations)
}
// Func creates a MigrationFunc that applies all operations all the changeset.
func (c *changeset) Func(m sqlschema.Migrator) MigrationFunc {
return func(ctx context.Context, db *bun.DB) error {
return c.apply(ctx, db, m)
}
}
// GetReverse returns a new changeset with each operation in it "reversed" and in reverse order.
func (c *changeset) GetReverse() *changeset {
var reverse changeset
for i := len(c.operations) - 1; i >= 0; i-- {
reverse.Add(c.operations[i].GetReverse())
}
return &reverse
}
// Up is syntactic sugar.
func (c *changeset) Up(m sqlschema.Migrator) MigrationFunc {
return c.Func(m)
}
// Down is syntactic sugar.
func (c *changeset) Down(m sqlschema.Migrator) MigrationFunc {
return c.GetReverse().Func(m)
}
// apply generates SQL for each operation and executes it.
func (c *changeset) apply(ctx context.Context, db *bun.DB, m sqlschema.Migrator) error {
if len(c.operations) == 0 {
return nil
}
for _, op := range c.operations {
if _, skip := op.(*Unimplemented); skip {
continue
}
b := internal.MakeQueryBytes()
b, err := m.AppendSQL(b, op)
if err != nil {
return fmt.Errorf("apply changes: %w", err)
}
query := internal.String(b)
if _, err = db.ExecContext(ctx, query); err != nil {
return fmt.Errorf("apply changes: %w", err)
}
}
return nil
}
func (c *changeset) WriteTo(w io.Writer, m sqlschema.Migrator) error {
var err error
b := internal.MakeQueryBytes()
for _, op := range c.operations {
if comment, isComment := op.(*Unimplemented); isComment {
b = append(b, "/*\n"...)
b = append(b, *comment...)
b = append(b, "\n*/"...)
continue
}
// Append each query separately, merge later.
// Dialects assume that the []byte only holds
// the contents of a single query and may be misled.
queryBytes := internal.MakeQueryBytes()
queryBytes, err = m.AppendSQL(queryBytes, op)
if err != nil {
return fmt.Errorf("write changeset: %w", err)
}
b = append(b, queryBytes...)
b = append(b, ";\n"...)
}
if _, err := w.Write(b); err != nil {
return fmt.Errorf("write changeset: %w", err)
}
return nil
}
func (c *changeset) ResolveDependencies() error {
if len(c.operations) <= 1 {
return nil
}
const (
unvisited = iota
current
visited
)
status := make(map[Operation]int, len(c.operations))
for _, op := range c.operations {
status[op] = unvisited
}
var resolved []Operation
var nextOp Operation
var visit func(op Operation) error
next := func() bool {
for op, s := range status {
if s == unvisited {
nextOp = op
return true
}
}
return false
}
// visit iterates over c.operations until it finds all operations that depend on the current one
// or runs into circular dependency, in which case it will return an error.
visit = func(op Operation) error {
switch status[op] {
case visited:
return nil
case current:
// TODO: add details (circle) to the error message
return errors.New("detected circular dependency")
}
status[op] = current
for _, another := range c.operations {
if dop, hasDeps := another.(interface {
DependsOn(Operation) bool
}); another == op || !hasDeps || !dop.DependsOn(op) {
continue
}
if err := visit(another); err != nil {
return err
}
}
status[op] = visited
// Any dependent nodes would've already been added to the list by now, so we prepend.
resolved = append([]Operation{op}, resolved...)
return nil
}
for next() {
if err := visit(nextOp); err != nil {
return err
}
}
c.operations = resolved
return nil
}
+426
View File
@@ -0,0 +1,426 @@
package migrate
import (
"github.com/uptrace/bun/internal/ordered"
"github.com/uptrace/bun/migrate/sqlschema"
)
// changeset is a set of changes to the database schema definition.
type changeset struct {
operations []Operation
}
// Add new operations to the changeset.
func (c *changeset) Add(op ...Operation) {
c.operations = append(c.operations, op...)
}
// diff calculates the diff between the current database schema and the target state.
// The changeset is not sorted -- the caller should resolve dependencies before applying the changes.
func diff(got, want sqlschema.Database, opts ...diffOption) *changeset {
d := newDetector(got, want, opts...)
return d.detectChanges()
}
func (d *detector) detectChanges() *changeset {
currentTables := toOrderedMap(d.current.GetTables())
targetTables := toOrderedMap(d.target.GetTables())
RenameCreate:
for _, wantPair := range targetTables.Pairs() {
wantName, wantTable := wantPair.Key, wantPair.Value
// A table with this name exists in the database. We assume that schema objects won't
// be renamed to an already existing name, nor do we support such cases.
// Simply check if the table definition has changed.
if haveTable, ok := currentTables.Load(wantName); ok {
d.detectColumnChanges(haveTable, wantTable, true)
d.detectConstraintChanges(haveTable, wantTable)
continue
}
// Find all renamed tables. We assume that renamed tables have the same signature.
for _, havePair := range currentTables.Pairs() {
haveName, haveTable := havePair.Key, havePair.Value
if _, exists := targetTables.Load(haveName); !exists && d.canRename(haveTable, wantTable) {
d.changes.Add(&RenameTableOp{
TableName: haveTable.GetName(),
NewName: wantName,
})
d.refMap.RenameTable(haveTable.GetName(), wantName)
// Find renamed columns, if any, and check if constraints (PK, UNIQUE) have been updated.
// We need not check wantTable any further.
d.detectColumnChanges(haveTable, wantTable, false)
d.detectConstraintChanges(haveTable, wantTable)
currentTables.Delete(haveName)
continue RenameCreate
}
}
// If wantTable does not exist in the database and was not renamed
// then we need to create this table in the database.
additional := wantTable.(*sqlschema.BunTable)
d.changes.Add(&CreateTableOp{
TableName: wantTable.GetName(),
Model: additional.Model,
})
}
// Drop any remaining "current" tables which do not have a model.
for _, tPair := range currentTables.Pairs() {
name, table := tPair.Key, tPair.Value
if _, keep := targetTables.Load(name); !keep {
d.changes.Add(&DropTableOp{
TableName: table.GetName(),
})
}
}
targetFKs := d.target.GetForeignKeys()
currentFKs := d.refMap.Deref()
for fk := range targetFKs {
if _, ok := currentFKs[fk]; !ok {
d.changes.Add(&AddForeignKeyOp{
ForeignKey: fk,
ConstraintName: "", // leave empty to let each dialect apply their convention
})
}
}
for fk, name := range currentFKs {
if _, ok := targetFKs[fk]; !ok {
d.changes.Add(&DropForeignKeyOp{
ConstraintName: name,
ForeignKey: fk,
})
}
}
return &d.changes
}
// toOrderedMap transforms a slice of objects to an ordered map, using return of GetName() as key.
func toOrderedMap[V interface{ GetName() string }](named []V) *ordered.Map[string, V] {
m := ordered.NewMap[string, V]()
for _, v := range named {
m.Store(v.GetName(), v)
}
return m
}
// detectColumnChanges finds renamed columns and, if checkType == true, columns with changed type.
func (d *detector) detectColumnChanges(current, target sqlschema.Table, checkType bool) {
currentColumns := toOrderedMap(current.GetColumns())
targetColumns := toOrderedMap(target.GetColumns())
ChangeRename:
for _, tPair := range targetColumns.Pairs() {
tName, tCol := tPair.Key, tPair.Value
// This column exists in the database, so it hasn't been renamed, dropped, or added.
// Still, we should not delete(columns, thisColumn), because later we will need to
// check that we do not try to rename a column to an already a name that already exists.
if cCol, ok := currentColumns.Load(tName); ok {
if checkType && !d.equalColumns(cCol, tCol) {
d.changes.Add(&ChangeColumnTypeOp{
TableName: target.GetName(),
Column: tName,
From: cCol,
To: d.makeTargetColDef(cCol, tCol),
})
}
continue
}
// Column tName does not exist in the database -- it's been either renamed or added.
// Find renamed columns first.
for _, cPair := range currentColumns.Pairs() {
cName, cCol := cPair.Key, cPair.Value
// Cannot rename if a column with this name already exists or the types differ.
if _, exists := targetColumns.Load(cName); exists || !d.equalColumns(tCol, cCol) {
continue
}
d.changes.Add(&RenameColumnOp{
TableName: target.GetName(),
OldName: cName,
NewName: tName,
})
d.refMap.RenameColumn(target.GetName(), cName, tName)
currentColumns.Delete(cName) // no need to check this column again
// Update primary key definition to avoid superficially recreating the constraint.
current.GetPrimaryKey().Columns.Replace(cName, tName)
continue ChangeRename
}
d.changes.Add(&AddColumnOp{
TableName: target.GetName(),
ColumnName: tName,
Column: tCol,
})
}
// Drop columns which do not exist in the target schema and were not renamed.
for _, cPair := range currentColumns.Pairs() {
cName, cCol := cPair.Key, cPair.Value
if _, keep := targetColumns.Load(cName); !keep {
d.changes.Add(&DropColumnOp{
TableName: target.GetName(),
ColumnName: cName,
Column: cCol,
})
}
}
}
func (d *detector) detectConstraintChanges(current, target sqlschema.Table) {
Add:
for _, want := range target.GetUniqueConstraints() {
for _, got := range current.GetUniqueConstraints() {
if got.Equals(want) {
continue Add
}
}
d.changes.Add(&AddUniqueConstraintOp{
TableName: target.GetName(),
Unique: want,
})
}
Drop:
for _, got := range current.GetUniqueConstraints() {
for _, want := range target.GetUniqueConstraints() {
if got.Equals(want) {
continue Drop
}
}
d.changes.Add(&DropUniqueConstraintOp{
TableName: target.GetName(),
Unique: got,
})
}
targetPK := target.GetPrimaryKey()
currentPK := current.GetPrimaryKey()
// Detect primary key changes
if targetPK == nil && currentPK == nil {
return
}
switch {
case targetPK == nil && currentPK != nil:
d.changes.Add(&DropPrimaryKeyOp{
TableName: target.GetName(),
PrimaryKey: *currentPK,
})
case currentPK == nil && targetPK != nil:
d.changes.Add(&AddPrimaryKeyOp{
TableName: target.GetName(),
PrimaryKey: *targetPK,
})
case targetPK.Columns != currentPK.Columns:
d.changes.Add(&ChangePrimaryKeyOp{
TableName: target.GetName(),
Old: *currentPK,
New: *targetPK,
})
}
}
func newDetector(got, want sqlschema.Database, opts ...diffOption) *detector {
cfg := &detectorConfig{
cmpType: func(c1, c2 sqlschema.Column) bool {
return c1.GetSQLType() == c2.GetSQLType() && c1.GetVarcharLen() == c2.GetVarcharLen()
},
}
for _, opt := range opts {
opt(cfg)
}
return &detector{
current: got,
target: want,
refMap: newRefMap(got.GetForeignKeys()),
cmpType: cfg.cmpType,
}
}
type diffOption func(*detectorConfig)
func withCompareTypeFunc(f CompareTypeFunc) diffOption {
return func(cfg *detectorConfig) {
cfg.cmpType = f
}
}
// detectorConfig controls how differences in the model states are resolved.
type detectorConfig struct {
cmpType CompareTypeFunc
}
// detector may modify the passed database schemas, so it isn't safe to re-use them.
type detector struct {
// current state represents the existing database schema.
current sqlschema.Database
// target state represents the database schema defined in bun models.
target sqlschema.Database
changes changeset
refMap refMap
// cmpType determines column type equivalence.
// Default is direct comparison with '==' operator, which is inaccurate
// due to the existence of dialect-specific type aliases. The caller
// should pass a concrete InspectorDialect.EquivalentType for robust comparison.
cmpType CompareTypeFunc
}
// canRename checks if t1 can be renamed to t2.
func (d detector) canRename(t1, t2 sqlschema.Table) bool {
return t1.GetSchema() == t2.GetSchema() && equalSignatures(t1, t2, d.equalColumns)
}
func (d detector) equalColumns(col1, col2 sqlschema.Column) bool {
return d.cmpType(col1, col2) &&
col1.GetDefaultValue() == col2.GetDefaultValue() &&
col1.GetIsNullable() == col2.GetIsNullable() &&
col1.GetIsAutoIncrement() == col2.GetIsAutoIncrement() &&
col1.GetIsIdentity() == col2.GetIsIdentity()
}
func (d detector) makeTargetColDef(current, target sqlschema.Column) sqlschema.Column {
// Avoid unnecessary type-change migrations if the types are equivalent.
if d.cmpType(current, target) {
target = &sqlschema.BaseColumn{
Name: target.GetName(),
DefaultValue: target.GetDefaultValue(),
IsNullable: target.GetIsNullable(),
IsAutoIncrement: target.GetIsAutoIncrement(),
IsIdentity: target.GetIsIdentity(),
SQLType: current.GetSQLType(),
VarcharLen: current.GetVarcharLen(),
}
}
return target
}
// CompareTypeFunc compares two column definitions and reports whether they have the same SQL type.
type CompareTypeFunc func(sqlschema.Column, sqlschema.Column) bool
// equalSignatures determines if two tables have the same "signature".
func equalSignatures(t1, t2 sqlschema.Table, eq CompareTypeFunc) bool {
sig1 := newSignature(t1, eq)
sig2 := newSignature(t2, eq)
return sig1.Equals(sig2)
}
// signature is a set of column definitions, which allows "relation/name-agnostic" comparison between them;
// meaning that two columns are considered equal if their types are the same.
type signature struct {
// underlying stores the number of occurrences for each unique column type.
// It helps to account for the fact that a table might have multiple columns that have the same type.
underlying map[sqlschema.BaseColumn]int
eq CompareTypeFunc
}
func newSignature(t sqlschema.Table, eq CompareTypeFunc) signature {
s := signature{
underlying: make(map[sqlschema.BaseColumn]int),
eq: eq,
}
s.scan(t)
return s
}
// scan iterates over table's field and counts occurrences of each unique column definition.
func (s *signature) scan(t sqlschema.Table) {
for _, icol := range t.GetColumns() {
scanCol := icol.(*sqlschema.BaseColumn)
// This is slightly more expensive than if the columns could be compared directly
// and we always did s.underlying[col]++, but we get type-equivalence in return.
col, count := s.getCount(*scanCol)
if count == 0 {
s.underlying[*scanCol] = 1
} else {
s.underlying[col]++
}
}
}
// getCount uses CompareTypeFunc to find a column with the same (equivalent) SQL type
// and returns its count. Count 0 means there are no columns with of this type.
func (s *signature) getCount(keyCol sqlschema.BaseColumn) (key sqlschema.BaseColumn, count int) {
for col, cnt := range s.underlying {
if s.eq(&col, &keyCol) {
return col, cnt
}
}
return keyCol, 0
}
// Equals returns true if 2 signatures share an identical set of columns.
func (s *signature) Equals(other signature) bool {
if len(s.underlying) != len(other.underlying) {
return false
}
for col, count := range s.underlying {
if _, countOther := other.getCount(col); countOther != count {
return false
}
}
return true
}
// refMap is a utility for tracking superficial changes in foreign keys,
// which do not require any modification in the database.
// Modern SQL dialects automatically updated foreign key constraints whenever
// a column or a table is renamed. Detector can use refMap to ignore any
// differences in foreign keys which were caused by renamed column/table.
type refMap map[*sqlschema.ForeignKey]string
func newRefMap(fks map[sqlschema.ForeignKey]string) refMap {
rm := make(map[*sqlschema.ForeignKey]string)
for fk, name := range fks {
rm[&fk] = name
}
return rm
}
// RenameT updates table name in all foreign key definions which depend on it.
func (rm refMap) RenameTable(tableName string, newName string) {
for fk := range rm {
switch tableName {
case fk.From.TableName:
fk.From.TableName = newName
case fk.To.TableName:
fk.To.TableName = newName
}
}
}
// RenameColumn updates column name in all foreign key definions which depend on it.
func (rm refMap) RenameColumn(tableName string, column, newName string) {
for fk := range rm {
if tableName == fk.From.TableName {
fk.From.Column.Replace(column, newName)
}
if tableName == fk.To.TableName {
fk.To.Column.Replace(column, newName)
}
}
}
// Deref returns copies of ForeignKey values to a map.
func (rm refMap) Deref() map[sqlschema.ForeignKey]string {
out := make(map[sqlschema.ForeignKey]string)
for fk, name := range rm {
out[*fk] = name
}
return out
}
+357
View File
@@ -0,0 +1,357 @@
package migrate
import (
"bufio"
"bytes"
"context"
"fmt"
"io"
"io/fs"
"slices"
"strings"
"text/template"
"time"
"github.com/uptrace/bun"
)
// Migration represents a single database migration with up and down functions.
type Migration struct {
bun.BaseModel
ID int64 `bun:",pk,autoincrement"`
Name string
Comment string `bun:"-"`
GroupID int64
MigratedAt time.Time `bun:",notnull,nullzero,default:current_timestamp"`
Up internalMigrationFunc `bun:"-"`
Down internalMigrationFunc `bun:"-"`
}
// String returns the migration name and comment.
func (m Migration) String() string {
return fmt.Sprintf("%s_%s", m.Name, m.Comment)
}
// IsApplied reports whether the migration has been applied.
func (m Migration) IsApplied() bool {
return m.ID > 0
}
// MigrationFunc is a function that executes a migration against a database.
type MigrationFunc func(ctx context.Context, db *bun.DB) error
type internalMigrationFunc func(ctx context.Context, migrator *Migrator, migration *Migration) error
func wrapGoMigrationFunc(fn MigrationFunc) internalMigrationFunc {
return func(ctx context.Context, migrator *Migrator, migration *Migration) error {
if migrator.beforeMigrationHook != nil {
if err := migrator.beforeMigrationHook(ctx, migrator.db, migration); err != nil {
return err
}
}
if err := fn(ctx, migrator.db); err != nil {
return err
}
if migrator.afterMigrationHook != nil {
if err := migrator.afterMigrationHook(ctx, migrator.db, migration); err != nil {
return err
}
}
return nil
}
}
func newSQLMigrationFunc(fsys fs.FS, name string) internalMigrationFunc {
return func(ctx context.Context, migrator *Migrator, migration *Migration) error {
sqlFile, err := fsys.Open(name)
if err != nil {
return err
}
contents, err := io.ReadAll(sqlFile)
if err != nil {
return err
}
var reader io.Reader = bytes.NewReader(contents)
if migrator.templateData != nil {
buf, err := renderTemplate(contents, migrator.templateData)
if err != nil {
return err
}
reader = buf
}
scanner := bufio.NewScanner(reader)
var queries []string
var query []byte
for scanner.Scan() {
b := scanner.Bytes()
const prefix = "--bun:"
if bytes.HasPrefix(b, []byte(prefix)) {
b = b[len(prefix):]
if bytes.Equal(b, []byte("split")) {
queries = append(queries, string(query))
query = query[:0]
continue
}
return fmt.Errorf("bun: unknown directive: %q", b)
}
query = append(query, b...)
query = append(query, '\n')
}
if len(query) > 0 {
queries = append(queries, string(query))
}
if err := scanner.Err(); err != nil {
return err
}
var idb bun.IConn
isTx := strings.HasSuffix(name, ".tx.up.sql") || strings.HasSuffix(name, ".tx.down.sql")
if isTx {
tx, err := migrator.db.BeginTx(ctx, nil)
if err != nil {
return err
}
idb = tx
} else {
conn, err := migrator.db.Conn(ctx)
if err != nil {
return err
}
idb = conn
}
var retErr error
var execErr error
defer func() {
if tx, ok := idb.(bun.Tx); ok {
if execErr != nil {
retErr = tx.Rollback()
} else {
retErr = tx.Commit()
}
return
}
if conn, ok := idb.(bun.Conn); ok {
retErr = conn.Close()
return
}
panic("not reached")
}()
execErr = migrator.exec(ctx, idb, migration, queries)
if execErr != nil {
return execErr
}
return retErr
}
}
func renderTemplate(contents []byte, templateData any) (*bytes.Buffer, error) {
tmpl, err := template.New("migration").Parse(string(contents))
if err != nil {
return nil, fmt.Errorf("failed to parse template: %w", err)
}
var rendered bytes.Buffer
if err := tmpl.Execute(&rendered, templateData); err != nil {
return nil, fmt.Errorf("failed to execute template: %w", err)
}
return &rendered, nil
}
const goTemplate = `package %s
import (
"context"
"fmt"
"github.com/uptrace/bun"
)
func init() {
Migrations.MustRegister(func(ctx context.Context, db *bun.DB) error {
fmt.Print(" [up migration] ")
return nil
}, func(ctx context.Context, db *bun.DB) error {
fmt.Print(" [down migration] ")
return nil
})
}
`
const sqlTemplate = `SET statement_timeout = 0;
--bun:split
SELECT 1
--bun:split
SELECT 2
`
const transactionalSQLTemplate = `SET statement_timeout = 0;
SELECT 1;
`
//------------------------------------------------------------------------------
// MigrationSlice is a slice of migrations that provides helper methods for filtering and grouping.
type MigrationSlice []Migration
func (ms MigrationSlice) String() string {
if len(ms) == 0 {
return "empty"
}
if len(ms) > 5 {
return fmt.Sprintf("%d migrations (%s ... %s)", len(ms), ms[0].Name, ms[len(ms)-1].Name)
}
var sb strings.Builder
for i := range ms {
if i > 0 {
sb.WriteString(", ")
}
sb.WriteString(ms[i].String())
}
return sb.String()
}
// Applied returns applied migrations in descending order
// (the order is important and is used in Rollback).
func (ms MigrationSlice) Applied() MigrationSlice {
var applied MigrationSlice
for i := range ms {
if ms[i].IsApplied() {
applied = append(applied, ms[i])
}
}
sortDesc(applied)
return applied
}
// Unapplied returns unapplied migrations in ascending order
// (the order is important and is used in Migrate).
func (ms MigrationSlice) Unapplied() MigrationSlice {
var unapplied MigrationSlice
for i := range ms {
if !ms[i].IsApplied() {
unapplied = append(unapplied, ms[i])
}
}
sortAsc(unapplied)
return unapplied
}
// LastGroupID returns the last applied migration group id.
// The id is 0 when there are no migration groups.
func (ms MigrationSlice) LastGroupID() int64 {
var lastGroupID int64
for i := range ms {
groupID := ms[i].GroupID
if groupID > lastGroupID {
lastGroupID = groupID
}
}
return lastGroupID
}
// LastGroup returns the last applied migration group.
func (ms MigrationSlice) LastGroup() *MigrationGroup {
group := &MigrationGroup{
ID: ms.LastGroupID(),
}
if group.ID == 0 {
return group
}
for i := range ms {
if ms[i].GroupID == group.ID {
group.Migrations = append(group.Migrations, ms[i])
}
}
return group
}
// MigrationGroup is a group of migrations that were applied together in a single Migrate call.
type MigrationGroup struct {
ID int64
Migrations MigrationSlice
}
// IsZero reports whether the group is empty.
func (g MigrationGroup) IsZero() bool {
return g.ID == 0 && len(g.Migrations) == 0
}
func (g MigrationGroup) String() string {
if g.IsZero() {
return "nil"
}
return fmt.Sprintf("group #%d (%s)", g.ID, g.Migrations)
}
// MigrationFile represents a generated migration file on disk.
type MigrationFile struct {
Name string
Path string
Content string
}
//------------------------------------------------------------------------------
type migrationConfig struct {
nop bool
}
func newMigrationConfig(opts []MigrationOption) *migrationConfig {
cfg := new(migrationConfig)
for _, opt := range opts {
opt(cfg)
}
return cfg
}
// MigrationOption configures how a migration is executed.
type MigrationOption func(cfg *migrationConfig)
// WithNopMigration creates a no-op migration that marks itself as applied without running.
func WithNopMigration() MigrationOption {
return func(cfg *migrationConfig) {
cfg.nop = true
}
}
//------------------------------------------------------------------------------
func sortAsc(ms MigrationSlice) {
slices.SortFunc(ms, func(a, b Migration) int {
return strings.Compare(a.Name, b.Name)
})
}
func sortDesc(ms MigrationSlice) {
slices.SortFunc(ms, func(a, b Migration) int {
return strings.Compare(b.Name, a.Name)
})
}
+177
View File
@@ -0,0 +1,177 @@
package migrate
import (
"errors"
"fmt"
"io/fs"
"os"
"path/filepath"
"regexp"
"runtime"
"strings"
)
// MigrationsOption configures a Migrations instance.
type MigrationsOption func(m *Migrations)
// WithMigrationsDirectory sets the directory where migration files are stored.
func WithMigrationsDirectory(directory string) MigrationsOption {
return func(m *Migrations) {
m.explicitDirectory = directory
}
}
// Migrations is a collection of registered migrations.
type Migrations struct {
ms MigrationSlice
explicitDirectory string
implicitDirectory string
}
// NewMigrations creates a new collection of migrations.
func NewMigrations(opts ...MigrationsOption) *Migrations {
m := new(Migrations)
for _, opt := range opts {
opt(m)
}
m.implicitDirectory = filepath.Dir(migrationFile())
return m
}
// Sorted returns a copy of the migrations sorted by name in ascending order.
func (m *Migrations) Sorted() MigrationSlice {
migrations := make(MigrationSlice, len(m.ms))
copy(migrations, m.ms)
sortAsc(migrations)
return migrations
}
// MustRegister is like Register but panics on error.
func (m *Migrations) MustRegister(up, down MigrationFunc) {
if err := m.Register(up, down); err != nil {
panic(err)
}
}
// Register registers up and down migration functions derived from the caller's file name.
func (m *Migrations) Register(up, down MigrationFunc) error {
fpath := migrationFile()
name, comment, err := extractMigrationName(fpath)
if err != nil {
return err
}
m.Add(Migration{
Name: name,
Comment: comment,
Up: wrapGoMigrationFunc(up),
Down: wrapGoMigrationFunc(down),
})
return nil
}
// Add appends a migration to the collection.
func (m *Migrations) Add(migration Migration) {
if migration.Name == "" {
panic("migration name is required")
}
m.ms = append(m.ms, migration)
}
// DiscoverCaller discovers SQL migration files in the caller's directory.
func (m *Migrations) DiscoverCaller() error {
dir := filepath.Dir(migrationFile())
return m.Discover(os.DirFS(dir))
}
// Discover discovers SQL migration files in the given filesystem.
func (m *Migrations) Discover(fsys fs.FS) error {
return fs.WalkDir(fsys, ".", func(path string, d fs.DirEntry, err error) error {
if err != nil {
return err
}
if d.IsDir() {
return nil
}
if !strings.HasSuffix(path, ".up.sql") && !strings.HasSuffix(path, ".down.sql") {
return nil
}
name, comment, err := extractMigrationName(path)
if err != nil {
return err
}
migration := m.getOrCreateMigration(name)
migration.Comment = comment
migrationFunc := newSQLMigrationFunc(fsys, path)
if strings.HasSuffix(path, ".up.sql") {
migration.Up = migrationFunc
return nil
}
if strings.HasSuffix(path, ".down.sql") {
migration.Down = migrationFunc
return nil
}
return errors.New("migrate: not reached")
})
}
func (m *Migrations) getOrCreateMigration(name string) *Migration {
for i := range m.ms {
mig := &m.ms[i]
if mig.Name == name {
return mig
}
}
m.ms = append(m.ms, Migration{Name: name})
return &m.ms[len(m.ms)-1]
}
func (m *Migrations) getDirectory() string {
if m.explicitDirectory != "" {
return m.explicitDirectory
}
if m.implicitDirectory != "" {
return m.implicitDirectory
}
return filepath.Dir(migrationFile())
}
func migrationFile() string {
const depth = 32
var pcs [depth]uintptr
n := runtime.Callers(1, pcs[:])
frames := runtime.CallersFrames(pcs[:n])
for {
f, ok := frames.Next()
if !ok {
break
}
if !strings.Contains(f.Function, "/bun/migrate.") {
return f.File
}
}
return ""
}
var fnameRE = regexp.MustCompile(`^(\d{1,14})_([0-9a-z_\-]+)\.`)
func extractMigrationName(fpath string) (string, string, error) {
fname := filepath.Base(fpath)
matches := fnameRE.FindStringSubmatch(fname)
if matches == nil {
return "", "", fmt.Errorf("migrate: unsupported migration name format: %q", fname)
}
return matches[1], matches[2], nil
}

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