diff --git a/roaringarray.go b/roaringarray.go index dfe80c4c..a3d126e8 100644 --- a/roaringarray.go +++ b/roaringarray.go @@ -689,6 +689,22 @@ func (ra *roaringArray) toBytes() ([]byte, error) { return buf.Bytes(), err } +func (ra *roaringArray) setNoRunContainer(i uint32, card int, buf []byte) { + if card > arrayDefaultMaxSize { + nb := bitmapContainer{ + cardinality: card, + bitmap: byteSliceAsUint64Slice(buf), + } + ra.containers[i] = &nb + return + } + + nb := arrayContainer{ + byteSliceAsUint16Slice(buf), + } + ra.containers[i] = &nb +} + // Reads a serialized roaringArray from a byte slice. func (ra *roaringArray) readFrom(stream internal.ByteInput, cookieHeader ...byte) (int64, error) { var cookie uint32 @@ -764,54 +780,75 @@ func (ra *roaringArray) readFrom(stream internal.ByteInput, cookieHeader ...byte ra.needCopyOnWrite = make([]bool, size) } - for i := uint32(0); i < size; i++ { - key := keycard[2*i] - card := int(keycard[2*i+1]) + 1 - ra.keys[i] = key - ra.needCopyOnWrite[i] = willNeedCopyOnWrite - - if isRunBitmap != nil && isRunBitmap[i/8]&(1<<(i%8)) != 0 { - // run container - nr, err := stream.ReadUInt16() - if err != nil { - return 0, fmt.Errorf("failed to read runtime container size: %s", err) - } - - buf, err := stream.Next(int(nr) * 4) - if err != nil { - return stream.GetReadBytes(), fmt.Errorf("failed to read runtime container content: %s", err) - } - - nb := runContainer16{ - iv: byteSliceAsInterval16Slice(buf), + _, isByteInputAdapter := stream.(*internal.ByteInputAdapter) + if isRunBitmap == nil && isByteInputAdapter { + for i := uint32(0); i < size; { + groupStart := i + groupSize := 0 + for i < size { + card := int(keycard[2*i+1]) + 1 + containerSize := getSizeInBytesFromCardinality(card) + if groupSize > 0 && (groupSize+containerSize > maxContainerGroupSize || + (card > arrayDefaultMaxSize && groupSize%8 != 0)) { + break + } + groupSize += containerSize + i++ } - ra.containers[i] = &nb - } else if card > arrayDefaultMaxSize { - // bitmap container - buf, err := stream.Next(arrayDefaultMaxSize * 2) + buf, err := stream.Next(groupSize) if err != nil { - return stream.GetReadBytes(), fmt.Errorf("failed to read bitmap container: %s", err) - } - - nb := bitmapContainer{ - cardinality: card, - bitmap: byteSliceAsUint64Slice(buf), + return stream.GetReadBytes(), fmt.Errorf("failed to read no-run container group: %s", err) } - - ra.containers[i] = &nb - } else { - // array container - buf, err := stream.Next(card * 2) - if err != nil { - return stream.GetReadBytes(), fmt.Errorf("failed to read array container: %s", err) + groupOffset := 0 + for j := groupStart; j < i; j++ { + card := int(keycard[2*j+1]) + 1 + containerSize := getSizeInBytesFromCardinality(card) + key := keycard[2*j] + ra.keys[j] = key + ra.needCopyOnWrite[j] = willNeedCopyOnWrite + + containerBuf := buf[groupOffset : groupOffset+containerSize : groupOffset+containerSize] + ra.setNoRunContainer(j, card, containerBuf) + groupOffset += containerSize } - - nb := arrayContainer{ - byteSliceAsUint16Slice(buf), + } + } else { + for i := uint32(0); i < size; i++ { + key := keycard[2*i] + card := int(keycard[2*i+1]) + 1 + ra.keys[i] = key + ra.needCopyOnWrite[i] = willNeedCopyOnWrite + + if isRunBitmap != nil && isRunBitmap[i/8]&(1<<(i%8)) != 0 { + // run container + nr, err := stream.ReadUInt16() + if err != nil { + return 0, fmt.Errorf("failed to read runtime container size: %s", err) + } + + buf, err := stream.Next(int(nr) * 4) + if err != nil { + return stream.GetReadBytes(), fmt.Errorf("failed to read runtime container content: %s", err) + } + + nb := runContainer16{ + iv: byteSliceAsInterval16Slice(buf), + } + + ra.containers[i] = &nb + } else { + containerSize := getSizeInBytesFromCardinality(card) + buf, err := stream.Next(containerSize) + if err != nil { + if card > arrayDefaultMaxSize { + return stream.GetReadBytes(), fmt.Errorf("failed to read bitmap container: %s", err) + } + return stream.GetReadBytes(), fmt.Errorf("failed to read array container: %s", err) + } + + ra.setNoRunContainer(i, card, buf) } - - ra.containers[i] = &nb } } diff --git a/serialization_read_benchmark_test.go b/serialization_read_benchmark_test.go new file mode 100644 index 00000000..eb9b4daf --- /dev/null +++ b/serialization_read_benchmark_test.go @@ -0,0 +1,99 @@ +package roaring + +import ( + "fmt" + "io" + "os" + "testing" +) + +type benchmarkCountingReader struct { + reader io.Reader + calls int +} + +func (r *benchmarkCountingReader) Read(p []byte) (int, error) { + r.calls++ + return r.reader.Read(p) +} + +type benchmarkShortReader struct { + reader io.Reader + max int +} + +func (r *benchmarkShortReader) Read(p []byte) (int, error) { + if len(p) > r.max { + p = p[:r.max] + } + return r.reader.Read(p) +} + +func benchmarkNoRunSparseData(containers int) []byte { + bitmap := New() + for i := uint32(0); i < uint32(containers); i++ { + bitmap.Add(i<<16 | 1) + } + data, err := bitmap.ToBytes() + if err != nil { + panic(err) + } + return data +} + +func BenchmarkReadFromNoRunSparse(b *testing.B) { + for _, containers := range []int{1024, 4096, 16384} { + data := benchmarkNoRunSparseData(containers) + b.Run(fmt.Sprintf("file-%d", containers), func(b *testing.B) { + benchmarkReadFromNoRunSparse(b, data, containers, 0) + }) + b.Run(fmt.Sprintf("short-%d", containers), func(b *testing.B) { + benchmarkReadFromNoRunSparse(b, data, containers, 257) + }) + } +} + +func benchmarkReadFromNoRunSparse(b *testing.B, data []byte, containers, maxRead int) { + b.Helper() + + file, err := os.CreateTemp(b.TempDir(), "roaring-benchmark-") + if err != nil { + b.Fatal(err) + } + if _, err := file.Write(data); err != nil { + file.Close() + b.Fatal(err) + } + if _, err := file.Seek(0, io.SeekStart); err != nil { + file.Close() + b.Fatal(err) + } + defer file.Close() + + countingReader := &benchmarkCountingReader{reader: file} + var reader io.Reader = countingReader + if maxRead > 0 { + reader = &benchmarkShortReader{reader: countingReader, max: maxRead} + } + + b.ResetTimer() + var totalReads int + for b.Loop() { + b.StopTimer() + if _, err := file.Seek(0, io.SeekStart); err != nil { + b.StartTimer() + b.Fatal(err) + } + countingReader.calls = 0 + b.StartTimer() + + bitmap := New() + if n, err := bitmap.ReadFrom(reader); err != nil { + b.Fatal(err) + } else if n != int64(len(data)) || bitmap.GetCardinality() != uint64(containers) { + b.Fatalf("unexpected decode: bytes=%d cardinality=%d", n, bitmap.GetCardinality()) + } + totalReads += countingReader.calls + } + b.ReportMetric(float64(totalReads)/float64(b.N), "read-calls/op") +} diff --git a/serialization_read_test.go b/serialization_read_test.go new file mode 100644 index 00000000..311fb30b --- /dev/null +++ b/serialization_read_test.go @@ -0,0 +1,202 @@ +package roaring + +import ( + "bytes" + "io" + "testing" + "unsafe" +) + +type readCountingReader struct { + reader io.Reader + calls int + max int +} + +func (r *readCountingReader) Read(p []byte) (int, error) { + r.calls++ + if r.max > 0 && len(p) > r.max { + p = p[:r.max] + } + return r.reader.Read(p) +} + +func noRunSparseBitmap(containers int) *Bitmap { + bitmap := New() + for i := uint32(0); i < uint32(containers); i++ { + bitmap.Add(i<<16 | 1) + } + return bitmap +} + +func mixedNoRunBitmap() *Bitmap { + bitmap := New() + bitmap.Add(1<<16 | 7) + for low := uint32(0); low <= 8192; low += 2 { + bitmap.Add(2<<16 | low) + } + bitmap.Add(3<<16 | 9) + return bitmap +} + +func serializedBitmap(t *testing.T, bitmap *Bitmap) []byte { + t.Helper() + data, err := bitmap.ToBytes() + if err != nil { + t.Fatal(err) + } + return data +} + +func TestReadFromNoRunBatchesPayloads(t *testing.T) { + want := noRunSparseBitmap(16384) + data := serializedBitmap(t, want) + withSentinel := append(append([]byte(nil), data...), []byte("sentinel")...) + source := bytes.NewReader(withSentinel) + reader := &readCountingReader{reader: source} + got := New() + + read, err := got.ReadFrom(reader) + if err != nil { + t.Fatal(err) + } + if read != int64(len(data)) { + t.Fatalf("ReadFrom consumed %d bytes, want %d", read, len(data)) + } + if !got.Equals(want) { + t.Fatal("decoded bitmap differs from source") + } + if reader.calls != 5 { + t.Fatalf("underlying reader calls = %d, want 5", reader.calls) + } + remaining, err := io.ReadAll(source) + if err != nil { + t.Fatal(err) + } + if string(remaining) != "sentinel" { + t.Fatalf("remaining stream = %q, want sentinel", remaining) + } +} + +func TestReadFromNoRunSplitsOnlyBetweenContainers(t *testing.T) { + want := noRunSparseBitmap(32769) + data := serializedBitmap(t, want) + source := bytes.NewReader(data) + reader := &readCountingReader{reader: source} + got := New() + + read, err := got.ReadFrom(reader) + if err != nil { + t.Fatal(err) + } + if read != int64(len(data)) { + t.Fatalf("ReadFrom consumed %d bytes, want %d", read, len(data)) + } + if !got.Equals(want) { + t.Fatal("decoded bitmap differs from source") + } + if reader.calls != 6 { + t.Fatalf("underlying reader calls = %d, want 6", reader.calls) + } +} + +func TestReadFromNoRunHandlesShortReads(t *testing.T) { + want := noRunSparseBitmap(16384) + data := serializedBitmap(t, want) + withSentinel := append(append([]byte(nil), data...), []byte("sentinel")...) + source := bytes.NewReader(withSentinel) + reader := &readCountingReader{reader: source, max: 3} + got := New() + + read, err := got.ReadFrom(reader) + if err != nil { + t.Fatal(err) + } + if read != int64(len(data)) { + t.Fatalf("ReadFrom consumed %d bytes, want %d", read, len(data)) + } + if !got.Equals(want) { + t.Fatal("decoded bitmap differs from source") + } + remaining, err := io.ReadAll(source) + if err != nil { + t.Fatal(err) + } + if string(remaining) != "sentinel" { + t.Fatalf("remaining stream = %q, want sentinel", remaining) + } +} + +func TestReadFromNoRunAlignsBitmapContainers(t *testing.T) { + want := New() + want.Add(1) + want.Add(2) + want.Add(3) + for low := uint32(0); low <= arrayDefaultMaxSize; low++ { + want.Add(1<<16 | low) + } + + data := serializedBitmap(t, want) + got := New() + if _, err := got.ReadFrom(bytes.NewReader(data)); err != nil { + t.Fatal(err) + } + + bitmap, ok := got.highlowcontainer.containers[1].(*bitmapContainer) + if !ok { + t.Fatalf("container 1 has type %T, want *bitmapContainer", got.highlowcontainer.containers[1]) + } + if address := uintptr(unsafe.Pointer(&bitmap.bitmap[0])); address%8 != 0 { + t.Fatalf("bitmap backing address = %#x, want 8-byte alignment", address) + } + if !got.Equals(want) { + t.Fatal("decoded bitmap differs from source") + } +} + +func TestReadFromNoRunPreservesContainerBoundaries(t *testing.T) { + want := mixedNoRunBitmap() + data := serializedBitmap(t, want) + got := New() + if _, err := got.ReadFrom(bytes.NewReader(data)); err != nil { + t.Fatal(err) + } + + want.Add(1<<16 | 8) + want.Add(2<<16 | 1) + want.Remove(3<<16 | 9) + got.Add(1<<16 | 8) + got.Add(2<<16 | 1) + got.Remove(3<<16 | 9) + if !got.Equals(want) { + t.Fatal("mutating one decoded container changed another container") + } +} + +func TestReadFromNoRunTruncationReportsConsumedBytes(t *testing.T) { + data := serializedBitmap(t, noRunSparseBitmap(16384)) + truncated := data[:len(data)-1] + source := bytes.NewReader(truncated) + got := New() + + read, err := got.ReadFrom(source) + if err == nil { + t.Fatal("truncated bitmap unexpectedly decoded") + } + if read != int64(len(truncated)) { + t.Fatalf("ReadFrom consumed %d bytes, want %d", read, len(truncated)) + } +} + +func TestFromBufferNoRunRetainsCopyOnWrite(t *testing.T) { + data := serializedBitmap(t, mixedNoRunBitmap()) + got := New() + if _, err := got.FromBuffer(data); err != nil { + t.Fatal(err) + } + for i, needsCopy := range got.highlowcontainer.needCopyOnWrite { + if !needsCopy { + t.Fatalf("container %d is not marked copy-on-write", i) + } + } +} diff --git a/util.go b/util.go index e727d6e0..e500928c 100644 --- a/util.go +++ b/util.go @@ -16,6 +16,7 @@ const ( invalidCardinality = -1 serialCookie = 12347 // runs, arrays, and bitmaps noOffsetThreshold = 4 + maxContainerGroupSize = 64 * 1024 // maximum payload size for a batched no-run read // MaxUint32 is the largest uint32 value. MaxUint32 = math.MaxUint32