package store import ( "bytes" "crypto/sha256" "encoding/binary" "encoding/hex" "errors" "fmt" "hash/crc32" "io" "os" "path/filepath" "regexp" "strconv" "sync" "github.com/klauspost/compress/zstd" ) const ( segmentVersion = uint16(1) segmentHeaderSize = uint16(128) segmentTrailerSize = uint64(4) headerCRCOffset = 116 segmentFileMode = 0o600 DefaultSegmentLimit = uint64(64 << 20) ) var ( segmentMagic = [8]byte{'R', 'V', 'B', 'O', 'X', 'S', 'E', 'G'} castagnoliTable = crc32.MakeTable(crc32.Castagnoli) segmentNamePattern = regexp.MustCompile(`^[0-9a-f]{32}-[0-9]{10}\.seg$`) ErrInvalidSegmentRecord = errors.New("invalid segment record") ErrSegmentChecksum = errors.New("segment checksum mismatch") ErrCommittedRangeMissing = errors.New("segment committed range is missing") ErrUnsafeSegmentReference = errors.New("unsafe segment reference") ) type PayloadKind uint16 const ( PayloadKindExecutionSpec PayloadKind = 1 PayloadKindCommandBlob PayloadKind = 2 PayloadKindCommandEvent PayloadKind = 3 PayloadKindStdin PayloadKind = 4 PayloadKindAudit PayloadKind = 5 ) type SegmentRecord struct { OwnerUUID [16]byte Sequence uint64 ObservedUnixNano int64 ReceiptUnixNano int64 Kind PayloadKind Stream uint16 Compression uint16 RawLength uint64 Payload []byte } type SegmentLimits struct { MaxStoredBytes uint64 MaxRawBytes uint64 MaxRecordBytes uint64 } func DefaultSegmentLimits() SegmentLimits { return SegmentLimits{MaxStoredBytes: DefaultSegmentLimit, MaxRawBytes: DefaultSegmentLimit, MaxRecordBytes: DefaultSegmentLimit + uint64(segmentHeaderSize) + segmentTrailerSize} } func EncodeSegmentRecord(record SegmentRecord, limits SegmentLimits) ([]byte, error) { if err := validateSegmentRecord(record, limits); err != nil { return nil, err } storedLength := uint64(len(record.Payload)) recordLength := uint64(segmentHeaderSize) + storedLength + segmentTrailerSize header := make([]byte, segmentHeaderSize) copy(header[0:8], segmentMagic[:]) binary.BigEndian.PutUint16(header[8:10], segmentVersion) binary.BigEndian.PutUint16(header[10:12], segmentHeaderSize) binary.BigEndian.PutUint64(header[12:20], recordLength) copy(header[20:36], record.OwnerUUID[:]) binary.BigEndian.PutUint64(header[36:44], record.Sequence) binary.BigEndian.PutUint64(header[44:52], uint64(record.ObservedUnixNano)) binary.BigEndian.PutUint64(header[52:60], uint64(record.ReceiptUnixNano)) binary.BigEndian.PutUint16(header[60:62], uint16(record.Kind)) binary.BigEndian.PutUint16(header[62:64], record.Stream) binary.BigEndian.PutUint16(header[64:66], record.Compression) binary.BigEndian.PutUint64(header[68:76], record.RawLength) binary.BigEndian.PutUint64(header[76:84], storedLength) payloadDigest := sha256Bytes(record.Payload) copy(header[84:116], payloadDigest[:]) binary.BigEndian.PutUint32(header[headerCRCOffset:headerCRCOffset+4], crc32.Checksum(header, castagnoliTable)) encoded := make([]byte, 0, recordLength) encoded = append(encoded, header...) encoded = append(encoded, record.Payload...) recordCRC := crc32.New(castagnoliTable) _, _ = recordCRC.Write(encoded) encoded = binary.BigEndian.AppendUint32(encoded, recordCRC.Sum32()) return encoded, nil } func DecodeSegmentRecord(reader io.Reader, limits SegmentLimits) (SegmentRecord, uint64, error) { var record SegmentRecord header := make([]byte, segmentHeaderSize) if _, err := io.ReadFull(reader, header); err != nil { return record, 0, err } if !bytes.Equal(header[0:8], segmentMagic[:]) || binary.BigEndian.Uint16(header[8:10]) != segmentVersion || binary.BigEndian.Uint16(header[10:12]) != segmentHeaderSize { return record, 0, ErrInvalidSegmentRecord } wantHeaderCRC := binary.BigEndian.Uint32(header[headerCRCOffset : headerCRCOffset+4]) headerForCRC := append([]byte(nil), header...) binary.BigEndian.PutUint32(headerForCRC[headerCRCOffset:headerCRCOffset+4], 0) if crc32.Checksum(headerForCRC, castagnoliTable) != wantHeaderCRC { return record, 0, ErrSegmentChecksum } if header[66] != 0 || header[67] != 0 || !allZero(header[120:128]) { return record, 0, ErrInvalidSegmentRecord } recordLength := binary.BigEndian.Uint64(header[12:20]) storedLength := binary.BigEndian.Uint64(header[76:84]) if err := validateEncodedLengths(storedLength, binary.BigEndian.Uint64(header[68:76]), recordLength, limits); err != nil { return record, 0, err } payload := make([]byte, int(storedLength)) if _, err := io.ReadFull(reader, payload); err != nil { return record, 0, err } var trailer [segmentTrailerSize]byte if _, err := io.ReadFull(reader, trailer[:]); err != nil { return record, 0, err } recordCRC := crc32.New(castagnoliTable) _, _ = recordCRC.Write(header) _, _ = recordCRC.Write(payload) if recordCRC.Sum32() != binary.BigEndian.Uint32(trailer[:]) { return record, 0, ErrSegmentChecksum } wantDigest := header[84:116] gotDigest := sha256Bytes(payload) if !bytes.Equal(wantDigest, gotDigest[:]) { return record, 0, ErrSegmentChecksum } copy(record.OwnerUUID[:], header[20:36]) record.Sequence = binary.BigEndian.Uint64(header[36:44]) record.ObservedUnixNano = int64(binary.BigEndian.Uint64(header[44:52])) record.ReceiptUnixNano = int64(binary.BigEndian.Uint64(header[52:60])) record.Kind = PayloadKind(binary.BigEndian.Uint16(header[60:62])) record.Stream = binary.BigEndian.Uint16(header[62:64]) record.Compression = binary.BigEndian.Uint16(header[64:66]) record.RawLength = binary.BigEndian.Uint64(header[68:76]) record.Payload = payload if err := validateSegmentRecord(record, limits); err != nil { return SegmentRecord{}, 0, err } return record, recordLength, nil } func validateSegmentRecord(record SegmentRecord, limits SegmentLimits) error { if err := validateLimits(limits); err != nil { return err } if record.Kind < PayloadKindExecutionSpec || record.Kind > PayloadKindAudit || record.Compression < 1 || record.Compression > 2 || record.Stream > 2 { return ErrInvalidSegmentRecord } if record.Kind != PayloadKindAudit && isZeroUUID(record.OwnerUUID) { return ErrInvalidSegmentRecord } if record.Kind == PayloadKindCommandEvent && record.Sequence == 0 { return ErrInvalidSegmentRecord } storedLength := uint64(len(record.Payload)) recordLength := uint64(segmentHeaderSize) + storedLength + segmentTrailerSize if err := validateEncodedLengths(storedLength, record.RawLength, recordLength, limits); err != nil { return err } return validateStoredPayload(record, limits) } func validateStoredPayload(record SegmentRecord, limits SegmentLimits) error { if record.Compression == 1 { if record.RawLength != uint64(len(record.Payload)) { return ErrInvalidSegmentRecord } return nil } decoder, err := zstd.NewReader(bytes.NewReader(record.Payload), zstd.WithDecoderConcurrency(1), zstd.WithDecoderMaxMemory(limits.MaxRawBytes+1)) if err != nil { return fmt.Errorf("%w: zstd header", ErrInvalidSegmentRecord) } defer decoder.Close() decodedBytes, err := io.Copy(io.Discard, io.LimitReader(decoder, int64(limits.MaxRawBytes)+1)) if err != nil || uint64(decodedBytes) != record.RawLength || uint64(decodedBytes) > limits.MaxRawBytes { return fmt.Errorf("%w: zstd payload", ErrInvalidSegmentRecord) } return nil } func validateEncodedLengths(storedLength, rawLength, recordLength uint64, limits SegmentLimits) error { if err := validateLimits(limits); err != nil { return err } if storedLength > limits.MaxStoredBytes || rawLength > limits.MaxRawBytes || recordLength > limits.MaxRecordBytes { return ErrInvalidSegmentRecord } if storedLength > ^uint64(0)-uint64(segmentHeaderSize)-segmentTrailerSize || recordLength != uint64(segmentHeaderSize)+storedLength+segmentTrailerSize { return ErrInvalidSegmentRecord } return nil } func validateLimits(limits SegmentLimits) error { maxInt := uint64(^uint(0) >> 1) if limits.MaxStoredBytes == 0 || limits.MaxRawBytes == 0 || limits.MaxRecordBytes < uint64(segmentHeaderSize)+segmentTrailerSize || limits.MaxStoredBytes > maxInt || limits.MaxRawBytes >= maxInt || limits.MaxRecordBytes > maxInt { return ErrInvalidSegmentRecord } return nil } type Segment struct { mu sync.Mutex file *os.File name string offset uint64 closed bool limits SegmentLimits } func CreateSegment(directory string, owner [16]byte, ordinal uint32, limits SegmentLimits) (*Segment, error) { if isZeroUUID(owner) { return nil, ErrUnsafeSegmentReference } if err := validateLimits(limits); err != nil { return nil, err } if err := ensurePrivateDirectory(directory); err != nil { return nil, err } name := hex.EncodeToString(owner[:]) + "-" + fmt.Sprintf("%010d", ordinal) + ".seg" path := filepath.Join(directory, name) file, err := os.OpenFile(path, os.O_CREATE|os.O_EXCL|os.O_RDWR, segmentFileMode) if err != nil { return nil, err } if err := file.Sync(); err != nil { _ = file.Close() return nil, err } if err := syncDirectory(directory); err != nil { _ = file.Close() return nil, err } return &Segment{file: file, name: name, limits: limits}, nil } func (segment *Segment) Name() string { return segment.name } func (segment *Segment) Append(record SegmentRecord) (uint64, error) { encoded, err := EncodeSegmentRecord(record, segment.limits) if err != nil { return 0, err } segment.mu.Lock() defer segment.mu.Unlock() if segment.closed { return 0, os.ErrClosed } if _, err := segment.file.WriteAt(encoded, int64(segment.offset)); err != nil { return 0, err } if err := segment.file.Sync(); err != nil { return 0, err } segment.offset += uint64(len(encoded)) return segment.offset, nil } func (segment *Segment) Close() error { if segment == nil { return nil } segment.mu.Lock() defer segment.mu.Unlock() if segment.closed { return nil } segment.closed = true return segment.file.Close() } type SegmentRecovery struct { Records []RecoveredSegmentRecord CommittedOffset uint64 OriginalSize uint64 TruncatedTail uint64 } type RecoveredSegmentRecord struct { Offset uint64 Length uint64 Record SegmentRecord } func RecoverSegment(directory, name string, committedOffset uint64, limits SegmentLimits) (SegmentRecovery, error) { var result SegmentRecovery if filepath.Base(name) != name { return result, ErrUnsafeSegmentReference } if _, _, err := parseSegmentName(name); err != nil { return result, err } if err := ensurePrivateDirectory(directory); err != nil { return result, err } path := filepath.Join(directory, name) info, err := os.Lstat(path) if err != nil { return result, err } if !info.Mode().IsRegular() || info.Mode()&os.ModeSymlink != 0 || info.Mode().Perm()&0o077 != 0 { return result, ErrUnsafeSegmentReference } result.OriginalSize = uint64(info.Size()) result.CommittedOffset = committedOffset if result.OriginalSize < committedOffset { return result, ErrCommittedRangeMissing } file, err := os.OpenFile(path, os.O_RDWR, 0) if err != nil { return result, err } defer file.Close() limited := io.LimitReader(file, int64(committedOffset)) var consumed uint64 for consumed < committedOffset { record, length, err := DecodeSegmentRecord(limited, limits) if err != nil { return result, fmt.Errorf("decode committed record at offset %d: %w", consumed, err) } if length == 0 || length > committedOffset-consumed { return result, ErrCommittedRangeMissing } consumed += length result.Records = append(result.Records, RecoveredSegmentRecord{Offset: consumed - length, Length: length, Record: record}) } if consumed != committedOffset { return result, ErrCommittedRangeMissing } if result.OriginalSize > committedOffset { if err := file.Truncate(int64(committedOffset)); err != nil { return result, err } if err := file.Sync(); err != nil { return result, err } result.TruncatedTail = result.OriginalSize - committedOffset } return result, nil } func syncDirectory(path string) error { directory, err := os.Open(path) if err != nil { return err } defer directory.Close() return directory.Sync() } func isZeroUUID(value [16]byte) bool { return value == [16]byte{} } func allZero(value []byte) bool { for _, current := range value { if current != 0 { return false } } return true } func sha256Bytes(value []byte) [sha256.Size]byte { return sha256.Sum256(value) } func parseSegmentName(name string) ([16]byte, uint32, error) { var owner [16]byte if !segmentNamePattern.MatchString(name) { return owner, 0, ErrUnsafeSegmentReference } decoded, err := hex.DecodeString(name[:32]) if err != nil { return owner, 0, ErrUnsafeSegmentReference } copy(owner[:], decoded) ordinal, err := strconv.ParseUint(name[33:43], 10, 32) if err != nil || isZeroUUID(owner) { return [16]byte{}, 0, ErrUnsafeSegmentReference } return owner, uint32(ordinal), nil }