This is an automated email from the ASF dual-hosted git repository.

zeroshade pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/arrow-go.git


The following commit(s) were added to refs/heads/main by this push:
     new e25ccd31 fix(parquet): reject truncated scalar bit reads (#944)
e25ccd31 is described below

commit e25ccd31e1175889681557761136eb8157854c04
Author: Minh Vu <[email protected]>
AuthorDate: Thu Jul 23 21:41:55 2026 +0200

    fix(parquet): reject truncated scalar bit reads (#944)
    
    BitReader.fillbuffer padded short reads with zeros without tracking how 
many bits were valid, so scalar and trailing reads could treat missing input as 
encoded zero values.
    
    - Track valid bits per buffer and reject scalar reads that extend beyond 
them.
    - Keep boolean batches on the bulk path, checking availability once per 
byte.
    - Decode scalar dictionary indexes in validated buffer-sized chunks.
    - Propagate seek and refill failures, and validate bulk discards.
    - Preserve zero-width reads without requiring input.
    
    Truncated scalar bit streams now return short-input errors rather than 
decoding missing data as zeros.
---
 parquet/internal/utils/bit_reader.go      | 189 +++++++++++++++++++++++-------
 parquet/internal/utils/bit_reader_test.go | 108 +++++++++++++++++
 2 files changed, 255 insertions(+), 42 deletions(-)

diff --git a/parquet/internal/utils/bit_reader.go 
b/parquet/internal/utils/bit_reader.go
index c7e2e79c..862394cc 100644
--- a/parquet/internal/utils/bit_reader.go
+++ b/parquet/internal/utils/bit_reader.go
@@ -69,6 +69,7 @@ type BitReader struct {
        buffer     uint64
        byteoffset int64
        bitoffset  uint
+       validBits  uint
        raw        [8]byte
 
        unpackBuf [buflen]uint32
@@ -92,6 +93,7 @@ func (b *BitReader) Reset(r reader) {
        b.buffer = 0
        b.byteoffset = 0
        b.bitoffset = 0
+       b.validBits = 0
 }
 
 // GetVlqInt reads a Vlq encoded int from the stream. The encoded value must 
start
@@ -150,8 +152,7 @@ func (b *BitReader) getAlignedUint8(nbytes int, v *uint8) 
bool {
 
        b.byteoffset += int64(nbytes)
        b.bitoffset = 0
-       b.fillbuffer()
-       return true
+       return b.fillbuffer() == nil
 }
 
 // getAlignedUint16 reads nbytes from the underlying stream into the passed 
uint16 value.
@@ -177,8 +178,7 @@ func (b *BitReader) getAlignedUint16(nbytes int, v *uint16) 
bool {
 
        b.byteoffset += int64(nbytes)
        b.bitoffset = 0
-       b.fillbuffer()
-       return true
+       return b.fillbuffer() == nil
 }
 
 // getAlignedUint32 reads nbytes from the underlying stream into the passed 
uint32 value.
@@ -204,8 +204,7 @@ func (b *BitReader) getAlignedUint32(nbytes int, v *uint32) 
bool {
 
        b.byteoffset += int64(nbytes)
        b.bitoffset = 0
-       b.fillbuffer()
-       return true
+       return b.fillbuffer() == nil
 }
 
 // getAlignedUint64 reads nbytes from the underlying stream into the passed 
uint64 value.
@@ -231,16 +230,16 @@ func (b *BitReader) getAlignedUint64(nbytes int, v 
*uint64) bool {
 
        b.byteoffset += int64(nbytes)
        b.bitoffset = 0
-       b.fillbuffer()
-       return true
+       return b.fillbuffer() == nil
 }
 
 // fillbuffer fills the uint64 buffer with bytes from the underlying stream
 func (b *BitReader) fillbuffer() error {
        n, err := b.reader.ReadAt(b.raw[:], b.byteoffset)
-       if err != nil && n == 0 && err != io.EOF {
+       if err != nil && err != io.EOF {
                return err
        }
+       b.validBits = uint(n * 8)
        for i := n; i < 8; i++ {
                b.raw[i] = 0
        }
@@ -250,18 +249,103 @@ func (b *BitReader) fillbuffer() error {
 
 // next reads an integral value from the next bits in the buffer
 func (b *BitReader) next(bits uint) (v uint64, err error) {
-       v = trailingBits(b.buffer, b.bitoffset+bits) >> b.bitoffset
-       b.bitoffset += bits
-       // if we need more bits to get what was requested then refill the buffer
-       if b.bitoffset >= 64 {
+       if bits == 0 {
+               return 0, nil
+       }
+
+       if b.bitoffset == 64 {
                b.byteoffset += 8
-               b.bitoffset -= 64
+               b.bitoffset = 0
                if err = b.fillbuffer(); err != nil {
                        return 0, err
                }
-               v |= trailingBits(b.buffer, b.bitoffset) << (bits - b.bitoffset)
        }
-       return
+
+       end := b.bitoffset + bits
+       if end <= 64 {
+               if end > b.validBits {
+                       return 0, io.ErrUnexpectedEOF
+               }
+               v = trailingBits(b.buffer, end) >> b.bitoffset
+               b.bitoffset = end
+               return v, nil
+       }
+
+       if b.validBits < 64 {
+               return 0, io.ErrUnexpectedEOF
+       }
+
+       firstBits := 64 - b.bitoffset
+       v = b.buffer >> b.bitoffset
+       b.byteoffset += 8
+       b.bitoffset = 0
+       if err = b.fillbuffer(); err != nil {
+               return 0, err
+       }
+
+       remaining := bits - firstBits
+       if remaining > b.validBits {
+               return 0, io.ErrUnexpectedEOF
+       }
+       v |= trailingBits(b.buffer, remaining) << firstBits
+       b.bitoffset = remaining
+       return v, nil
+}
+
+// nextBool reads one bit after the caller has verified that its byte is
+// available. GetBatchBools performs that check once per byte so its unaligned
+// and trailing reads do not pay the scalar validBits check for every value.
+func (b *BitReader) nextBool() (bool, error) {
+       if b.bitoffset == 64 {
+               b.byteoffset += 8
+               b.bitoffset = 0
+               if err := b.fillbuffer(); err != nil {
+                       return false, err
+               }
+       }
+
+       v := b.buffer&(uint64(1)<<b.bitoffset) != 0
+       b.bitoffset++
+       return v, nil
+}
+
+// getBatchIndexScalar decodes indexes from the buffered scalar path. It checks
+// availability once per buffer and decodes all complete values in that buffer
+// without repeating the validBits check for every index.
+func (b *BitReader) getBatchIndexScalar(bits uint, out []IndexType, 
stopAtBufferBoundary bool) (int, error) {
+       i := 0
+       for i < len(out) {
+               if stopAtBufferBoundary && (b.bitoffset == 0 || b.bitoffset == 
64) {
+                       if b.bitoffset == 64 {
+                               b.byteoffset += 8
+                               b.bitoffset = 0
+                       }
+                       return i, nil
+               }
+
+               if b.bitoffset < b.validBits {
+                       available := int((b.validBits - b.bitoffset) / bits)
+                       available = min(available, len(out)-i)
+                       for range available {
+                               end := b.bitoffset + bits
+                               out[i] = IndexType(trailingBits(b.buffer, end) 
>> b.bitoffset)
+                               b.bitoffset = end
+                               i++
+                       }
+                       if available > 0 {
+                               continue
+                       }
+               }
+
+               val, err := b.next(bits)
+               if err != nil {
+                       return i, err
+               }
+               out[i] = IndexType(val)
+               i++
+       }
+
+       return i, nil
 }
 
 // GetBatchIndex is like GetBatch but for IndexType (used for dictionary 
decoding)
@@ -271,20 +355,28 @@ func (b *BitReader) GetBatchIndex(bits uint, out 
[]IndexType) (i int, err error)
        if bits > 32 {
                return 0, errors.New("must be 32 bits or less per read")
        }
-
-       var val uint64
+       if bits == 0 {
+               clear(out)
+               return len(out), nil
+       }
 
        length := len(out)
-       // if we aren't currently byte-aligned, read bits until we are 
byte-aligned.
-       for ; i < length && b.bitoffset != 0; i++ {
-               val, err = b.next(bits)
-               out[i] = IndexType(val)
+       // If the buffer is partially consumed, read indexes until the next 
buffer boundary.
+       if b.bitoffset != 0 {
+               var n int
+               n, err = b.getBatchIndexScalar(bits, out, true)
+               i += n
                if err != nil {
-                       return
+                       return i, err
+               }
+               if i == length {
+                       return i, nil
                }
        }
 
-       b.reader.Seek(b.byteoffset, io.SeekStart)
+       if _, err = b.reader.Seek(b.byteoffset, io.SeekStart); err != nil {
+               return i, err
+       }
        // grab as many 32 byte chunks as possible in one shot
        if i < length { // IndexType should be a 32 bit value so we can do 
quick unpacking right into the output
                numUnpacked, unpackErr := unpack32(b.reader, 
(*(*[]uint32)(unsafe.Pointer(&out)))[i:], int(bits))
@@ -296,21 +388,17 @@ func (b *BitReader) GetBatchIndex(bits uint, out 
[]IndexType) (i int, err error)
        }
 
        // re-fill our buffer just in case.
-       b.fillbuffer()
-       // grab the remaining values that aren't 32 byte aligned
-       for ; i < length; i++ {
-               val, err = b.next(bits)
-               out[i] = IndexType(val)
-               if err != nil {
-                       break
-               }
+       if err = b.fillbuffer(); err != nil {
+               return i, err
        }
-       return
+       // grab the remaining values that aren't 32 byte aligned
+       n, err := b.getBatchIndexScalar(bits, out[i:], false)
+       i += n
+       return i, err
 }
 
 // GetBatchBools is like GetBatch but optimized for reading bits as boolean 
values
 func (b *BitReader) GetBatchBools(out []bool) (int, error) {
-       bits := uint(1)
        length := len(out)
 
        i := 0
@@ -321,14 +409,16 @@ func (b *BitReader) GetBatchBools(out []bool) (int, 
error) {
                                return i, err
                        }
                }
-               val, err := b.next(bits)
-               out[i] = val != 0
+               val, err := b.nextBool()
+               out[i] = val
                if err != nil {
                        return i, err
                }
        }
 
-       b.reader.Seek(b.byteoffset, io.SeekStart)
+       if _, err := b.reader.Seek(b.byteoffset, io.SeekStart); err != nil {
+               return i, err
+       }
        buf := arrow.Uint32Traits.CastToBytes(b.unpackBuf[:])
        blen := buflen * 8
        for length-i >= 8 {
@@ -378,8 +468,8 @@ func (b *BitReader) GetBatchBools(out []bool) (int, error) {
                                return i, err
                        }
                }
-               val, err := b.next(bits)
-               out[i] = val != 0
+               val, err := b.nextBool()
+               out[i] = val
                if err != nil {
                        return i, err
                }
@@ -419,11 +509,22 @@ func (b *BitReader) Discard(bits uint, n int) (int, 
error) {
                toSkip := (n - i) / 32 * 32
 
                bytesToSkip := bitutil.BytesForBits(int64(toSkip * int(bits)))
+               if bytesToSkip > 0 {
+                       var last [1]byte
+                       if nread, err := b.reader.ReadAt(last[:], 
b.byteoffset+int64(bytesToSkip)-1); nread != 1 {
+                               if err == nil {
+                                       err = io.ErrUnexpectedEOF
+                               }
+                               return i, err
+                       }
+               }
                b.byteoffset += int64(bytesToSkip)
                i += toSkip
        }
 
-       b.fillbuffer()
+       if err := b.fillbuffer(); err != nil {
+               return i, err
+       }
        for ; i < n; i++ {
                if _, err := b.next(bits); err != nil {
                        return i, err
@@ -454,7 +555,9 @@ func (b *BitReader) GetBatch(bits uint, out []uint64) (int, 
error) {
                }
        }
 
-       b.reader.Seek(b.byteoffset, io.SeekStart)
+       if _, err := b.reader.Seek(b.byteoffset, io.SeekStart); err != nil {
+               return i, err
+       }
        for i < length {
                // unpack groups of 32 bytes at a time into a buffer since it's 
more efficient
                unpackSize := utils.Min(buflen, length-i)
@@ -473,7 +576,9 @@ func (b *BitReader) GetBatch(bits uint, out []uint64) (int, 
error) {
                }
        }
 
-       b.fillbuffer()
+       if err := b.fillbuffer(); err != nil {
+               return i, err
+       }
        // and then the remaining trailing values
        for ; i < length; i++ {
                val, err := b.next(bits)
diff --git a/parquet/internal/utils/bit_reader_test.go 
b/parquet/internal/utils/bit_reader_test.go
index 9df2f617..502bbbfb 100644
--- a/parquet/internal/utils/bit_reader_test.go
+++ b/parquet/internal/utils/bit_reader_test.go
@@ -171,6 +171,22 @@ func TestBitReaderGetBatchBools(t *testing.T) {
                }
        })
 
+       t.Run("unaligned across buffer", func(t *testing.T) {
+               data := bytes.Repeat([]byte{0xAA, 0xCC, 0xF0}, 4)
+               reader := utils.NewBitReader(bytes.NewReader(data))
+               _, ok := reader.GetValue(1)
+               assert.True(t, ok)
+
+               out := make([]bool, 80)
+               n, err := reader.GetBatchBools(out)
+               assert.NoError(t, err)
+               assert.Equal(t, len(out), n)
+               for i, got := range out {
+                       bit := i + 1
+                       assert.Equal(t, data[bit/8]&(1<<uint(bit%8)) != 0, got)
+               }
+       })
+
        t.Run("no progress", func(t *testing.T) {
                reader := 
utils.NewBitReader(&stalledReader{bytes.NewReader(nil)})
                n, err := reader.GetBatchBools(make([]bool, 8))
@@ -201,6 +217,98 @@ func TestBitReader(t *testing.T) {
        }
 }
 
+func TestBitReaderRejectsMissingBits(t *testing.T) {
+       t.Run("single value", func(t *testing.T) {
+               reader := utils.NewBitReader(bytes.NewReader(nil))
+               _, ok := reader.GetValue(1)
+               assert.False(t, ok)
+       })
+
+       t.Run("partial batch", func(t *testing.T) {
+               reader := utils.NewBitReader(bytes.NewReader([]byte{0xAB}))
+               out := make([]uint64, 2)
+               n, err := reader.GetBatch(8, out)
+               assert.ErrorIs(t, err, io.ErrUnexpectedEOF)
+               assert.Equal(t, 1, n)
+               assert.Equal(t, uint64(0xAB), out[0])
+       })
+
+       t.Run("crosses buffer", func(t *testing.T) {
+               reader := 
utils.NewBitReader(bytes.NewReader(bytes.Repeat([]byte{0xFF}, 8)))
+               out := make([]uint64, 2)
+               n, err := reader.GetBatch(63, out)
+               assert.ErrorIs(t, err, io.ErrUnexpectedEOF)
+               assert.Equal(t, 1, n)
+       })
+
+       t.Run("discard", func(t *testing.T) {
+               reader := utils.NewBitReader(bytes.NewReader([]byte{0xFF}))
+               n, err := reader.Discard(1, 64)
+               assert.Error(t, err)
+               assert.Zero(t, n)
+       })
+}
+
+func TestBitReaderAvailableBits(t *testing.T) {
+       for width := 1; width <= 64; width++ {
+               for byteLen := 0; byteLen <= 16; byteLen++ {
+                       t.Run(fmt.Sprintf("width=%d/bytes=%d", width, byteLen), 
func(t *testing.T) {
+                               reader := 
utils.NewBitReader(bytes.NewReader(bytes.Repeat([]byte{0xFF}, byteLen)))
+                               out := make([]uint64, 31)
+                               n, err := reader.GetBatch(uint(width), out)
+
+                               expected := min(len(out), byteLen*8/width)
+                               assert.Equal(t, expected, n)
+                               if expected < len(out) {
+                                       assert.ErrorIs(t, err, 
io.ErrUnexpectedEOF)
+                               } else {
+                                       assert.NoError(t, err)
+                               }
+
+                               want := uint64(math.MaxUint64)
+                               if width < 64 {
+                                       want = 1<<uint(width) - 1
+                               }
+                               for _, got := range out[:n] {
+                                       assert.Equal(t, want, got)
+                               }
+
+                               if width <= 32 {
+                                       reader = 
utils.NewBitReader(bytes.NewReader(bytes.Repeat([]byte{0xFF}, byteLen)))
+                                       indices := make([]utils.IndexType, 
len(out))
+                                       for i := range indices {
+                                               indices[i] = math.MinInt32
+                                       }
+                                       n, err = 
reader.GetBatchIndex(uint(width), indices)
+                                       assert.Equal(t, expected, n)
+                                       if expected < len(indices) {
+                                               assert.ErrorIs(t, err, 
io.ErrUnexpectedEOF)
+                                       } else {
+                                               assert.NoError(t, err)
+                                       }
+                                       for _, got := range indices[:n] {
+                                               assert.Equal(t, int32(want), 
got)
+                                       }
+                                       for _, got := range indices[n:] {
+                                               assert.Equal(t, 
int32(math.MinInt32), got)
+                                       }
+                               }
+                       })
+               }
+       }
+
+       reader := utils.NewBitReader(bytes.NewReader(nil))
+       out := make([]uint64, 31)
+       n, err := reader.GetBatch(0, out)
+       assert.NoError(t, err)
+       assert.Equal(t, len(out), n)
+
+       indices := make([]utils.IndexType, len(out))
+       n, err = reader.GetBatchIndex(0, indices)
+       assert.NoError(t, err)
+       assert.Equal(t, len(indices), n)
+}
+
 func TestBitArrayVals(t *testing.T) {
        tests := []struct {
                name  string

Reply via email to