package mfer import ( "bytes" "context" "crypto/sha256" "errors" "fmt" "io" "strings" "github.com/klauspost/compress/zstd" "github.com/spf13/afero" "google.golang.org/protobuf/encoding/protowire" "google.golang.org/protobuf/proto" "sneak.berlin/go/mfer/internal/bork" "sneak.berlin/go/mfer/internal/log" ) var ( errInvalidUUIDLength = errors.New("invalid UUID length") errUnknownVersion = errors.New("unknown version") errUnknownCompression = errors.New("unknown compression type") errCompressedHashWrong = errors.New("compressed data hash mismatch") errSignatureNoPubKey = errors.New("signature present but no public key") errDecompressedTooLarge = errors.New("decompressed data exceeds maximum allowed size") errUUIDMismatch = errors.New("outer and inner UUID mismatch") errInvalidFileFormat = errors.New("invalid file format") errInvalidManifestPath = errors.New("manifest contains invalid path") errDecodedTooLarge = errors.New( "manifest would take too much memory to decode") errSignerNotSigningKey = errors.New( "signer is not the fingerprint of the key that made the signature") ) // validateUUID checks that the byte slice is the 16 bytes of a binary UUID. // Any 16 bytes are one, so the length is all there is to check. func validateUUID(data []byte) error { if len(data) != uuidLength { return errInvalidUUIDLength } return nil } // validateOuterHeader checks the outer message's version, compression // type, and UUID. func (m *manifest) validateOuterHeader() error { if m.pbOuter.GetVersion() != MFFileOuter_VERSION_ONE { return errUnknownVersion } if m.pbOuter.GetCompressionType() != MFFileOuter_COMPRESSION_ZSTD { return errUnknownCompression } // Validate outer UUID before any decompression err := validateUUID(m.pbOuter.GetUuid()) if err != nil { return fmt.Errorf("outer UUID invalid: %w", err) } return nil } // verifyOuterIntegrity checks the hash of the compressed payload and, if a // signature is present, verifies it against the embedded public key, which // must be one key, and checks that the signer field is that key's // fingerprint. func (m *manifest) verifyOuterIntegrity() error { h := sha256.New() _, err := h.Write(m.pbOuter.GetInnerMessage()) if err != nil { return fmt.Errorf("deserialize: hash write: %w", err) } sha256Hash := h.Sum(nil) if !bytes.Equal(sha256Hash, m.pbOuter.GetSha256()) { return errCompressedHashWrong } if len(m.pbOuter.GetSignature()) == 0 { return nil } if len(m.pbOuter.GetSigningPubKey()) == 0 { return errSignatureNoPubKey } sigString, err := m.signatureString() if err != nil { return fmt.Errorf( "failed to generate signature string for verification: %w", err, ) } // Loading a manifest takes no context; gpgTimeout still bounds gpg. signingKey, err := gpgVerify( context.Background(), []byte(sigString), m.pbOuter.GetSignature(), m.pbOuter.GetSigningPubKey(), ) if err != nil { return fmt.Errorf("signature verification failed: %w", err) } if !strings.EqualFold(string(m.pbOuter.GetSigner()), signingKey) { return fmt.Errorf("%w: signer %q, signing key %s", errSignerNotSigningKey, m.pbOuter.GetSigner(), signingKey) } log.Infof("signature verified successfully") return nil } // decompressInner decompresses the inner payload, enforcing size limits // to prevent decompression bombs. func (m *manifest) decompressInner() ([]byte, error) { bb := bytes.NewBuffer(m.pbOuter.GetInnerMessage()) // By default the decoder decodes a payload under 128 KiB in full, // each frame up to the decoder's limit, before the LimitReader below // reads any of it. Decoding synchronously and never in full makes it // decode only what the LimitReader asks for. It also sets aside a new // buffer for each frame that asks for a larger window than the frames // before it, even a frame holding no data. Refusing windows above // zstdWindowSize keeps each buffer to a little over zstdWindowSize, and // the buffers of frames holding no data to about 16 times it in total. zr, err := zstd.NewReader(bb, zstd.WithDecoderConcurrency(1), zstd.WithDecodeBuffersBelow(0), zstd.WithDecoderMaxWindow(zstdWindowSize)) if err != nil { return nil, fmt.Errorf("deserialize: zstd reader: %w", err) } defer zr.Close() // Limit decompressed size to prevent decompression bombs. // Use declared size + 1 byte to detect overflow, capped at MaxDecompressedSize. maxSize := MaxDecompressedSize if m.pbOuter.GetSize() > 0 && m.pbOuter.GetSize() < maxSize { maxSize = m.pbOuter.GetSize() + 1 } limitedReader := io.LimitReader(zr, maxSize) dat, err := io.ReadAll(limitedReader) if err != nil { return nil, fmt.Errorf("deserialize: decompress: %w", err) } if int64(len(dat)) >= MaxDecompressedSize { return nil, fmt.Errorf( "%w of %d bytes", errDecompressedTooLarge, MaxDecompressedSize, ) } return dat, nil } // checkDecodedSize refuses an encoded inner message whose file entries, // hashes, timestamps and MIME types would take more than maxDecodedGrowth // times its size to decode. Decoding sets aside a fixed amount for each, // however short its encoding, so a message of empty ones would take about // 50 times its size. func checkDecodedSize(inner []byte) error { limit := maxDecodedGrowth * int64(len(inner)) var decoded int64 add := func(size int64) error { decoded += size if decoded > limit { return errDecodedTooLarge } return nil } return forEachBytesField(inner, func(num protowire.Number, entry []byte) error { if num != filesFieldNumber { return nil } err := add(decodedFileEntrySize) if err != nil { return err } return forEachBytesField(entry, func(num protowire.Number, _ []byte) error { if num == hashesFieldNumber { return add(decodedHashSize) } if num == mtimeFieldNumber || num == ctimeFieldNumber { return add(decodedTimestampSize) } if num == mimeTypeFieldNumber { return add(decodedMIMETypeSize) } return nil }) }) } // forEachBytesField calls fn with the number and value of each // length-delimited field in the encoded message msg, and fails if msg is // malformed. func forEachBytesField( msg []byte, fn func(num protowire.Number, value []byte) error, ) error { for len(msg) > 0 { num, wireType, tagLen := protowire.ConsumeTag(msg) if tagLen < 0 { return protowire.ParseError(tagLen) } valueLen := protowire.ConsumeFieldValue(num, wireType, msg[tagLen:]) if valueLen < 0 { return protowire.ParseError(valueLen) } if wireType == protowire.BytesType { value, _ := protowire.ConsumeBytes(msg[tagLen:]) err := fn(num, value) if err != nil { return err } } msg = msg[tagLen+valueLen:] } return nil } func (m *manifest) deserializeInner() error { err := m.validateOuterHeader() if err != nil { return err } err = m.verifyOuterIntegrity() if err != nil { return err } dat, err := m.decompressInner() if err != nil { return err } isize := len(dat) if int64(isize) != m.pbOuter.GetSize() { log.Debugf("truncated data, got %d expected %d", isize, m.pbOuter.GetSize()) return bork.ErrFileTruncated } err = checkDecodedSize(dat) if err != nil { return fmt.Errorf("deserialize: unmarshal inner: %w", err) } // Deserialize inner message m.pbInner = new(MFFile) // Unknown fields would cost memory; mfer never writes a loaded manifest out. err = proto.UnmarshalOptions{DiscardUnknown: true}.Unmarshal(dat, m.pbInner) if err != nil { return fmt.Errorf("deserialize: unmarshal inner: %w", err) } // Validate inner UUID err = validateUUID(m.pbInner.GetUuid()) if err != nil { return fmt.Errorf("inner UUID invalid: %w", err) } // Verify UUIDs match if !bytes.Equal(m.pbOuter.GetUuid(), m.pbInner.GetUuid()) { return errUUIDMismatch } // Enforce the manifest path invariants on every entry as it is loaded, // so that no consumer of a manifest — Checker today, any restore or // extract path tomorrow — acts on a traversal or absolute path from an // untrusted .mf. Reject loudly on the first offender rather than // dropping entries, which would let a hostile manifest hide files from a // check. A path listed twice is refused too: check would check the one // file against both entries. seen := make(map[string]bool, len(m.pbInner.GetFiles())) for _, f := range m.pbInner.GetFiles() { err = ValidatePath(f.GetPath()) if err != nil { return fmt.Errorf("%w: %w", errInvalidManifestPath, err) } if seen[f.GetPath()] { return fmt.Errorf("%w %q", errDuplicatePath, f.GetPath()) } seen[f.GetPath()] = true } log.Infof("loaded manifest with %d files", len(m.pbInner.GetFiles())) return nil } func validateMagic(dat []byte) bool { ml := len([]byte(MAGIC)) if len(dat) < ml { return false } got := dat[0:ml] expected := []byte(MAGIC) return bytes.Equal(got, expected) } // NewManifestFromReader reads a manifest from an io.Reader. // //nolint:revive // unexported-return: exporting manifest is owner question 13 func NewManifestFromReader(input io.Reader) (*manifest, error) { m := &manifest{} dat, err := io.ReadAll(input) if err != nil { return nil, err } if !validateMagic(dat) { return nil, errInvalidFileFormat } // remove magic bytes prefix: ml := len([]byte(MAGIC)) bb := bytes.NewBuffer(dat[ml:]) dat = bb.Bytes() // deserialize outer: m.pbOuter = new(MFFileOuter) // Unknown fields would cost memory; mfer never writes a loaded manifest out. err = proto.UnmarshalOptions{DiscardUnknown: true}.Unmarshal(dat, m.pbOuter) if err != nil { return nil, err } // deserialize inner: err = m.deserializeInner() if err != nil { return nil, err } return m, nil } // ManifestFromFileOptions configures NewManifestFromFile. type ManifestFromFileOptions struct { // Path is the manifest file to read (required). Path string // Fs is the filesystem to use, defaults to OsFs if nil. Fs afero.Fs } // NewManifestFromFile reads a manifest from a file. It returns an error if // opts is nil or its path is empty. // //nolint:revive // unexported-return: exporting manifest is owner question 13 func NewManifestFromFile(opts *ManifestFromFileOptions) (*manifest, error) { if opts == nil || opts.Path == "" { return nil, errManifestPathEmpty } fs := opts.Fs if fs == nil { fs = afero.NewOsFs() } f, err := fs.Open(opts.Path) if err != nil { return nil, err } defer func() { _ = f.Close() }() return NewManifestFromReader(f) }