diff --git a/driver/pgdriver/column.go b/driver/pgdriver/column.go index c964bc0f9..fe6f37a9f 100644 --- a/driver/pgdriver/column.go +++ b/driver/pgdriver/column.go @@ -31,7 +31,12 @@ const ( ) func readColumnValue(rd *reader, dataType int32, dataLen int) (any, error) { - if dataLen == -1 { + if dataLen < 0 { + if dataLen == -1 { + return nil, nil + } + return nil, errInvalidMessageLength // TODO or nil? + } return nil, nil } @@ -61,6 +66,13 @@ func readColumnValue(rd *reader, dataType int32, dataLen int) (any, error) { return readBytesCol(rd, dataLen) } + if dataLen < 0 { + // dataLen == -1 signals SQL NULL; any other negative value is invalid + // and must not reach make([]byte, dataLen) (would panic). dataLen == 0 + // is fine and keeps its previous behavior (empty value). + return nil, nil + } + b := make([]byte, dataLen) if _, err := io.ReadFull(rd, b); err != nil { return nil, err diff --git a/driver/pgdriver/driver.go b/driver/pgdriver/driver.go index ef04b00b5..a2efff737 100644 --- a/driver/pgdriver/driver.go +++ b/driver/pgdriver/driver.go @@ -174,6 +174,14 @@ func (cn *Conn) reader(ctx context.Context, timeout time.Duration) *reader { } func (cn *Conn) write(ctx context.Context, wb *writeBuffer) error { + // Refuse to send a message that failed to build (e.g. it exceeded the + // protocol size limit); sending it would desync the connection. + if wb.err != nil { + err := wb.err + wb.Reset() + return err + } + cn.setWriteDeadline(ctx, -1) n, err := cn.netConn.Write(wb.Bytes) diff --git a/driver/pgdriver/error.go b/driver/pgdriver/error.go index 38a6708ff..efbad6060 100644 --- a/driver/pgdriver/error.go +++ b/driver/pgdriver/error.go @@ -84,4 +84,15 @@ func isBadConn(err error, allowTimeout bool) bool { var ( errRequiresParameter = errors.New("pgdriver: requires that standard_conforming_strings=on and client_encoding=UTF8") + + // errMessageTooLarge is returned instead of sending a message whose length + // prefix would overflow the protocol's 32-bit size field. Without this guard + // the length wraps and the server desyncs, which can be abused for SQL + // injection (cf. jackc/pgx CVE-2024-27304). + errMessageTooLarge = errors.New("pgdriver: message size exceeds the maximum of 4 GiB") + + // errInvalidMessageLength is returned when the server sends a message whose + // length field is smaller than the 4-byte length prefix itself (which would + // otherwise lead to a negative-length slice/allocation panic). + errInvalidMessageLength = errors.New("pgdriver: server sent a message with an invalid length") ) diff --git a/driver/pgdriver/message_size_test.go b/driver/pgdriver/message_size_test.go new file mode 100644 index 000000000..ff948f138 --- /dev/null +++ b/driver/pgdriver/message_size_test.go @@ -0,0 +1,84 @@ +package pgdriver + +import ( + "bytes" + "context" + "database/sql/driver" + "encoding/binary" + "errors" + "math" + "testing" +) + +func TestFinishMessage_TooLargeSetsError(t *testing.T) { + orig := maxMessageSize + maxMessageSize = 8 + defer func() { maxMessageSize = orig }() + + wb := newWriteBuffer() + wb.StartMessage(queryMsg) + wb.WriteString("some query that exceeds the tiny test limit") + wb.FinishMessage() + + if !errors.Is(wb.err, errMessageTooLarge) { + t.Fatalf("wb.err = %v, want errMessageTooLarge", wb.err) + } +} + +func TestConnWrite_RefusesOversizedMessage(t *testing.T) { + wb := newWriteBuffer() + wb.setErr(errMessageTooLarge) + + // wb.err is set, so write must return it without touching the (nil) netConn. + if err := (&Conn{}).write(context.Background(), wb); !errors.Is(err, errMessageTooLarge) { + t.Fatalf("write() = %v, want errMessageTooLarge", err) + } +} + +func TestFinishParam_TooLargeSetsError(t *testing.T) { + orig := maxMessageSize + maxMessageSize = 4 + defer func() { maxMessageSize = orig }() + + wb := newWriteBuffer() + wb.StartParam() + wb.Write([]byte("larger than the limit")) + wb.FinishParam() + + if !errors.Is(wb.err, errMessageTooLarge) { + t.Fatalf("wb.err = %v, want errMessageTooLarge", wb.err) + } +} + +func TestWriteBindExecute_TooManyParams(t *testing.T) { + args := make([]driver.NamedValue, math.MaxInt16+1) + err := writeBindExecute(context.Background(), &Conn{}, "", args) + if err == nil { + t.Fatal("writeBindExecute with > MaxInt16 params: got nil error, want an error") + } +} + +func TestReadMessageType_RejectsShortLength(t *testing.T) { + // type byte + 4-byte length field of 0 (< 4, invalid). + buf := []byte{'X', 0, 0, 0, 0} + rd := newReader(bytes.NewReader(buf), 1024) + if _, _, err := readMessageType(rd); !errors.Is(err, errInvalidMessageLength) { + t.Fatalf("readMessageType short length: err = %v, want errInvalidMessageLength", err) + } + + // A valid length (>= 4) is accepted and the body length is length-4. + valid := []byte{'X', 0, 0, 0, 0} + binary.BigEndian.PutUint32(valid[1:], 10) + rd = newReader(bytes.NewReader(valid), 1024) + c, n, err := readMessageType(rd) + if err != nil || c != 'X' || n != 6 { + t.Fatalf("readMessageType valid: (%q, %d, %v), want ('X', 6, nil)", c, n, err) + } +} + +func TestReadTemp_RejectsNegative(t *testing.T) { + rd := newReader(bytes.NewReader(nil), 1024) + if _, err := rd.ReadTemp(-1); !errors.Is(err, errInvalidMessageLength) { + t.Fatalf("ReadTemp(-1) = %v, want errInvalidMessageLength", err) + } +} diff --git a/driver/pgdriver/proto.go b/driver/pgdriver/proto.go index cf573b2fe..b59e444d1 100644 --- a/driver/pgdriver/proto.go +++ b/driver/pgdriver/proto.go @@ -94,6 +94,9 @@ func newReader(r io.Reader, size int) *reader { } func (r *reader) ReadTemp(n int) ([]byte, error) { + if n < 0 { + return nil, errInvalidMessageLength + } if n <= len(r.buf) { b := r.buf[:n] _, err := io.ReadFull(r.Reader, b) @@ -827,6 +830,10 @@ func writeBindExecute(ctx context.Context, cn *Conn, name string, args []driver. wb := getWriteBuffer() defer putWriteBuffer(wb) + if len(args) > math.MaxInt16 { + return fmt.Errorf("pgdriver: too many bind parameters: %d (max %d)", len(args), math.MaxInt16) + } + wb.StartMessage(bindMsg) wb.WriteString("") wb.WriteString(name) @@ -1020,6 +1027,12 @@ func readMessageType(rd *reader) (byte, int, error) { if err != nil { return 0, 0, err } + if l < 4 { + // The length field includes its own 4 bytes, so it can never be < 4. + // A smaller value would yield a negative body length and panic callers + // that allocate/slice with it. + return 0, 0, errInvalidMessageLength + } return c, int(l) - 4, nil } diff --git a/driver/pgdriver/write_buffer.go b/driver/pgdriver/write_buffer.go index 1e7489fe7..a4a926bd8 100644 --- a/driver/pgdriver/write_buffer.go +++ b/driver/pgdriver/write_buffer.go @@ -3,9 +3,16 @@ package pgdriver import ( "encoding/binary" "io" + "math" "sync" ) +// maxMessageSize is the largest value that fits in the protocol's 32-bit message +// length field. It is an int64 (so the constant does not overflow int on 32-bit +// platforms) and a var (not a const) so tests can lower it. Do not mutate it +// outside of tests. +var maxMessageSize int64 = math.MaxUint32 + var wbPool = sync.Pool{ New: func() any { return newWriteBuffer() @@ -27,6 +34,11 @@ type writeBuffer struct { msgStart int paramStart int + + // err records the first fatal error (e.g. a message that exceeds the + // protocol size limit). It is checked before the buffer is sent so an + // oversized/overflowing message is never written to the wire. + err error } func newWriteBuffer() *writeBuffer { @@ -37,6 +49,14 @@ func newWriteBuffer() *writeBuffer { func (b *writeBuffer) Reset() { b.Bytes = b.Bytes[:0] + b.err = nil +} + +// setErr records the first error seen while building a message. +func (b *writeBuffer) setErr(err error) { + if b.err == nil { + b.err = err + } } func (b *writeBuffer) StartMessage(c byte) { @@ -50,8 +70,12 @@ func (b *writeBuffer) StartMessage(c byte) { } func (b *writeBuffer) FinishMessage() { - binary.BigEndian.PutUint32( - b.Bytes[b.msgStart:], uint32(len(b.Bytes)-b.msgStart)) + n := len(b.Bytes) - b.msgStart + if int64(n) > maxMessageSize { + b.setErr(errMessageTooLarge) + return + } + binary.BigEndian.PutUint32(b.Bytes[b.msgStart:], uint32(n)) } func (b *writeBuffer) Query() []byte { @@ -64,8 +88,12 @@ func (b *writeBuffer) StartParam() { } func (b *writeBuffer) FinishParam() { - binary.BigEndian.PutUint32( - b.Bytes[b.paramStart:], uint32(len(b.Bytes)-b.paramStart-4)) + n := len(b.Bytes) - b.paramStart - 4 + if int64(n) > maxMessageSize { + b.setErr(errMessageTooLarge) + return + } + binary.BigEndian.PutUint32(b.Bytes[b.paramStart:], uint32(n)) } var nullParamLength = int32(-1)