feat: add crash-safe segment framing
This commit is contained in:
@@ -0,0 +1,149 @@
|
||||
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"),
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user