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
This commit is contained in:
+217
@@ -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))
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user