diff --git a/src/k1link/media_fragments.py b/src/k1link/media_fragments.py new file mode 100644 index 0000000..8b0b4d7 --- /dev/null +++ b/src/k1link/media_fragments.py @@ -0,0 +1,408 @@ +"""Bounded ISO-BMFF timing shared by archive inspection and stream decode. + +Pure bytes only: no session store, source paths, materializer or decoder import. +""" + +from __future__ import annotations + +from collections.abc import Iterator +from dataclasses import dataclass + +MAX_SAFE_INTEGER = (1 << 53) - 1 +MAX_MP4_BOXES = 100_000 +MAX_MP4_SAMPLES_PER_FRAGMENT = 1_000_000 + + +class Mp4IntegrityError(ValueError): + """An ISO-BMFF fragment has absent, ambiguous or unbounded timing.""" + + +@dataclass(frozen=True, slots=True) +class VideoTiming: + track_id: int + timescale: int + default_sample_duration: int | None + default_sample_flags: int | None + + +@dataclass(frozen=True, slots=True) +class VideoFragmentTiming: + base_decode_time: int + duration_units: int + random_access: bool + + @property + def end_decode_time(self) -> int: + return self.base_decode_time + self.duration_units + + +@dataclass(slots=True) +class ParseBudget: + boxes_remaining: int = MAX_MP4_BOXES + samples_remaining: int = MAX_MP4_SAMPLES_PER_FRAGMENT + + def consume_box(self) -> None: + self.boxes_remaining -= 1 + if self.boxes_remaining < 0: + raise Mp4IntegrityError("recorded media ISO-BMFF box budget was exceeded") + + def consume_samples(self, count: int) -> None: + if count < 0 or count > self.samples_remaining: + raise Mp4IntegrityError("recorded media ISO-BMFF sample budget was exceeded") + self.samples_remaining -= count + + +def video_timing(payload: bytes, budget: ParseBudget) -> VideoTiming: + moov_payloads = [ + box_payload + for box_type, box_payload in _iter_mp4_boxes(payload, budget) + if box_type == b"moov" + ] + if len(moov_payloads) != 1: + raise Mp4IntegrityError("recorded media init has no unique moov box") + moov = moov_payloads[0] + moov_boxes = tuple(_iter_mp4_boxes(moov, budget)) + defaults: dict[int, tuple[int, int]] = {} + for box_type, box_payload in moov_boxes: + if box_type != b"mvex": + continue + for child_type, child_payload in _iter_mp4_boxes(box_payload, budget): + if child_type != b"trex": + continue + track_id, trex_default_duration, trex_default_flags = _parse_trex(child_payload) + if track_id in defaults: + raise Mp4IntegrityError("recorded media init repeats a trex track") + defaults[track_id] = (trex_default_duration, trex_default_flags) + + video_tracks: list[tuple[int, int]] = [] + for box_type, trak_payload in moov_boxes: + if box_type != b"trak": + continue + track_id = _trak_track_id(trak_payload, budget) + media = _trak_media_timing(trak_payload, budget) + if media is not None: + video_tracks.append((track_id, media)) + if len(video_tracks) != 1: + raise Mp4IntegrityError("recorded media init has no unique video track") + track_id, timescale = video_tracks[0] + default_sample = defaults.get(track_id) + default_duration = None if default_sample is None else default_sample[0] + return VideoTiming( + track_id=track_id, + timescale=timescale, + default_sample_duration=( + default_duration if default_duration is not None and default_duration > 0 else None + ), + default_sample_flags=None if default_sample is None else default_sample[1], + ) + + +def _trak_track_id(payload: bytes, budget: ParseBudget) -> int: + track_ids = [ + _parse_tkhd_track_id(box_payload) + for box_type, box_payload in _iter_mp4_boxes(payload, budget) + if box_type == b"tkhd" + ] + if len(track_ids) != 1 or track_ids[0] <= 0: + raise Mp4IntegrityError("recorded media track id is invalid") + return track_ids[0] + + +def _trak_media_timing(payload: bytes, budget: ParseBudget) -> int | None: + media_boxes = [ + box_payload + for box_type, box_payload in _iter_mp4_boxes(payload, budget) + if box_type == b"mdia" + ] + if len(media_boxes) != 1: + raise Mp4IntegrityError("recorded media track has no unique mdia box") + children = tuple(_iter_mp4_boxes(media_boxes[0], budget)) + handlers = [ + _parse_hdlr_type(box_payload) for box_type, box_payload in children if box_type == b"hdlr" + ] + if len(handlers) != 1: + raise Mp4IntegrityError("recorded media track handler is ambiguous") + if handlers[0] != b"vide": + return None + timescales = [ + _parse_mdhd_timescale(box_payload) + for box_type, box_payload in children + if box_type == b"mdhd" + ] + if len(timescales) != 1 or timescales[0] <= 0: + raise Mp4IntegrityError("recorded media video timescale is invalid") + return timescales[0] + + +def video_fragment_timing( + payload: bytes, + timing: VideoTiming, + budget: ParseBudget, +) -> VideoFragmentTiming: + moof_payloads = [ + box_payload + for box_type, box_payload in _iter_mp4_boxes(payload, budget) + if box_type == b"moof" + ] + if len(moof_payloads) != 1: + raise Mp4IntegrityError("recorded media fragment has no unique moof box") + matching_timings: list[VideoFragmentTiming] = [] + for box_type, traf_payload in _iter_mp4_boxes(moof_payloads[0], budget): + if box_type != b"traf": + continue + boxes = tuple(_iter_mp4_boxes(traf_payload, budget)) + tfhd_payloads = [box for kind, box in boxes if kind == b"tfhd"] + if len(tfhd_payloads) != 1: + raise Mp4IntegrityError("recorded media fragment tfhd is ambiguous") + track_id, fragment_default_duration, fragment_default_flags = _parse_tfhd(tfhd_payloads[0]) + if track_id != timing.track_id: + continue + tfdt_payloads = [box for kind, box in boxes if kind == b"tfdt"] + if len(tfdt_payloads) != 1: + raise Mp4IntegrityError("recorded media fragment tfdt is ambiguous") + base_decode_time = _parse_tfdt(tfdt_payloads[0]) + trun_payloads = [box for kind, box in boxes if kind == b"trun"] + if not trun_payloads: + raise Mp4IntegrityError("recorded media video fragment has no trun box") + default_duration = fragment_default_duration or timing.default_sample_duration + default_flags = ( + fragment_default_flags + if fragment_default_flags is not None + else timing.default_sample_flags + ) + trun_descriptors = tuple( + _parse_trun_descriptor(trun, default_duration, default_flags, budget) + for trun in trun_payloads + ) + fragment_sample_count = sum(item[2] for item in trun_descriptors) + if fragment_sample_count != 1: + raise Mp4IntegrityError("recorded media video fragment must contain exactly one sample") + duration = sum(item[0] for item in trun_descriptors) + if base_decode_time > MAX_SAFE_INTEGER - duration: + raise Mp4IntegrityError("recorded media fragment decode time is outside bounds") + matching_timings.append( + VideoFragmentTiming( + base_decode_time=base_decode_time, + duration_units=duration, + random_access=(trun_descriptors[0][1] & 0x00010000) == 0, + ) + ) + if len(matching_timings) != 1: + raise Mp4IntegrityError("recorded media fragment video track is ambiguous") + return matching_timings[0] + + +def _iter_mp4_boxes( + payload: bytes, + budget: ParseBudget, +) -> Iterator[tuple[bytes, bytes]]: + offset = 0 + payload_length = len(payload) + while offset < payload_length: + if payload_length - offset < 8: + raise Mp4IntegrityError("recorded media ISO-BMFF box is truncated") + size = int.from_bytes(payload[offset : offset + 4], "big") + box_type = payload[offset + 4 : offset + 8] + header_length = 8 + if size == 1: + if payload_length - offset < 16: + raise Mp4IntegrityError("recorded media ISO-BMFF box is truncated") + size = int.from_bytes(payload[offset + 8 : offset + 16], "big") + header_length = 16 + elif size == 0: + size = payload_length - offset + if size < header_length or size > payload_length - offset: + raise Mp4IntegrityError("recorded media ISO-BMFF box size is invalid") + budget.consume_box() + end = offset + size + yield box_type, payload[offset + header_length : end] + offset = end + + +def single_sample_payload(payload: bytes, timing: VideoTiming) -> memoryview: + """Extract the reviewed low-delay fMP4 subset without opening a demuxer. + + Exactly moof+mdat, one traf/trun/sample, explicit default sample size, + default-base-is-moof, no composition offset, no extra tracks or padding. + This intentionally accepts less than the archive timing inspector. Other + valid MP4 layouts require explicit qualification, never guessed offsets. + Call video_fragment_timing separately to validate continuity and duration. + """ + budget = ParseBudget(128, 1) + top = tuple(_iter_mp4_boxes(payload, budget)) + if [kind for kind, _ in top] != [b"moof", b"mdat"]: + raise Mp4IntegrityError("stream fragment requires exactly moof then mdat") + moof, mdat = top[0][1], top[1][1] + moof_size = int.from_bytes(payload[:4], "big") + if moof_size != len(moof) + 8 or moof_size < 8: + raise Mp4IntegrityError("stream fragment requires explicit 32-bit box sizes") + mdat_size = int.from_bytes(payload[moof_size : moof_size + 4], "big") + if mdat_size != len(mdat) + 8 or moof_size + mdat_size != len(payload): + raise Mp4IntegrityError("stream mdat size is ambiguous") + children = tuple(_iter_mp4_boxes(moof, budget)) + if [kind for kind, _ in children] != [b"mfhd", b"traf"]: + raise Mp4IntegrityError("stream moof requires one mfhd and one traf") + boxes = tuple(_iter_mp4_boxes(children[1][1], budget)) + if [kind for kind, _ in boxes] != [b"tfhd", b"tfdt", b"trun"]: + raise Mp4IntegrityError("stream traf layout is unsupported") + tfhd, trun = boxes[0][1], boxes[2][1] + if len(tfhd) != 20 or tfhd[:4] != b"\x00\x02\x00\x38": + raise Mp4IntegrityError("explicit moof-relative sample defaults required") + if _read_u32(tfhd, 4, "track id") != timing.track_id: + raise Mp4IntegrityError("stream sample track changed") + if not mdat or _read_u32(tfhd, 12, "sample size") != len(mdat): + raise Mp4IntegrityError("stream sample size does not match mdat") + flags = _full_box_flags(trun) + if trun[0] != 0 or flags not in (1, 5) or len(trun) != (16 if flags == 5 else 12): + raise Mp4IntegrityError("stream requires a single sample without composition offset") + if _read_u32(trun, 4, "sample count") != 1: + raise Mp4IntegrityError("stream requires exactly one sample") + if _read_u32(trun, 8, "data offset") != moof_size + 8: + raise Mp4IntegrityError("stream sample offset does not address mdat") + return memoryview(payload)[moof_size + 8 :] + + +def _parse_tkhd_track_id(payload: bytes) -> int: + version = _full_box_version(payload) + offset = 20 if version == 1 else 12 if version == 0 else -1 + return _read_u32(payload, offset, "tkhd track id") + + +def _parse_mdhd_timescale(payload: bytes) -> int: + version = _full_box_version(payload) + offset = 20 if version == 1 else 12 if version == 0 else -1 + return _read_u32(payload, offset, "mdhd timescale") + + +def _parse_hdlr_type(payload: bytes) -> bytes: + _full_box_version(payload) + if len(payload) < 12: + raise Mp4IntegrityError("recorded media hdlr box is truncated") + return payload[8:12] + + +def _parse_trex(payload: bytes) -> tuple[int, int, int]: + _full_box_version(payload) + return ( + _read_u32(payload, 4, "trex track id"), + _read_u32(payload, 12, "trex default sample duration"), + _read_u32(payload, 20, "trex default sample flags"), + ) + + +def _parse_tfhd(payload: bytes) -> tuple[int, int | None, int | None]: + flags = _full_box_flags(payload) + track_id = _read_u32(payload, 4, "tfhd track id") + cursor = 8 + for flag, width in ((0x000001, 8), (0x000002, 4)): + if flags & flag: + cursor = _advance_box_cursor(payload, cursor, width, "tfhd optional field") + default_duration: int | None = None + if flags & 0x000008: + default_duration = _read_u32(payload, cursor, "tfhd default sample duration") + cursor += 4 + if flags & 0x000010: + cursor = _advance_box_cursor(payload, cursor, 4, "tfhd default sample size") + default_flags: int | None = None + if flags & 0x000020: + default_flags = _read_u32(payload, cursor, "tfhd default sample flags") + return ( + track_id, + default_duration if default_duration and default_duration > 0 else None, + default_flags, + ) + + +def _parse_tfdt(payload: bytes) -> int: + version = _full_box_version(payload) + if version == 0: + return _read_u32(payload, 4, "tfdt base decode time") + if version == 1: + if len(payload) < 12: + raise Mp4IntegrityError("recorded media tfdt base decode time is truncated") + return int.from_bytes(payload[4:12], "big") + raise Mp4IntegrityError("recorded media tfdt version is unsupported") + + +def _parse_trun_descriptor( + payload: bytes, + default_duration: int | None, + default_sample_flags: int | None, + budget: ParseBudget, +) -> tuple[int, int, int]: + flags = _full_box_flags(payload) + sample_count = _read_u32(payload, 4, "trun sample count") + if sample_count < 1: + raise Mp4IntegrityError("recorded media trun has no samples") + budget.consume_samples(sample_count) + cursor = 8 + if flags & 0x000001: + cursor = _advance_box_cursor(payload, cursor, 4, "trun data offset") + first_sample_flags: int | None = None + if flags & 0x000004: + first_sample_flags = _read_u32(payload, cursor, "trun first sample flags") + cursor += 4 + if first_sample_flags is not None and flags & 0x000400: + raise Mp4IntegrityError("recorded media trun sample flags are ambiguous") + + duration = 0 + first_per_sample_flags: int | None = None + for sample_index in range(sample_count): + if flags & 0x000100: + sample_duration = _read_u32(payload, cursor, "trun sample duration") + if sample_duration <= 0: + raise Mp4IntegrityError("recorded media sample duration is invalid") + duration += sample_duration + cursor += 4 + elif default_duration is None or default_duration <= 0: + raise Mp4IntegrityError("recorded media sample duration is unavailable") + else: + duration += default_duration + if flags & 0x000200: + cursor = _advance_box_cursor(payload, cursor, 4, "trun sample size") + if flags & 0x000400: + sample_flags = _read_u32(payload, cursor, "trun sample flags") + if sample_index == 0: + first_per_sample_flags = sample_flags + cursor += 4 + if flags & 0x000800: + cursor = _advance_box_cursor( + payload, + cursor, + 4, + "trun sample composition time offset", + ) + + effective_flags = ( + first_per_sample_flags + if first_per_sample_flags is not None + else first_sample_flags + if first_sample_flags is not None + else default_sample_flags + ) + if effective_flags is None: + raise Mp4IntegrityError("recorded media sample flags are unavailable") + return duration, effective_flags, sample_count + + +def _full_box_version(payload: bytes) -> int: + if len(payload) < 4: + raise Mp4IntegrityError("recorded media full box is truncated") + return payload[0] + + +def _full_box_flags(payload: bytes) -> int: + _full_box_version(payload) + return int.from_bytes(payload[1:4], "big") + + +def _read_u32(payload: bytes, offset: int, description: str) -> int: + if offset < 0 or offset + 4 > len(payload): + raise Mp4IntegrityError(f"recorded media {description} is truncated") + return int.from_bytes(payload[offset : offset + 4], "big") + + +def _advance_box_cursor(payload: bytes, cursor: int, width: int, description: str) -> int: + if cursor < 0 or width < 0 or cursor + width > len(payload): + raise Mp4IntegrityError(f"recorded media {description} is truncated") + return cursor + width diff --git a/src/k1link/perception/streaming_decoder.py b/src/k1link/perception/streaming_decoder.py new file mode 100644 index 0000000..bf8bbe4 --- /dev/null +++ b/src/k1link/perception/streaming_decoder.py @@ -0,0 +1,110 @@ +"""One-sample fMP4 → persistent H.264 decoder, no recording/file dependency. + +Reuses the archive's bounded ISO-BMFF timing validator, not its materializer. +Only the supplied init and current complete fragment are inspected. Codec state +persists across fragments; no later fragment or source EOF is supplied to make +the current frame appear. Other codecs/reordered/multi-sample streams fail +closed and require a different explicit profile. Run inside a supervised, +memory-limited CPU child; native decoder allocations are not Python byte counts. +""" + +from __future__ import annotations + +import importlib +import io +from fractions import Fraction +from typing import Any + +from k1link.media_fragments import ( + ParseBudget, + single_sample_payload, + video_fragment_timing, + video_timing, +) + + +class StreamingDecodeError(ValueError): + pass + + +class FragmentDecoder: + def __init__(self) -> None: + self._av = importlib.import_module("av") + if self._av.__version__ != "18.0.0": + raise StreamingDecodeError("unqualified PyAV version") + # Import before readiness, not on the first source frame's deadline. + importlib.import_module("numpy") + self._codec: Any = None + self._init: bytes | None = None + self._timing: Any = None + self._next_dts: int | None = None + self.frames = 0 + self.failed = False + + def configure(self, payload: bytes) -> None: + if self._init is not None or self.failed or not 0 < len(payload) <= 65536: + self.failed = True + raise StreamingDecodeError("invalid or repeated codec initialization") + try: + timing = video_timing(payload, ParseBudget(128, 1)) + with self._av.open(io.BytesIO(payload), format="mp4", mode="r") as container: + if len(container.streams) != 1 or len(container.streams.video) != 1: + raise StreamingDecodeError("one video track is required") + description = container.streams.video[0].codec_context + if description.name != "h264" or (description.width, description.height) != ( + 800, + 600, + ): + raise StreamingDecodeError("codec or calibrated native raster changed") + extra = description.extradata + if not extra or len(extra) > 65536: + raise StreamingDecodeError("bounded H264 codec configuration required") + codec = self._av.CodecContext.create("h264", "r") + codec.extradata = extra + codec.thread_count = 1 + codec.thread_type = "SLICE" + codec.open() + self._codec, self._init, self._timing = codec, bytes(payload), timing + except Exception: + self.failed = True + raise + + def decode(self, payload: bytes) -> Any: + if self.failed or self._init is None or not 0 < len(payload) <= 1024 * 1024: + self.failed = True + raise StreamingDecodeError("uninitialized, failed or oversized fragment") + try: + timing = video_fragment_timing(payload, self._timing, ParseBudget(128, 1)) + if self.frames == 0 and not timing.random_access: + raise StreamingDecodeError("epoch must start at a random-access fragment") + if self._next_dts is not None and timing.base_decode_time != self._next_dts: + raise StreamingDecodeError("fragment codec timeline is discontinuous") + # Opening a new demuxer per fragment repeats FFmpeg stream probing. + # Validate the narrow explicit layout, then feed the one AVCC sample + # to the persistent codec. No future packet or flush is supplied. + packet = self._av.Packet(single_sample_payload(payload, self._timing)) + packet.pts = packet.dts = timing.base_decode_time + packet.duration = timing.duration_units + packet.time_base = Fraction(1, self._timing.timescale) + frames = self._codec.decode(packet) + if len(frames) != 1 or self._codec.has_b_frames: + raise StreamingDecodeError("current frame requires future input or reordering") + frame = frames[0] + if (frame.width, frame.height) != (800, 600) or frame.pts != packet.pts: + raise StreamingDecodeError("decoded raster or frame identity changed") + if self.frames == 0 and not frame.key_frame: + raise StreamingDecodeError("decoder did not confirm the initial keyframe") + image = frame.to_ndarray(format="bgr24") + if image.shape != (600, 800, 3) or image.nbytes != 1_440_000: + raise StreamingDecodeError("invalid native decoder output") + image.setflags(write=False) + self._next_dts = timing.base_decode_time + timing.duration_units + self.frames += 1 + return image + except Exception: + self.failed = True + raise + + def close(self) -> None: + self.failed = True + self._codec = self._init = self._timing = None diff --git a/src/k1link/sessions/media.py b/src/k1link/sessions/media.py index 26424fb..6f5a0e3 100644 --- a/src/k1link/sessions/media.py +++ b/src/k1link/sessions/media.py @@ -8,12 +8,14 @@ import re import secrets import stat import threading -from collections.abc import Iterable, Iterator +from collections.abc import Iterable from contextlib import suppress from dataclasses import dataclass from pathlib import Path from typing import Any, TypeGuard +from k1link import media_fragments as _fragment_timing + from .models import RecordedMediaArtifact, ReplayCommand, SessionIntegrityError CAMERA_ARCHIVE_SCHEMA = "missioncore.camera-recording/v1" @@ -25,8 +27,8 @@ MAX_MEDIA_INDEX_LINE_BYTES = 64 * 1024 MAX_INIT_BYTES = 8 * 1024 * 1024 MAX_MEDIA_SEGMENT_BYTES = 64 * 1024 * 1024 MAX_SAFE_INTEGER = (1 << 53) - 1 -MAX_MP4_BOXES = 100_000 -MAX_MP4_SAMPLES_PER_FRAGMENT = 1_000_000 +MAX_MP4_BOXES = _fragment_timing.MAX_MP4_BOXES +MAX_MP4_SAMPLES_PER_FRAGMENT = _fragment_timing.MAX_MP4_SAMPLES_PER_FRAGMENT MAX_MP4_FRAGMENT_DURATION_SECONDS = 3_600.0 MEDIA_RECORDING_TIMELINE_TOLERANCE_SECONDS = 0.05 MAX_MEDIA_EPOCHS = 4_096 @@ -88,39 +90,9 @@ class _CachedManifest: manifest: RecordedMediaManifest -@dataclass(frozen=True, slots=True) -class _Mp4VideoTiming: - track_id: int - timescale: int - default_sample_duration: int | None - default_sample_flags: int | None - - -@dataclass(frozen=True, slots=True) -class _Mp4VideoFragmentTiming: - base_decode_time: int - duration_units: int - random_access: bool - - @property - def end_decode_time(self) -> int: - return self.base_decode_time + self.duration_units - - -@dataclass(slots=True) -class _Mp4ParseBudget: - boxes_remaining: int = MAX_MP4_BOXES - samples_remaining: int = MAX_MP4_SAMPLES_PER_FRAGMENT - - def consume_box(self) -> None: - self.boxes_remaining -= 1 - if self.boxes_remaining < 0: - raise SessionIntegrityError("recorded media ISO-BMFF box budget was exceeded") - - def consume_samples(self, count: int) -> None: - if count < 0 or count > self.samples_remaining: - raise SessionIntegrityError("recorded media ISO-BMFF sample budget was exceeded") - self.samples_remaining -= count +_Mp4VideoTiming = _fragment_timing.VideoTiming +_Mp4VideoFragmentTiming = _fragment_timing.VideoFragmentTiming +_Mp4ParseBudget = _fragment_timing.ParseBudget class RecordedMediaInspector: @@ -1155,85 +1127,10 @@ def _checked_fragment_duration_seconds(duration_units: int, timescale: int) -> f def _mp4_video_timing(payload: bytes, budget: _Mp4ParseBudget) -> _Mp4VideoTiming: - moov_payloads = [ - box_payload - for box_type, box_payload in _iter_mp4_boxes(payload, budget) - if box_type == b"moov" - ] - if len(moov_payloads) != 1: - raise SessionIntegrityError("recorded media init has no unique moov box") - moov = moov_payloads[0] - moov_boxes = tuple(_iter_mp4_boxes(moov, budget)) - defaults: dict[int, tuple[int, int]] = {} - for box_type, box_payload in moov_boxes: - if box_type != b"mvex": - continue - for child_type, child_payload in _iter_mp4_boxes(box_payload, budget): - if child_type != b"trex": - continue - track_id, trex_default_duration, trex_default_flags = _parse_trex(child_payload) - if track_id in defaults: - raise SessionIntegrityError("recorded media init repeats a trex track") - defaults[track_id] = (trex_default_duration, trex_default_flags) - - video_tracks: list[tuple[int, int]] = [] - for box_type, trak_payload in moov_boxes: - if box_type != b"trak": - continue - track_id = _trak_track_id(trak_payload, budget) - media = _trak_media_timing(trak_payload, budget) - if media is not None: - video_tracks.append((track_id, media)) - if len(video_tracks) != 1: - raise SessionIntegrityError("recorded media init has no unique video track") - track_id, timescale = video_tracks[0] - default_sample = defaults.get(track_id) - default_duration = None if default_sample is None else default_sample[0] - return _Mp4VideoTiming( - track_id=track_id, - timescale=timescale, - default_sample_duration=( - default_duration if default_duration is not None and default_duration > 0 else None - ), - default_sample_flags=None if default_sample is None else default_sample[1], - ) - - -def _trak_track_id(payload: bytes, budget: _Mp4ParseBudget) -> int: - track_ids = [ - _parse_tkhd_track_id(box_payload) - for box_type, box_payload in _iter_mp4_boxes(payload, budget) - if box_type == b"tkhd" - ] - if len(track_ids) != 1 or track_ids[0] <= 0: - raise SessionIntegrityError("recorded media track id is invalid") - return track_ids[0] - - -def _trak_media_timing(payload: bytes, budget: _Mp4ParseBudget) -> int | None: - media_boxes = [ - box_payload - for box_type, box_payload in _iter_mp4_boxes(payload, budget) - if box_type == b"mdia" - ] - if len(media_boxes) != 1: - raise SessionIntegrityError("recorded media track has no unique mdia box") - children = tuple(_iter_mp4_boxes(media_boxes[0], budget)) - handlers = [ - _parse_hdlr_type(box_payload) for box_type, box_payload in children if box_type == b"hdlr" - ] - if len(handlers) != 1: - raise SessionIntegrityError("recorded media track handler is ambiguous") - if handlers[0] != b"vide": - return None - timescales = [ - _parse_mdhd_timescale(box_payload) - for box_type, box_payload in children - if box_type == b"mdhd" - ] - if len(timescales) != 1 or timescales[0] <= 0: - raise SessionIntegrityError("recorded media video timescale is invalid") - return timescales[0] + try: + return _fragment_timing.video_timing(payload, budget) + except _fragment_timing.Mp4IntegrityError as exc: + raise SessionIntegrityError(str(exc)) from exc def _mp4_video_fragment_timing( @@ -1241,232 +1138,10 @@ def _mp4_video_fragment_timing( timing: _Mp4VideoTiming, budget: _Mp4ParseBudget, ) -> _Mp4VideoFragmentTiming: - moof_payloads = [ - box_payload - for box_type, box_payload in _iter_mp4_boxes(payload, budget) - if box_type == b"moof" - ] - if len(moof_payloads) != 1: - raise SessionIntegrityError("recorded media fragment has no unique moof box") - matching_timings: list[_Mp4VideoFragmentTiming] = [] - for box_type, traf_payload in _iter_mp4_boxes(moof_payloads[0], budget): - if box_type != b"traf": - continue - boxes = tuple(_iter_mp4_boxes(traf_payload, budget)) - tfhd_payloads = [box for kind, box in boxes if kind == b"tfhd"] - if len(tfhd_payloads) != 1: - raise SessionIntegrityError("recorded media fragment tfhd is ambiguous") - track_id, fragment_default_duration, fragment_default_flags = _parse_tfhd(tfhd_payloads[0]) - if track_id != timing.track_id: - continue - tfdt_payloads = [box for kind, box in boxes if kind == b"tfdt"] - if len(tfdt_payloads) != 1: - raise SessionIntegrityError("recorded media fragment tfdt is ambiguous") - base_decode_time = _parse_tfdt(tfdt_payloads[0]) - trun_payloads = [box for kind, box in boxes if kind == b"trun"] - if not trun_payloads: - raise SessionIntegrityError("recorded media video fragment has no trun box") - default_duration = fragment_default_duration or timing.default_sample_duration - default_flags = ( - fragment_default_flags - if fragment_default_flags is not None - else timing.default_sample_flags - ) - trun_descriptors = tuple( - _parse_trun_descriptor(trun, default_duration, default_flags, budget) - for trun in trun_payloads - ) - fragment_sample_count = sum(item[2] for item in trun_descriptors) - if fragment_sample_count != 1: - raise SessionIntegrityError( - "recorded media video fragment must contain exactly one sample" - ) - duration = sum(item[0] for item in trun_descriptors) - if base_decode_time > MAX_SAFE_INTEGER - duration: - raise SessionIntegrityError("recorded media fragment decode time is outside bounds") - matching_timings.append( - _Mp4VideoFragmentTiming( - base_decode_time=base_decode_time, - duration_units=duration, - random_access=(trun_descriptors[0][1] & 0x00010000) == 0, - ) - ) - if len(matching_timings) != 1: - raise SessionIntegrityError("recorded media fragment video track is ambiguous") - return matching_timings[0] - - -def _iter_mp4_boxes( - payload: bytes, - budget: _Mp4ParseBudget, -) -> Iterator[tuple[bytes, bytes]]: - offset = 0 - payload_length = len(payload) - while offset < payload_length: - if payload_length - offset < 8: - raise SessionIntegrityError("recorded media ISO-BMFF box is truncated") - size = int.from_bytes(payload[offset : offset + 4], "big") - box_type = payload[offset + 4 : offset + 8] - header_length = 8 - if size == 1: - if payload_length - offset < 16: - raise SessionIntegrityError("recorded media ISO-BMFF box is truncated") - size = int.from_bytes(payload[offset + 8 : offset + 16], "big") - header_length = 16 - elif size == 0: - size = payload_length - offset - if size < header_length or size > payload_length - offset: - raise SessionIntegrityError("recorded media ISO-BMFF box size is invalid") - budget.consume_box() - end = offset + size - yield box_type, payload[offset + header_length : end] - offset = end - - -def _parse_tkhd_track_id(payload: bytes) -> int: - version = _full_box_version(payload) - offset = 20 if version == 1 else 12 if version == 0 else -1 - return _read_u32(payload, offset, "tkhd track id") - - -def _parse_mdhd_timescale(payload: bytes) -> int: - version = _full_box_version(payload) - offset = 20 if version == 1 else 12 if version == 0 else -1 - return _read_u32(payload, offset, "mdhd timescale") - - -def _parse_hdlr_type(payload: bytes) -> bytes: - _full_box_version(payload) - if len(payload) < 12: - raise SessionIntegrityError("recorded media hdlr box is truncated") - return payload[8:12] - - -def _parse_trex(payload: bytes) -> tuple[int, int, int]: - _full_box_version(payload) - return ( - _read_u32(payload, 4, "trex track id"), - _read_u32(payload, 12, "trex default sample duration"), - _read_u32(payload, 20, "trex default sample flags"), - ) - - -def _parse_tfhd(payload: bytes) -> tuple[int, int | None, int | None]: - flags = _full_box_flags(payload) - track_id = _read_u32(payload, 4, "tfhd track id") - cursor = 8 - for flag, width in ((0x000001, 8), (0x000002, 4)): - if flags & flag: - cursor = _advance_box_cursor(payload, cursor, width, "tfhd optional field") - default_duration: int | None = None - if flags & 0x000008: - default_duration = _read_u32(payload, cursor, "tfhd default sample duration") - cursor += 4 - if flags & 0x000010: - cursor = _advance_box_cursor(payload, cursor, 4, "tfhd default sample size") - default_flags: int | None = None - if flags & 0x000020: - default_flags = _read_u32(payload, cursor, "tfhd default sample flags") - return ( - track_id, - default_duration if default_duration and default_duration > 0 else None, - default_flags, - ) - - -def _parse_tfdt(payload: bytes) -> int: - version = _full_box_version(payload) - if version == 0: - return _read_u32(payload, 4, "tfdt base decode time") - if version == 1: - if len(payload) < 12: - raise SessionIntegrityError("recorded media tfdt base decode time is truncated") - return int.from_bytes(payload[4:12], "big") - raise SessionIntegrityError("recorded media tfdt version is unsupported") - - -def _parse_trun_descriptor( - payload: bytes, - default_duration: int | None, - default_sample_flags: int | None, - budget: _Mp4ParseBudget, -) -> tuple[int, int, int]: - flags = _full_box_flags(payload) - sample_count = _read_u32(payload, 4, "trun sample count") - if sample_count < 1: - raise SessionIntegrityError("recorded media trun has no samples") - budget.consume_samples(sample_count) - cursor = 8 - if flags & 0x000001: - cursor = _advance_box_cursor(payload, cursor, 4, "trun data offset") - first_sample_flags: int | None = None - if flags & 0x000004: - first_sample_flags = _read_u32(payload, cursor, "trun first sample flags") - cursor += 4 - if first_sample_flags is not None and flags & 0x000400: - raise SessionIntegrityError("recorded media trun sample flags are ambiguous") - - duration = 0 - first_per_sample_flags: int | None = None - for sample_index in range(sample_count): - if flags & 0x000100: - sample_duration = _read_u32(payload, cursor, "trun sample duration") - if sample_duration <= 0: - raise SessionIntegrityError("recorded media sample duration is invalid") - duration += sample_duration - cursor += 4 - elif default_duration is None or default_duration <= 0: - raise SessionIntegrityError("recorded media sample duration is unavailable") - else: - duration += default_duration - if flags & 0x000200: - cursor = _advance_box_cursor(payload, cursor, 4, "trun sample size") - if flags & 0x000400: - sample_flags = _read_u32(payload, cursor, "trun sample flags") - if sample_index == 0: - first_per_sample_flags = sample_flags - cursor += 4 - if flags & 0x000800: - cursor = _advance_box_cursor( - payload, - cursor, - 4, - "trun sample composition time offset", - ) - - effective_flags = ( - first_per_sample_flags - if first_per_sample_flags is not None - else first_sample_flags - if first_sample_flags is not None - else default_sample_flags - ) - if effective_flags is None: - raise SessionIntegrityError("recorded media sample flags are unavailable") - return duration, effective_flags, sample_count - - -def _full_box_version(payload: bytes) -> int: - if len(payload) < 4: - raise SessionIntegrityError("recorded media full box is truncated") - return payload[0] - - -def _full_box_flags(payload: bytes) -> int: - _full_box_version(payload) - return int.from_bytes(payload[1:4], "big") - - -def _read_u32(payload: bytes, offset: int, description: str) -> int: - if offset < 0 or offset + 4 > len(payload): - raise SessionIntegrityError(f"recorded media {description} is truncated") - return int.from_bytes(payload[offset : offset + 4], "big") - - -def _advance_box_cursor(payload: bytes, cursor: int, width: int, description: str) -> int: - if cursor < 0 or width < 0 or cursor + width > len(payload): - raise SessionIntegrityError(f"recorded media {description} is truncated") - return cursor + width + try: + return _fragment_timing.video_fragment_timing(payload, timing, budget) + except _fragment_timing.Mp4IntegrityError as exc: + raise SessionIntegrityError(str(exc)) from exc def _mp4_media_type(init_payload: bytes) -> str: diff --git a/tests/test_perception_streaming_decoder.py b/tests/test_perception_streaming_decoder.py new file mode 100644 index 0000000..5bc0c3e --- /dev/null +++ b/tests/test_perception_streaming_decoder.py @@ -0,0 +1,231 @@ +"""Synthetic boxes and fake codec only; real decoding belongs on Worker 006.""" + +import subprocess +import sys +from types import SimpleNamespace + +import pytest + +from k1link import media_fragments as mp4 +from k1link.perception import streaming_decoder as decoder_module + + +def box(kind, data=b""): + return (len(data) + 8).to_bytes(4, "big") + kind + data + + +def u32(value): + return value.to_bytes(4, "big") + + +def fixture(*, dts=0, sync=True, size=5, offset_delta=0, count=1, trun_flags=1): + tkhd = box(b"tkhd", bytes(12) + u32(1)) + mdhd = box(b"mdhd", bytes(12) + u32(1000)) + hdlr = box(b"hdlr", bytes(8) + b"vide") + init = box(b"moov", box(b"trak", tkhd + box(b"mdia", mdhd + hdlr))) + tfhd = box( + b"tfhd", b"\x00\x02\x00\x38" + u32(1) + u32(100) + u32(size) + u32(0 if sync else 0x10000) + ) + tfdt = box(b"tfdt", bytes(4) + u32(dts)) + mfhd = box(b"mfhd", bytes(8)) + + def moof(offset): + trun = box(b"trun", u32(trun_flags) + u32(count) + u32(offset)) + return box(b"moof", mfhd + box(b"traf", tfhd + tfdt + trun)) + + return init, moof(len(moof(0)) + 8 + offset_delta) + box(b"mdat", b"frame") + + +def test_shared_parser_and_direct_sample_are_source_neutral(): + init, fragment = fixture() + timing = mp4.video_timing(init, mp4.ParseBudget(128, 1)) + frame = mp4.video_fragment_timing(fragment, timing, mp4.ParseBudget(128, 1)) + assert (timing.track_id, timing.timescale) == (1, 1000) + assert (frame.base_decode_time, frame.duration_units, frame.random_access) == (0, 100, True) + sample = mp4.single_sample_payload(fragment, timing) + assert sample == b"frame" and sample.obj is fragment + script = ( + "import sys; import k1link.perception.streaming_decoder; " + "assert 'k1link.sessions' not in sys.modules; " + "assert 'k1link.compute' not in sys.modules; assert 'av' not in sys.modules" + ) + subprocess.run([sys.executable, "-c", script], check=True, timeout=5) + + +@pytest.mark.parametrize( + "changes,match", + [ + ({"size": 4}, "size"), + ({"offset_delta": 1}, "offset"), + ({"count": 2}, "one sample"), + ({"trun_flags": 0x801}, "composition"), + ], +) +def test_direct_sample_rejects_guessed_offsets_or_sizes(changes, match): + init, fragment = fixture(**changes) + with pytest.raises(mp4.Mp4IntegrityError, match=match): + mp4.single_sample_payload(fragment, mp4.video_timing(init, mp4.ParseBudget())) + + +def test_parser_bounds_and_truncation(): + init, fragment = fixture(count=2) + timing = mp4.video_timing(init, mp4.ParseBudget()) + with pytest.raises(mp4.Mp4IntegrityError, match="sample budget"): + mp4.video_fragment_timing(fragment, timing, mp4.ParseBudget(128, 1)) + with pytest.raises(mp4.Mp4IntegrityError, match="box budget"): + mp4.video_timing(init, mp4.ParseBudget(1, 1)) + with pytest.raises(mp4.Mp4IntegrityError, match="box size"): + mp4.single_sample_payload(fragment[:-1], timing) + with pytest.raises(mp4.Mp4IntegrityError, match="moof then mdat"): + mp4.single_sample_payload(fragment + box(b"free"), timing) + + +@pytest.fixture +def fake_av(monkeypatch): + class Image: + shape, nbytes = (600, 800, 3), 1_440_000 + + def setflags(self, *, write): + assert write is False + + class Codec: + has_b_frames = False + + def open(self): + pass + + def decode(self, packet): + self.calls.append(packet) + return ( + [] + if self.future + else [ + SimpleNamespace( + width=self.width, + height=600, + pts=packet.pts + self.pts_delta, + key_frame=self.key_frame, + to_ndarray=lambda **kwargs: Image(), + ) + ] + ) + + codec = Codec() + codec.calls, codec.future, codec.width, codec.pts_delta, codec.key_frame = ( + [], + False, + 800, + 0, + True, + ) + + class Streams(list): + @property + def video(self): + return self + + class Container: + def __enter__(self): + return SimpleNamespace( + streams=Streams( + [ + SimpleNamespace( + codec_context=SimpleNamespace( + name="h264", width=800, height=600, extradata=b"config" + ) + ) + ] + ) + ) + + def __exit__(self, *args): + pass + + opens = [] + + def open_container(stream, **kwargs): + opens.append(stream.getvalue()) + return Container() + + av = SimpleNamespace( + __version__="18.0.0", + open=open_container, + CodecContext=SimpleNamespace(create=lambda *args: codec), + Packet=lambda raw: SimpleNamespace(raw=bytes(raw)), + ) + original = decoder_module.importlib.import_module + monkeypatch.setattr( + decoder_module.importlib, + "import_module", + lambda name: av if name == "av" else original(name), + ) + return codec, av, opens + + +def test_persistent_codec_emits_current_frame_without_demux_or_flush(fake_av): + codec, _, opens = fake_av + init, fragment = fixture() + decoder = decoder_module.FragmentDecoder() + decoder.configure(init) + for sequence in range(3): + decoder.decode(fixture(dts=sequence * 100)[1]) + assert len(opens) == 1 and len(codec.calls) == decoder.frames == 3 + assert codec.thread_count == 1 and codec.thread_type == "SLICE" + assert all(p.raw == b"frame" and p.pts == p.dts for p in codec.calls) + decoder.close() + with pytest.raises(decoder_module.StreamingDecodeError, match="failed"): + decoder.decode(fragment) + + +@pytest.mark.parametrize( + "mutation,match", + [ + (("future", True), "future"), + (("has_b_frames", True), "reordering"), + (("width", 801), "raster"), + (("pts_delta", 1), "identity"), + (("key_frame", False), "keyframe"), + ], +) +def test_ambiguous_native_output_poisons_decoder(fake_av, mutation, match): + codec, _, _ = fake_av + setattr(codec, *mutation) + decoder = decoder_module.FragmentDecoder() + init, fragment = fixture() + decoder.configure(init) + with pytest.raises(decoder_module.StreamingDecodeError, match=match): + decoder.decode(fragment) + assert decoder.failed and decoder.frames == 0 + + +def test_codec_gap_does_not_use_later_frame_as_current(fake_av): + codec, _, _ = fake_av + decoder = decoder_module.FragmentDecoder() + decoder.configure(fixture()[0]) + decoder.decode(fixture()[1]) + with pytest.raises(decoder_module.StreamingDecodeError, match="discontinuous"): + decoder.decode(fixture(dts=200)[1]) + assert len(codec.calls) == 1 and decoder.failed + + +@pytest.mark.parametrize("case", ["no-init", "reinit", "non-keyframe", "oversized", "version"]) +def test_invalid_stream_setup_fails_closed(fake_av, case): + _, av, _ = fake_av + if case == "version": + av.__version__ = "other" + with pytest.raises(decoder_module.StreamingDecodeError, match="version"): + decoder_module.FragmentDecoder() + return + decoder = decoder_module.FragmentDecoder() + if case != "no-init": + decoder.configure(fixture()[0]) + with pytest.raises(decoder_module.StreamingDecodeError): + if case == "reinit": + decoder.configure(fixture()[0]) + else: + decoder.decode( + b"x" * (1024 * 1024 + 1) + if case == "oversized" + else fixture(sync=case != "non-keyframe")[1] + ) + assert decoder.failed