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"), } }