150 lines
4.7 KiB
Go
150 lines
4.7 KiB
Go
package store
|
|
|
|
import (
|
|
"bytes"
|
|
"encoding/binary"
|
|
"errors"
|
|
"hash/crc32"
|
|
"io"
|
|
"reflect"
|
|
"testing"
|
|
|
|
"github.com/klauspost/compress/zstd"
|
|
)
|
|
|
|
func TestSegmentRecordRoundTrip_HP_STORE_02(t *testing.T) {
|
|
t.Parallel()
|
|
|
|
record := sampleSegmentRecord()
|
|
encoded, err := EncodeSegmentRecord(record, DefaultSegmentLimits())
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
decoded, length, err := DecodeSegmentRecord(bytes.NewReader(encoded), DefaultSegmentLimits())
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if length != uint64(len(encoded)) || !reflect.DeepEqual(decoded, record) {
|
|
t.Fatalf("decoded = (%+v, %d), want (%+v, %d)", decoded, length, record, len(encoded))
|
|
}
|
|
encoder, err := zstd.NewWriter(nil, zstd.WithEncoderConcurrency(1))
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
zstdRecord := record
|
|
zstdRecord.Compression = 2
|
|
zstdRecord.RawLength = 5
|
|
zstdRecord.Payload = encoder.EncodeAll([]byte("zstd!"), nil)
|
|
encoder.Close()
|
|
encoded, err = EncodeSegmentRecord(zstdRecord, DefaultSegmentLimits())
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
decoded, _, err = DecodeSegmentRecord(bytes.NewReader(encoded), DefaultSegmentLimits())
|
|
if err != nil || !reflect.DeepEqual(decoded, zstdRecord) {
|
|
t.Fatalf("zstd decoded = (%+v, %v)", decoded, err)
|
|
}
|
|
}
|
|
|
|
func TestSegmentRecordChecksumFailures_BH_STORE_04(t *testing.T) {
|
|
t.Parallel()
|
|
|
|
encoded, err := EncodeSegmentRecord(sampleSegmentRecord(), DefaultSegmentLimits())
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
for name, offset := range map[string]int{
|
|
"header": 20,
|
|
"payload": int(segmentHeaderSize),
|
|
"trailer": len(encoded) - 1,
|
|
} {
|
|
t.Run(name, func(t *testing.T) {
|
|
corrupt := append([]byte(nil), encoded...)
|
|
corrupt[offset] ^= 0xff
|
|
if _, _, err := DecodeSegmentRecord(bytes.NewReader(corrupt), DefaultSegmentLimits()); !errors.Is(err, ErrSegmentChecksum) {
|
|
t.Fatalf("Decode error = %v, want checksum failure", err)
|
|
}
|
|
})
|
|
}
|
|
for _, length := range []int{0, 7, int(segmentHeaderSize) - 1, len(encoded) - 1} {
|
|
if _, _, err := DecodeSegmentRecord(bytes.NewReader(encoded[:length]), DefaultSegmentLimits()); err == nil {
|
|
t.Errorf("truncated length %d accepted", length)
|
|
}
|
|
}
|
|
}
|
|
|
|
func TestSegmentRecordLengthBoundsBeforeAllocation_BOUND_STORE_01(t *testing.T) {
|
|
t.Parallel()
|
|
|
|
limits := SegmentLimits{MaxStoredBytes: 64, MaxRawBytes: 128, MaxRecordBytes: 256}
|
|
header := make([]byte, segmentHeaderSize)
|
|
copy(header[:8], segmentMagic[:])
|
|
binary.BigEndian.PutUint16(header[8:10], segmentVersion)
|
|
binary.BigEndian.PutUint16(header[10:12], segmentHeaderSize)
|
|
binary.BigEndian.PutUint64(header[12:20], uint64(segmentHeaderSize)+65+segmentTrailerSize)
|
|
header[20] = 1
|
|
binary.BigEndian.PutUint16(header[60:62], uint16(PayloadKindCommandBlob))
|
|
binary.BigEndian.PutUint16(header[64:66], 1)
|
|
binary.BigEndian.PutUint64(header[68:76], 65)
|
|
binary.BigEndian.PutUint64(header[76:84], 65)
|
|
binary.BigEndian.PutUint32(header[headerCRCOffset:headerCRCOffset+4], crc32.Checksum(header, castagnoliTable))
|
|
reader := &countingReader{Reader: bytes.NewReader(header)}
|
|
if _, _, err := DecodeSegmentRecord(reader, limits); !errors.Is(err, ErrInvalidSegmentRecord) {
|
|
t.Fatalf("Decode error = %v", err)
|
|
}
|
|
if reader.read != int64(segmentHeaderSize) {
|
|
t.Fatalf("read %d bytes, want header only", reader.read)
|
|
}
|
|
|
|
bad := sampleSegmentRecord()
|
|
bad.Kind = 0
|
|
if _, err := EncodeSegmentRecord(bad, DefaultSegmentLimits()); !errors.Is(err, ErrInvalidSegmentRecord) {
|
|
t.Fatalf("invalid kind error = %v", err)
|
|
}
|
|
bad = sampleSegmentRecord()
|
|
bad.OwnerUUID = [16]byte{}
|
|
if _, err := EncodeSegmentRecord(bad, DefaultSegmentLimits()); !errors.Is(err, ErrInvalidSegmentRecord) {
|
|
t.Fatalf("zero owner error = %v", err)
|
|
}
|
|
bad = sampleSegmentRecord()
|
|
bad.Compression = 2
|
|
if _, err := EncodeSegmentRecord(bad, DefaultSegmentLimits()); !errors.Is(err, ErrInvalidSegmentRecord) {
|
|
t.Fatalf("invalid zstd error = %v", err)
|
|
}
|
|
}
|
|
|
|
func FuzzDecodeSegmentRecordDoesNotEscapeBounds(f *testing.F) {
|
|
encoded, err := EncodeSegmentRecord(sampleSegmentRecord(), DefaultSegmentLimits())
|
|
if err != nil {
|
|
f.Fatal(err)
|
|
}
|
|
f.Add(encoded)
|
|
f.Add([]byte("RVBOXSEG"))
|
|
f.Fuzz(func(t *testing.T, data []byte) {
|
|
limits := SegmentLimits{MaxStoredBytes: 1024, MaxRawBytes: 2048, MaxRecordBytes: 4096}
|
|
_, _, _ = DecodeSegmentRecord(bytes.NewReader(data), limits)
|
|
})
|
|
}
|
|
|
|
type countingReader struct {
|
|
io.Reader
|
|
read int64
|
|
}
|
|
|
|
func (reader *countingReader) Read(data []byte) (int, error) {
|
|
count, err := reader.Reader.Read(data)
|
|
reader.read += int64(count)
|
|
return count, err
|
|
}
|
|
|
|
func sampleSegmentRecord() SegmentRecord {
|
|
var owner [16]byte
|
|
for index := range owner {
|
|
owner[index] = byte(index + 1)
|
|
}
|
|
return SegmentRecord{
|
|
OwnerUUID: owner, Sequence: 7, ObservedUnixNano: 1_700_000_000, ReceiptUnixNano: 1_700_000_001,
|
|
Kind: PayloadKindCommandEvent, Stream: 1, Compression: 1, RawLength: 10, Payload: []byte("compressed"),
|
|
}
|
|
}
|