Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
14 changes: 2 additions & 12 deletions pgtype/array_codec.go
Original file line number Diff line number Diff line change
Expand Up @@ -131,7 +131,7 @@ func (p *encodePlanArrayCodecText) Encode(value any, buf []byte) (newBuf []byte,
elem := array.Index(i)
var elemBuf []byte
isNil, callNilDriverValuer := isNilDriverValuer(elem)
if !isNil {
if !isNil || callNilDriverValuer {
elemType := reflect.TypeOf(elem)
if lastElemType != elemType {
lastElemType = elemType
Expand All @@ -144,11 +144,6 @@ func (p *encodePlanArrayCodecText) Encode(value any, buf []byte) (newBuf []byte,
if err != nil {
return nil, err
}
} else if callNilDriverValuer {
elemBuf, err = (&encodePlanDriverValuer{m: p.m, oid: p.ac.ElementType.OID, formatCode: TextFormatCode}).Encode(elem, inElemBuf)
if err != nil {
return nil, err
}
}

if elemBuf == nil {
Expand Down Expand Up @@ -201,7 +196,7 @@ func (p *encodePlanArrayCodecBinary) Encode(value any, buf []byte) (newBuf []byt
elem := array.Index(i)
var elemBuf []byte
isNil, callNilDriverValuer := isNilDriverValuer(elem)
if !isNil {
if !isNil || callNilDriverValuer {
elemType := reflect.TypeOf(elem)
if lastElemType != elemType {
lastElemType = elemType
Expand All @@ -214,11 +209,6 @@ func (p *encodePlanArrayCodecBinary) Encode(value any, buf []byte) (newBuf []byt
if err != nil {
return nil, err
}
} else if callNilDriverValuer {
elemBuf, err = (&encodePlanDriverValuer{m: p.m, oid: p.ac.ElementType.OID, formatCode: BinaryFormatCode}).Encode(elem, buf)
if err != nil {
return nil, err
}
}

if elemBuf == nil {
Expand Down
88 changes: 88 additions & 0 deletions pgtype/array_codec_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@ package pgtype_test
import (
"context"
"database/sql/driver"
"encoding/binary"
"encoding/hex"
"reflect"
"strings"
Expand Down Expand Up @@ -147,6 +148,93 @@ func (v jsonbValuerSlice) Value() (driver.Value, error) {
return []byte("[1]"), nil
}

const (
codecValuerUUIDOID uint32 = 910001
codecValuerUUIDArrayOID uint32 = 910002
codecValuerCompositeOID uint32 = 910003
)

type codecValuerUUID []byte

func (v codecValuerUUID) Value() (driver.Value, error) {
if v == nil {
return "", nil
}
return "driver-valuer", nil
}

type codecValuerUUIDCodec struct{}

func (codecValuerUUIDCodec) FormatSupported(format int16) bool {
return format == pgtype.TextFormatCode || format == pgtype.BinaryFormatCode
}

func (codecValuerUUIDCodec) PreferredFormat() int16 {
return pgtype.BinaryFormatCode
}

func (codecValuerUUIDCodec) PlanEncode(m *pgtype.Map, oid uint32, format int16, value any) pgtype.EncodePlan {
if _, ok := value.(codecValuerUUID); !ok {
return nil
}

return encodePlanCodecValuerUUID{}
}

func (codecValuerUUIDCodec) PlanScan(m *pgtype.Map, oid uint32, format int16, target any) pgtype.ScanPlan {
return nil
}

func (codecValuerUUIDCodec) DecodeDatabaseSQLValue(m *pgtype.Map, oid uint32, format int16, src []byte) (driver.Value, error) {
return nil, nil
}

func (codecValuerUUIDCodec) DecodeValue(m *pgtype.Map, oid uint32, format int16, src []byte) (any, error) {
return nil, nil
}

type encodePlanCodecValuerUUID struct{}

func (encodePlanCodecValuerUUID) Encode(value any, buf []byte) ([]byte, error) {
v := value.(codecValuerUUID)
if v == nil {
return nil, nil
}
return append(buf, v...), nil
}

func newCodecValuerTestMap() *pgtype.Map {
m := pgtype.NewMap()
elementType := &pgtype.Type{Name: "codec_valuer_uuid", OID: codecValuerUUIDOID, Codec: codecValuerUUIDCodec{}}
m.RegisterType(elementType)
m.RegisterType(&pgtype.Type{Name: "_codec_valuer_uuid", OID: codecValuerUUIDArrayOID, Codec: &pgtype.ArrayCodec{ElementType: elementType}})
m.RegisterType(&pgtype.Type{Name: "codec_valuer_composite", OID: codecValuerCompositeOID, Codec: &pgtype.CompositeCodec{Fields: []pgtype.CompositeCodecField{
{Name: "id", Type: elementType},
}}})
return m
}

func TestArrayCodecTypedNilElementUsesCodecBeforeDriverValuer(t *testing.T) {
m := newCodecValuerTestMap()
input := []codecValuerUUID{codecValuerUUID("codec-value"), nil}

textBuf, err := m.Encode(codecValuerUUIDArrayOID, pgtype.TextFormatCode, input, nil)
require.NoError(t, err)
require.Equal(t, `{codec-value,NULL}`, string(textBuf))

binaryBuf, err := m.Encode(codecValuerUUIDArrayOID, pgtype.BinaryFormatCode, input, nil)
require.NoError(t, err)
require.GreaterOrEqual(t, len(binaryBuf), 39)
require.Equal(t, int32(1), int32(binary.BigEndian.Uint32(binaryBuf[0:4])))
require.Equal(t, int32(1), int32(binary.BigEndian.Uint32(binaryBuf[4:8])))
require.Equal(t, codecValuerUUIDOID, binary.BigEndian.Uint32(binaryBuf[8:12]))
require.Equal(t, int32(2), int32(binary.BigEndian.Uint32(binaryBuf[12:16])))
require.Equal(t, int32(1), int32(binary.BigEndian.Uint32(binaryBuf[16:20])))
require.Equal(t, int32(len("codec-value")), int32(binary.BigEndian.Uint32(binaryBuf[20:24])))
require.Equal(t, "codec-value", string(binaryBuf[24:35]))
require.Equal(t, int32(-1), int32(binary.BigEndian.Uint32(binaryBuf[35:39])))
}

func TestArrayCodecTypedNilElementWithValuer(t *testing.T) {
pgxtest.RunWithQueryExecModes(context.Background(), t, defaultConnTestRunner, pgxtest.KnownOIDQueryExecModes, func(ctx context.Context, t testing.TB, conn *pgx.Conn) {
input := []jsonbValuerSlice{nil, nil}
Expand Down
26 changes: 8 additions & 18 deletions pgtype/composite.go
Original file line number Diff line number Diff line change
Expand Up @@ -492,15 +492,10 @@ func (b *CompositeBinaryBuilder) AppendValue(oid uint32, field any) {
return
}

var plan EncodePlan
if isNil {
plan = &encodePlanDriverValuer{m: b.m, oid: oid, formatCode: BinaryFormatCode}
} else {
plan = b.m.PlanEncode(oid, BinaryFormatCode, field)
if plan == nil {
b.err = fmt.Errorf("unable to encode %v into OID %d in binary format", field, oid)
return
}
plan := b.m.PlanEncode(oid, BinaryFormatCode, field)
if plan == nil {
b.err = fmt.Errorf("unable to encode %v into OID %d in binary format", field, oid)
return
}

b.buf = pgio.AppendUint32(b.buf, oid)
Expand Down Expand Up @@ -553,15 +548,10 @@ func (b *CompositeTextBuilder) AppendValue(oid uint32, field any) {
return
}

var plan EncodePlan
if isNil {
plan = &encodePlanDriverValuer{m: b.m, oid: oid, formatCode: TextFormatCode}
} else {
plan = b.m.PlanEncode(oid, TextFormatCode, field)
if plan == nil {
b.err = fmt.Errorf("unable to encode %v into OID %d in text format", field, oid)
return
}
plan := b.m.PlanEncode(oid, TextFormatCode, field)
if plan == nil {
b.err = fmt.Errorf("unable to encode %v into OID %d in text format", field, oid)
return
}

fieldBuf, err := plan.Encode(field, b.fieldBuf[0:0])
Expand Down
8 changes: 8 additions & 0 deletions pgtype/composite_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -110,6 +110,14 @@ create type ct_stringer_test as (
})
}

func TestCompositeCodecTypedNilFieldUsesCodecBeforeDriverValuer(t *testing.T) {
m := newCodecValuerTestMap()

buf, err := m.Encode(codecValuerCompositeOID, pgtype.TextFormatCode, pgtype.CompositeFields{codecValuerUUID(nil)}, nil)
require.NoError(t, err)
require.Equal(t, `()`, string(buf))
}

type point3d struct {
X, Y, Z float64
}
Expand Down