Skip to content
73 changes: 50 additions & 23 deletions src/rootfilespec/rntuple/envelope.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@
from dataclasses import dataclass, field
from typing import Annotated, Generic, TypeVar, cast

import xxhash # type: ignore[import-not-found]
from typing_extensions import Self

from rootfilespec.bootstrap.compression import decompress
Expand Down Expand Up @@ -74,12 +75,8 @@ class REnvelope(ROOTSerializable):
@classmethod
def read(cls, buffer: ReadBuffer) -> tuple[Self, ReadBuffer]:
"""Reads an REnvelope from the given buffer."""
#### Save initial buffer position (for checking unknown bytes)
payload_start_pos = buffer.relpos

#### Get the first 64bit integer (lengthType) which contains the length and type of the envelope
# lengthType, buffer = buffer.consume(8)
(lengthType,), buffer = buffer.unpack("<Q")
(lengthType,), _ = buffer.unpack("<Q")

# Envelope type, encoded in the 16 least significant bits
typeID = lengthType & 0xFFFF
Expand All @@ -88,30 +85,44 @@ def read(cls, buffer: ReadBuffer) -> tuple[Self, ReadBuffer]:
msg = f"Envelope type {typeID} read does not match passed class {cls.__name__}"
raise ValueError(msg)

# Envelope size (uncompressed), encoded in the 48 most significant bits
# Envelope size (uncompressed), encoded in the 48 most significant bits.
# It includes the preamble and the checksum, so it is at least 16 bytes
length = lengthType >> 16
# Ensure that the length of the envelope matches the buffer length
if length - 8 != len(buffer):
msg = f"Length of envelope ({length} minus 8) of type {typeID} does not match buffer length ({len(buffer)})"
if length < 16:
msg = f"Length of envelope ({length}) of type {typeID} is shorter than 16 bytes"
raise ValueError(msg)
if length != len(buffer):
msg = f"Length of envelope ({length}) of type {typeID} does not match buffer length ({len(buffer)})"
raise ValueError(msg)

#### Split the envelope once: the bytes the checksum covers, then the checksum
# The checksum covers [0, length - 8): the preamble, the payload and any
# unknown trailing bytes (root-io-spec ERRATA 5; "Checksum verification
# ... must include both known and unknown contents")
covered, trailer = buffer[: length - 8], buffer[length - 8 :]
(checksum,), rest = trailer.unpack("<Q")

#### Verify it before trusting the payload, as ROOT does
# (RNTupleSerialize.cxx:909-939)
computed = xxhash.xxh3_64_intdigest(covered.data)
if computed != checksum:
msg = (
f"{cls.__name__} checksum mismatch: "
f"stored {checksum:#018x}, computed {computed:#018x}"
)
raise ValueError(msg)

members = {"typeID": typeID, "length": length}
#### Get the payload
members, buffer = cls.update_members(members, buffer)
#### Get the payload, after the 8-byte preamble
_, payload = covered.consume(8)
members = {"typeID": typeID, "length": length, "checksum": checksum}
members, payload = cls.update_members(members, payload)

#### Consume any unknown trailing information in the envelope
_unknown, buffer = buffer.consume(
length - (buffer.relpos - payload_start_pos) - 8
)
# Unknown Bytes = Envelope Size - Envelope Bytes Read - Checksum (8 bytes)
# Envelope Bytes Read = buffer.relpos - payload_start_pos
#### Keep any unknown trailing information in the envelope
_unknown, _ = payload.consume(len(payload))

#### Get the checksum (appended to envelope when writing to disk)
(checksum,), buffer = buffer.unpack("<Q") # Last 8 bytes of the envelope
members["checksum"] = checksum
envelope = cls(**members)
envelope._unknown = _unknown
return envelope, buffer
return envelope, rest


EnvType = TypeVar("EnvType", bound=REnvelope)
Expand Down Expand Up @@ -149,8 +160,24 @@ def read_from(self, buffer: ReadBuffer) -> EnvType:

Envelopes are compressed, so this decompresses and deserializes.
"""
if len(buffer) != self.size:
msg = (
f"{self.envtype.__name__} at {self.locator}: expected {self.size} "
f"bytes, got {len(buffer)}"
)
raise ValueError(msg)

#### Decompress the buffer if necessary
if len(buffer) != self.length:
# RNTuple decompression tests equality of the stored size (the locator's)
# and the length: equal means stored raw, smaller compressed, and larger
# is an error (root-io-spec NOTES 2; RNTupleZip.hxx:106-113)
if self.size > self.length:
msg = (
f"{self.envtype.__name__} at {self.locator}: stored size "
f"{self.size} is larger than its uncompressed length {self.length}"
)
raise ValueError(msg)
if self.size < self.length:
buffer = decompress(buffer, self.length)

#### Now read the envelope
Expand Down
161 changes: 161 additions & 0 deletions tests/test_envelope_checksums.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,161 @@
from pathlib import Path

import pytest

from rootfilespec.bootstrap import BOOTSTRAP_CONTEXT, ROOT3a3aRNTuple
from rootfilespec.reader import open_path
from rootfilespec.rntuple.envelope import REnvelopeLocator
from rootfilespec.rntuple.header import HeaderEnvelope
from rootfilespec.rntuple.RLocator import StandardLocator
from rootfilespec.serializable import BufferContext, ReadBuffer

DATA = Path(__file__).parent.parent / "reference" / "root-io-spec" / "data" / "rntuple"
pytestmark = pytest.mark.skipif(
not DATA.exists(), reason="reference/root-io-spec not checked out"
)


def _anchors(path: Path) -> list[ROOT3a3aRNTuple]:
with open_path(path) as reader:
keylist = reader.keylist()
return [
reader.fetch(keylist[name])
for name in keylist
if keylist[name].fClassName == b"ROOT::RNTuple"
]


def _buffer(raw: bytes, offset: int, size: int) -> ReadBuffer:
return ReadBuffer(
memoryview(raw[offset : offset + size]),
0,
BOOTSTRAP_CONTEXT,
BufferContext(abspos=offset),
)


def _read(raw: bytes, loc):
return loc.read_from(_buffer(raw, loc.offset, loc.size))


@pytest.mark.parametrize("name", sorted(p.name for p in DATA.glob("*.root")))
def test_every_envelope_verifies(name: str):
"""Issue #117: every envelope of every root-io-spec RNTuple fixture is checked

Including rntuple/compressed.root, whose envelopes are compressed: the
checksum is over the uncompressed envelope.
"""
path = DATA / name
raw = path.read_bytes()
for anchor in _anchors(path):
header = _read(raw, anchor.header_locator)
footer = _read(raw, anchor.footer_locator)
assert footer.headerChecksum == header.checksum
for loc in footer.pagelist_locators:
assert _read(raw, loc).headerChecksum == header.checksum


def test_header_checksum_is_the_pinned_value():
"""rntuple/anchor.root's case.toml pins the header envelope's checksum at 500
and the footer's copy of it at 854"""
path = DATA / "anchor.root"
raw = path.read_bytes()
(anchor,) = _anchors(path)
pinned = bytes([0x2A, 0x13, 0xA2, 0x84, 0xB3, 0x59, 0xD3, 0xDD])
assert raw[500:508] == raw[854:862] == pinned
header = _read(raw, anchor.header_locator)
footer = _read(raw, anchor.footer_locator)
assert header.checksum == footer.headerChecksum == int.from_bytes(pinned, "little")


def test_corrupted_header_envelope_raises():
"""rntuple/anchor.root's header envelope is 268..508 (length 240, checksum at
500..508). Changing any byte before the checksum must be caught."""
path = DATA / "anchor.root"
raw = bytearray(path.read_bytes())
(anchor,) = _anchors(path)
loc = anchor.header_locator
assert (loc.offset, loc.size, loc.length) == (268, 240, 240)
# The "n" of the ntuple's name "ntpl" at 288: flipping its case keeps the
# envelope parseable, so only the checksum can catch it
assert raw[288:292] == b"ntpl"
raw[288] ^= 0x20
with pytest.raises(ValueError, match="HeaderEnvelope checksum mismatch"):
_read(bytes(raw), loc)


def test_corrupted_checksum_raises():
path = DATA / "anchor.root"
raw = bytearray(path.read_bytes())
(anchor,) = _anchors(path)
loc = anchor.header_locator
raw[loc.offset + loc.length - 1] ^= 0x01
with pytest.raises(ValueError, match="HeaderEnvelope checksum mismatch"):
_read(bytes(raw), loc)


def test_stored_size_larger_than_length_raises():
"""root-io-spec NOTES 2: RNTuple decompression tests equality; a stored size
larger than the uncompressed length is an error, not a raw payload"""
loc = REnvelopeLocator(
length=16, locator=StandardLocator(size=24, offset=0), envtype=HeaderEnvelope
)
with pytest.raises(ValueError, match="larger than its uncompressed length"):
loc.read_from(_buffer(bytes(24), 0, 24))


def test_envelope_shorter_than_16_bytes_raises():
"""The length counts the 8-byte preamble and the 8-byte checksum, so it is
at least 16, as ROOT's DeserializeEnvelope requires"""
loc = REnvelopeLocator(
length=8, locator=StandardLocator(size=8, offset=0), envtype=HeaderEnvelope
)
# The preamble alone: type 1 in the low 16 bits, length 8 in the upper 48
preamble = (8 << 16 | 0x01).to_bytes(8, "little")
with pytest.raises(ValueError, match="shorter than 16 bytes"):
loc.read_from(_buffer(preamble, 0, 8))


def test_envelope_length_not_matching_the_locator_raises():
"""rntuple/anchor.root's header envelope with the length in its preamble
(bytes 270..276) changed from 240 to 248: the anchor says 240 were stored"""
path = DATA / "anchor.root"
raw = bytearray(path.read_bytes())
(anchor,) = _anchors(path)
loc = anchor.header_locator
assert raw[268:276] == (240 << 16 | 0x01).to_bytes(8, "little")
raw[268:276] = (248 << 16 | 0x01).to_bytes(8, "little")
with pytest.raises(
ValueError, match=r"Length of envelope \(248\) .* buffer length \(240\)"
):
_read(bytes(raw), loc)


def test_every_corrupted_byte_is_a_checksum_mismatch():
"""The checksum is verified before the payload is parsed, as ROOT does, so
corrupting any byte the checksum covers reports the checksum, not whatever
the parser trips over first (an out-of-range slice, unknown feature flags,
...)."""
path = DATA / "anchor.root"
raw = path.read_bytes()
(anchor,) = _anchors(path)
loc = anchor.header_locator
# The 8-byte preamble is checked first (type, length); everything after it,
# up to the checksum, is covered only by the checksum
for pos in range(loc.offset + 8, loc.offset + loc.length - 8):
corrupted = bytearray(raw)
corrupted[pos] ^= 0x80
with pytest.raises(ValueError, match="HeaderEnvelope checksum mismatch"):
_read(bytes(corrupted), loc)


@pytest.mark.parametrize("size", [0, 200, 239])
def test_short_read_raises(size: int):
"""Fewer bytes than the locator's size (a truncated file) are reported as
such, not taken for a compressed envelope"""
path = DATA / "anchor.root"
raw = path.read_bytes()
(anchor,) = _anchors(path)
loc = anchor.header_locator
with pytest.raises(ValueError, match=f"expected 240 bytes, got {size}"):
loc.read_from(_buffer(raw, loc.offset, size))
Loading