Files
rvbox/internal/server/store/segment.go
T

396 lines
12 KiB
Go

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 []SegmentRecord
CommittedOffset uint64
OriginalSize uint64
TruncatedTail uint64
}
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, 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
}