feat: add crash-safe segment framing

This commit is contained in:
2026-08-31 09:23:31 +00:00
parent 8a88a72e62
commit 43dcd56313
4 changed files with 766 additions and 0 deletions
+149
View File
@@ -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"),
}
}