# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
#     http://www.apache.org/licenses/LICENSE-2.0

"""MP4 indexing and in-memory episode synthesis primitives."""

from __future__ import annotations

import struct
from collections.abc import Callable, Iterable
from dataclasses import dataclass
from functools import cached_property
from typing import Any, Literal, overload

import numpy as np
from numpy.typing import NDArray

# First read when indexing a source MP4. A faststart movie box is usually a few hundred KiB;
# larger boxes cost one more exact read.
DEFAULT_HEADER_PROBE_BYTES = 512 * 1024


@dataclass(frozen=True)
class Box:
    """Top-level or nested ISO BMFF box bounds."""

    type: bytes
    start: int
    header_size: int
    end: int

    @property
    def payload_start(self) -> int:
        """Return the absolute offset immediately after the box header."""
        return self.start + self.header_size

    @property
    def size(self) -> int:
        """Return the total box size in bytes."""
        return self.end - self.start


@dataclass(frozen=True)
class Mp4SampleSlice:
    """Contiguous source-byte range covering an indexed sample interval."""

    sample_lo: int
    sample_hi: int
    byte_offset: int
    byte_length: int
    source_start_pts: float


@dataclass(frozen=True)
class Mp4Index:
    """Serializable video-track sample table for one MP4 source file."""

    file_path: str
    file_size: int
    ftyp: bytes
    moov_offset: int
    mdat_offset: int
    mdat_payload_offset: int
    mdat_payload_size: int
    faststart: bool
    codec: str
    timescale: int
    duration: int
    track_id: int
    width: int
    height: int
    stsd_body: bytes
    sample_pts: NDArray[np.float64]
    sample_durations: NDArray[np.int64]
    sample_composition_offsets: NDArray[np.int64]
    sample_sizes: NDArray[np.int64]
    sample_offsets: NDArray[np.int64]
    sync_samples: NDArray[np.int64]

    @cached_property
    def _has_composition_offsets(self) -> bool:
        """Cache whether presentation order differs from decode order."""
        return bool(np.any(self.sample_composition_offsets))

    def sample_slice(
        self,
        from_ts: float,
        to_ts: float,
        *,
        keyframe_pad_s: float = 0.1,
        keyframe_pad_fraction: float = 0.05,
        file_size: int | None = None,
    ) -> Mp4SampleSlice:
        """Select a keyframe-aligned sample range around a timestamp span."""
        if to_ts < from_ts:
            raise ValueError(f"Invalid timestamp span: {from_ts=} {to_ts=}")
        if len(self.sample_pts) == 0:
            raise ValueError(f"{self.file_path} contains no indexed samples")

        pad = max(keyframe_pad_s, (to_ts - from_ts) * keyframe_pad_fraction)
        lo_ts = max(0.0, from_ts - pad)
        hi_ts = to_ts + pad
        if not self._has_composition_offsets:
            # Keep logarithmic lookup for the common no-reordering video path.
            lo = int(np.searchsorted(self.sample_pts, lo_ts, side="left"))
            hi = int(np.searchsorted(self.sample_pts, hi_ts, side="right")) - 1
            lo = min(max(lo, 0), len(self.sample_pts) - 1)
            hi = min(max(hi, lo), len(self.sample_pts) - 1)
            first_pts = float(self.sample_pts[lo])
        else:
            # Sample tables are in decode order, not presentation order (B-frames).
            selected = np.flatnonzero((self.sample_pts >= lo_ts) & (self.sample_pts <= hi_ts))
            if len(selected):
                lo, hi = int(selected[0]), int(selected[-1])
                first_pts = float(self.sample_pts[selected].min())
            else:
                lo = hi = int(np.argmin(np.abs(self.sample_pts - from_ts)))
                first_pts = float(self.sample_pts[lo])

        if len(self.sync_samples):
            # Open-GOP leading B-frames can be decoded after a keyframe while
            # being presented before it. Keep the preceding GOP in that case.
            prev_sync = self.sync_samples[
                (self.sync_samples <= lo) & (self.sample_pts[self.sync_samples] <= first_pts)
            ]
            if len(prev_sync):
                lo = int(prev_sync[-1])
            else:
                lo = int(self.sync_samples[0])
                if lo > hi:
                    hi = lo

        if self._has_composition_offsets:
            # Finish the GOP so later reference pictures do not leave holes in the
            # presentation sequence. Approximate frame-index decoders require a
            # complete, contiguous sequence, not just the requested B-frame packets.
            next_sync = self.sync_samples[self.sync_samples > hi]
            hi = int(next_sync[0]) - 1 if len(next_sync) else len(self.sample_pts) - 1

        offsets = self.sample_offsets[lo : hi + 1]
        sizes = self.sample_sizes[lo : hi + 1]
        slice_lo = int(offsets.min())
        slice_hi = int((offsets + sizes).max())
        if file_size is not None:
            slice_hi = min(slice_hi, int(file_size))
        return Mp4SampleSlice(
            sample_lo=lo,
            sample_hi=hi,
            byte_offset=slice_lo,
            byte_length=slice_hi - slice_lo,
            source_start_pts=float(self.sample_pts[lo]),
        )

    def to_dict(self) -> dict[str, Any]:
        """Serialize scalar MP4 metadata; sample arrays are stored separately."""
        return {
            "file_path": self.file_path,
            "file_size": self.file_size,
            "ftyp": self.ftyp.hex(),
            "moov_offset": self.moov_offset,
            "mdat_offset": self.mdat_offset,
            "mdat_payload_offset": self.mdat_payload_offset,
            "mdat_payload_size": self.mdat_payload_size,
            "faststart": self.faststart,
            "codec": self.codec,
            "timescale": self.timescale,
            "duration": self.duration,
            "track_id": self.track_id,
            "width": self.width,
            "height": self.height,
            "stsd_body": self.stsd_body.hex(),
        }

    @classmethod
    def from_dict(cls, data: dict[str, Any], arrays: dict[str, NDArray[Any]]) -> Mp4Index:
        """Reconstruct an MP4 index from scalar metadata and sample arrays."""
        return cls(
            file_path=data["file_path"],
            file_size=int(data["file_size"]),
            ftyp=bytes.fromhex(data["ftyp"]),
            moov_offset=int(data["moov_offset"]),
            mdat_offset=int(data["mdat_offset"]),
            mdat_payload_offset=int(data["mdat_payload_offset"]),
            mdat_payload_size=int(data["mdat_payload_size"]),
            faststart=bool(data["faststart"]),
            codec=data["codec"],
            timescale=int(data["timescale"]),
            duration=int(data["duration"]),
            track_id=int(data["track_id"]),
            width=int(data["width"]),
            height=int(data["height"]),
            stsd_body=bytes.fromhex(data["stsd_body"]),
            sample_pts=arrays["sample_pts"],
            sample_durations=arrays["sample_durations"],
            sample_composition_offsets=arrays["sample_composition_offsets"],
            sample_sizes=arrays["sample_sizes"],
            sample_offsets=arrays["sample_offsets"],
            sync_samples=arrays["sync_samples"],
        )


def fetch_mp4_index(
    path: str,
    read_range: Callable[[str, int, int], bytes],
    *,
    file_size: int,
    header_probe_bytes: int = DEFAULT_HEADER_PROBE_BYTES,
    max_probe_bytes: int = 64 * 1024 * 1024,
) -> Mp4Index:
    """Fetch enough source bytes to parse a complete MP4 video-track index.

    Box headers give exact sizes, so the prefix grows to the end of a truncated box instead of
    doubling, and a movie box stored after the media payload is read from the tail at once.
    """
    probe_limit = min(max_probe_bytes, file_size)
    data = read_range(path, 0, min(header_probe_bytes, probe_limit))
    while True:
        top = list(iter_boxes(data, 0, len(data), absolute_base=0, allow_truncated=True))
        has_mdat = any(box.type == b"mdat" for box in top)
        moov_box = next((box for box in top if box.type == b"moov"), None)
        has_moov = moov_box is not None and moov_box.end <= len(data)
        if has_mdat and has_moov:
            return parse_mp4_index(path, data, file_size=file_size)
        if has_mdat and moov_box is None:
            # The payload comes first: probing further would read media bytes, not the index.
            tail_index = _fetch_tail_moov_index(path, read_range, data, top, file_size, max_probe_bytes)
            if tail_index is not None:
                return tail_index
        if (has_mdat and moov_box is None) or len(data) >= probe_limit:
            missing = []
            if not has_mdat:
                missing.append("mdat")
            if not has_moov:
                missing.append("moov")
            raise ValueError(
                f"Could not find complete {'/'.join(missing)} in first {len(data)} bytes of {path}"
            )
        last_box = top[-1] if top else None
        # Read the rest of a truncated box plus the largest possible next box header.
        truncated = last_box is not None and last_box.end > len(data)
        target = last_box.end + 16 if truncated and last_box is not None else 2 * len(data)
        chunk = read_range(path, len(data), min(target, probe_limit) - len(data))
        if not chunk:
            raise ValueError(f"Empty range read at byte {len(data)} of {path}")
        data += chunk


def _fetch_tail_moov_index(
    path: str,
    read_range: Callable[[str, int, int], bytes],
    prefix: bytes,
    top_boxes: list[Box],
    file_size: int,
    max_probe_bytes: int,
) -> Mp4Index | None:
    """Probe after the media payload for a complete movie box, if available."""
    mdat_box = _one(top_boxes, b"mdat")
    if mdat_box is None or mdat_box.end >= file_size:
        return None
    tail_offset = mdat_box.end
    tail_length = min(max_probe_bytes, file_size - tail_offset)
    tail = read_range(path, tail_offset, tail_length)
    tail_boxes = list(iter_boxes(tail, 0, len(tail), absolute_base=tail_offset, allow_truncated=True))
    moov_box = next(
        (box for box in tail_boxes if box.type == b"moov" and box.end <= tail_offset + len(tail)), None
    )
    if moov_box is None:
        return None
    ftyp_box = _one(top_boxes, b"ftyp", required=False)
    ftyp = (
        prefix[ftyp_box.start : ftyp_box.end]
        if ftyp_box is not None
        else _box(b"ftyp", b"isom\0\0\2\0isomiso2mp41")
    )
    moov_start = moov_box.payload_start - tail_offset
    moov_end = moov_box.end - tail_offset
    return _parse_mp4_index_from_layout(
        path,
        file_size=file_size,
        ftyp=ftyp,
        moov_offset=moov_box.start,
        moov=tail[moov_start:moov_end],
        mdat_box=mdat_box,
    )


def parse_mp4_index(path: str, data: bytes, *, file_size: int | None = None) -> Mp4Index:
    """Parse an MP4 video-track index from an in-memory file prefix."""
    if file_size is None:
        file_size = len(data)
    top = list(iter_boxes(data, 0, len(data), absolute_base=0, allow_truncated=True))
    ftyp_box = _one(top, b"ftyp", required=False)
    moov_box = _one(top, b"moov")
    mdat_box = _one(top, b"mdat")
    if moov_box.end > len(data):
        raise ValueError(f"{path}: moov box is truncated")

    moov = data[moov_box.payload_start : moov_box.end]
    ftyp = (
        data[ftyp_box.start : ftyp_box.end]
        if ftyp_box is not None
        else _box(b"ftyp", b"isom\0\0\2\0isomiso2mp41")
    )
    return _parse_mp4_index_from_layout(
        path,
        file_size=file_size,
        ftyp=ftyp,
        moov_offset=moov_box.start,
        moov=moov,
        mdat_box=mdat_box,
    )


def _parse_mp4_index_from_layout(
    path: str,
    *,
    file_size: int,
    ftyp: bytes,
    moov_offset: int,
    moov: bytes,
    mdat_box: Box,
) -> Mp4Index:
    """Build a video-track index from movie metadata and media payload bounds."""
    mvhd_timescale, mvhd_duration = _parse_mvhd(_find_descendant(moov, [b"mvhd"]))
    trak_box, trak_payload = _find_video_trak(moov)
    _ = trak_box
    tkhd = _parse_tkhd(_find_descendant(trak_payload, [b"tkhd"]))
    mdhd_timescale, mdhd_duration = _parse_mdhd(_find_descendant(trak_payload, [b"mdia", b"mdhd"]))
    stbl = _find_descendant(trak_payload, [b"mdia", b"minf", b"stbl"])

    stsd = _find_child(stbl, b"stsd")
    stsd_body = stbl[stsd.payload_start : stsd.end]
    codec = _parse_stsd_codec(stsd_body)
    stts = _parse_stts(_payload(stbl, b"stts"))
    sample_sizes = _parse_stsz(_payload(stbl, b"stsz"))
    stsc = _parse_stsc(_payload(stbl, b"stsc"))
    chunk_offsets = _parse_chunk_offsets(stbl)
    sync_samples = _parse_stss(stbl, len(sample_sizes))

    sample_durations = _expand_stts(stts, len(sample_sizes))
    sample_pts_units: NDArray[np.int64] = np.empty(len(sample_durations), dtype=np.int64)
    if len(sample_durations):
        sample_pts_units[0] = 0
        if len(sample_durations) > 1:
            sample_pts_units[1:] = np.cumsum(sample_durations[:-1], dtype=np.int64)
    composition_offsets = _parse_ctts(stbl, len(sample_sizes))
    edit_offset = _parse_edit_offset(trak_payload, mvhd_timescale, mdhd_timescale)
    sample_pts = (sample_pts_units + composition_offsets).astype(np.float64) / float(mdhd_timescale)
    sample_pts += edit_offset
    sample_offsets = _sample_offsets(stsc, chunk_offsets, sample_sizes)

    return Mp4Index(
        file_path=path,
        file_size=file_size,
        ftyp=ftyp,
        moov_offset=moov_offset,
        mdat_offset=mdat_box.start,
        mdat_payload_offset=mdat_box.payload_start,
        mdat_payload_size=mdat_box.end - mdat_box.payload_start
        if mdat_box.end <= file_size
        else file_size - mdat_box.payload_start,
        faststart=moov_offset < mdat_box.start,
        codec=codec,
        timescale=mdhd_timescale,
        duration=mdhd_duration or mvhd_duration,
        track_id=tkhd["track_id"],
        width=tkhd["width"],
        height=tkhd["height"],
        stsd_body=stsd_body,
        sample_pts=sample_pts,
        sample_durations=sample_durations,
        sample_composition_offsets=composition_offsets,
        sample_sizes=sample_sizes,
        sample_offsets=sample_offsets,
        sync_samples=sync_samples,
    )


def synthesize_mp4(index: Mp4Index, sample_slice: Mp4SampleSlice, mdat_payload: bytes) -> bytes:
    """Build a seekable episode-local MP4 around one indexed media slice."""
    lo = sample_slice.sample_lo
    hi = sample_slice.sample_hi + 1
    if lo < 0 or hi > len(index.sample_sizes) or lo >= hi:
        raise ValueError(f"Invalid sample range [{lo}, {hi}) for {index.file_path}")

    offsets = index.sample_offsets[lo:hi]
    sizes = index.sample_sizes[lo:hi]
    rel_offsets: NDArray[np.int64] = offsets - sample_slice.byte_offset
    if int(rel_offsets.min()) != 0:
        raise ValueError("Sample slice must start at the minimum referenced sample offset")
    if int((rel_offsets + sizes).max()) > len(mdat_payload):
        raise ValueError("Sample slice does not cover all referenced samples")

    durations = index.sample_durations[lo:hi]
    sync = index.sync_samples[(index.sync_samples >= lo) & (index.sync_samples < hi)] - lo + 1
    composition_offsets = index.sample_composition_offsets[lo:hi]
    moov = _make_moov(index, durations, sizes, rel_offsets, sync, composition_offsets, mdat_data_offset=0)
    header_size = len(index.ftyp) + len(moov)
    mdat_header_size = 8 if len(mdat_payload) + 8 <= 0xFFFFFFFF else 16
    moov = _make_moov(
        index,
        durations,
        sizes,
        rel_offsets,
        sync,
        composition_offsets,
        mdat_data_offset=header_size + mdat_header_size,
    )
    return index.ftyp + moov + _box(b"mdat", mdat_payload)


def synthesized_mp4_size(index: Mp4Index, sample_slice: Mp4SampleSlice) -> int:
    """Return the exact synthesized mini-MP4 size without fetching its media payload."""
    lo = sample_slice.sample_lo
    hi = sample_slice.sample_hi + 1
    if lo < 0 or hi > len(index.sample_sizes) or lo >= hi:
        raise ValueError(f"Invalid sample range [{lo}, {hi}) for {index.file_path}")

    offsets = index.sample_offsets[lo:hi]
    sizes = index.sample_sizes[lo:hi]
    rel_offsets: NDArray[np.int64] = offsets - sample_slice.byte_offset
    if int(rel_offsets.min()) != 0:
        raise ValueError("Sample slice must start at the minimum referenced sample offset")
    if int((rel_offsets + sizes).max()) > sample_slice.byte_length:
        raise ValueError("Sample slice does not cover all referenced samples")

    durations = index.sample_durations[lo:hi]
    sync = index.sync_samples[(index.sync_samples >= lo) & (index.sync_samples < hi)] - lo + 1
    composition_offsets = index.sample_composition_offsets[lo:hi]
    sync = _safe_sync_samples(durations, sync, composition_offsets)
    duration = int(durations.sum())
    sample_count = len(sizes)

    def box_size(payload_size: int) -> int:
        return payload_size + (8 if payload_size + 8 <= 0xFFFFFFFF else 16)

    def timing_size(values: NDArray[np.int64]) -> int:
        runs = 1 + int(np.count_nonzero(values[1:] != values[:-1]))
        return box_size(8 + 8 * runs)

    # Count the same boxes as _make_moov without packing a Python object per
    # sample. Admission needs their lengths, not serialized copies of the tables.
    tables = (
        box_size(len(index.stsd_body))
        + timing_size(durations)
        + (timing_size(composition_offsets) if np.any(composition_offsets) else 0)
        + len(_stsc_one_sample_per_chunk(sample_count))
        + box_size(12 + 4 * sample_count)
        + (box_size(8 + 4 * len(sync)) if len(sync) else 0)
    )
    media_headers = len(_mdhd(index.timescale, duration)) + len(_hdlr())
    track_header = len(_tkhd(index.track_id, duration, index.width, index.height))
    movie_header = len(_mvhd(index.timescale, duration, index.track_id + 1))
    edit_size = 44 if int(composition_offsets[0]) else 0

    def moov_size(offset_width: int) -> int:
        stbl = box_size(tables + box_size(8 + offset_width * sample_count))
        minf = box_size(len(_vmhd()) + len(_dinf()) + stbl)
        mdia = box_size(media_headers + minf)
        return box_size(movie_header + box_size(track_header + edit_size + mdia))

    max_offset = int(rel_offsets.max())
    first_size = moov_size(8 if max_offset > 0xFFFFFFFF else 4)
    mdat_header_size = box_size(sample_slice.byte_length) - sample_slice.byte_length
    data_offset = len(index.ftyp) + first_size + mdat_header_size
    final_size = moov_size(8 if max_offset + data_offset > 0xFFFFFFFF else 4)
    return len(index.ftyp) + final_size + mdat_header_size + sample_slice.byte_length


def iter_boxes(
    data: bytes,
    start: int,
    end: int,
    *,
    absolute_base: int = 0,
    allow_truncated: bool = False,
) -> Iterable[Box]:
    """Iterate ISO BMFF boxes within a byte interval."""
    pos = start
    while pos + 8 <= end:
        size = struct.unpack_from(">I", data, pos)[0]
        typ = data[pos + 4 : pos + 8]
        header_size = 8
        if size == 1:
            if pos + 16 > end:
                break
            size = struct.unpack_from(">Q", data, pos + 8)[0]
            header_size = 16
        elif size == 0:
            size = end - pos
        if size < header_size:
            break
        box_end = pos + size
        if box_end > end and not allow_truncated:
            break
        yield Box(typ, absolute_base + pos, header_size, absolute_base + box_end)
        pos = box_end


def _find_video_trak(moov: bytes) -> tuple[Box, bytes]:
    """Return the first video track and its payload, or raise if absent."""
    for trak in _children(moov, 0, len(moov)):
        if trak.type != b"trak":
            continue
        payload = moov[trak.payload_start : trak.end]
        hdlr = _find_descendant(payload, [b"mdia", b"hdlr"])
        if hdlr[8:12] == b"vide":
            return trak, payload
    raise ValueError("No video track found")


def _find_descendant(data: bytes, path: list[bytes]) -> bytes:
    """Follow a sequence of box types and return the innermost payload."""
    current = data
    for typ in path:
        box = _find_child(current, typ)
        current = current[box.payload_start : box.end]
    return current


def _find_child(data: bytes, typ: bytes) -> Box:
    """Find the first direct child of a given type, or raise if absent."""
    for box in _children(data, 0, len(data)):
        if box.type == typ:
            return box
    raise ValueError(f"Missing MP4 box {typ.decode('latin1')}")


def _children(data: bytes, start: int, end: int) -> Iterable[Box]:
    """Iterate child boxes using offsets relative to the supplied buffer."""
    return iter_boxes(data, start, end, absolute_base=0)


@overload
def _one(boxes: list[Box], typ: bytes, *, required: Literal[True] = True) -> Box: ...


@overload
def _one(boxes: list[Box], typ: bytes, *, required: Literal[False]) -> Box | None: ...


def _one(boxes: list[Box], typ: bytes, *, required: bool = True) -> Box | None:
    """Return the first matching box, raising when a required box is absent."""
    matches = [box for box in boxes if box.type == typ]
    if not matches and required:
        raise ValueError(f"Missing MP4 box {typ.decode('latin1')}")
    return matches[0] if matches else None


def _payload(parent: bytes, typ: bytes) -> bytes:
    """Return a required child box's contents without its header."""
    box = _find_child(parent, typ)
    return parent[box.payload_start : box.end]


def _parse_mvhd(payload: bytes) -> tuple[int, int]:
    """Read the movie timescale and duration in movie ticks."""
    version = payload[0]
    if version == 1:
        return struct.unpack_from(">IQ", payload, 20)
    return struct.unpack_from(">II", payload, 12)


def _parse_mdhd(payload: bytes) -> tuple[int, int]:
    """Read the media timescale and duration in media ticks."""
    version = payload[0]
    if version == 1:
        return struct.unpack_from(">IQ", payload, 20)
    return struct.unpack_from(">II", payload, 12)


def _parse_tkhd(payload: bytes) -> dict[str, int]:
    """Read the track identifier, duration and integer pixel dimensions."""
    version = payload[0]
    if version == 1:
        track_id = struct.unpack_from(">I", payload, 20)[0]
        duration = struct.unpack_from(">Q", payload, 28)[0]
        width, height = struct.unpack_from(">II", payload, 88)
    else:
        track_id = struct.unpack_from(">I", payload, 12)[0]
        duration = struct.unpack_from(">I", payload, 20)[0]
        width, height = struct.unpack_from(">II", payload, 76)
    return {"track_id": track_id, "duration": duration, "width": width >> 16, "height": height >> 16}


def _parse_stsd_codec(stsd_body: bytes) -> str:
    """Read the first sample entry's four-character codec identifier."""
    if len(stsd_body) < 16:
        return "unknown"
    return stsd_body[12:16].decode("latin1")


def _parse_stts(payload: bytes) -> list[tuple[int, int]]:
    """Read run-length encoded sample counts and decode durations."""
    count = struct.unpack_from(">I", payload, 4)[0]
    out = []
    offset = 8
    for _ in range(count):
        out.append(struct.unpack_from(">II", payload, offset))
        offset += 8
    return out


def _parse_ctts(stbl: bytes, sample_count: int) -> NDArray[np.int64]:
    """Expand composition offsets to signed media ticks for every sample."""
    box = _one(list(_children(stbl, 0, len(stbl))), b"ctts", required=False)
    if box is None:
        return np.zeros(sample_count, dtype=np.int64)
    payload = stbl[box.payload_start : box.end]
    version = payload[0]
    if version not in (0, 1):
        raise ValueError(f"Unsupported ctts version {version}")
    count = struct.unpack_from(">I", payload, 4)[0]
    entries = [
        struct.unpack_from(">Ii" if version == 1 else ">II", payload, 8 + idx * 8) for idx in range(count)
    ]
    return _expand_stts(entries, sample_count)


def _parse_edit_offset(trak: bytes, movie_timescale: int, media_timescale: int) -> float:
    """Resolve a rate-one edit timeline to a seconds offset, rejecting other layouts."""
    edts = _one(list(_children(trak, 0, len(trak))), b"edts", required=False)
    if edts is None:
        return 0.0
    payload = _find_descendant(trak[edts.payload_start : edts.end], [b"elst"])
    version = payload[0]
    if version not in (0, 1):
        raise ValueError(f"Unsupported elst version {version}")
    count = struct.unpack_from(">I", payload, 4)[0]
    entry_size = 20 if version == 1 else 12
    entries = [
        struct.unpack_from(">Qqhh" if version == 1 else ">Iihh", payload, 8 + idx * entry_size)
        for idx in range(count)
    ]
    # A single rate-one edit, optionally preceded by empty time, is the usual
    # encoder delay layout. Repeated/trimmed timelines need more than a constant
    # timestamp translation and must not silently produce incorrect images.
    empty_duration = 0
    if len(entries) == 2 and entries[0][1] == -1 and entries[0][2:] == (1, 0):
        empty_duration = entries.pop(0)[0]
    if len(entries) != 1 or entries[0][1] < 0 or entries[0][2:] != (1, 0):
        raise ValueError("Unsupported MP4 edit list: expected one rate-one media edit")
    return empty_duration / movie_timescale - entries[0][1] / media_timescale


def _expand_stts(entries: list[tuple[int, int]], sample_count: int) -> NDArray[np.int64]:
    """Expand timing runs and validate their total sample count."""
    values: NDArray[np.int64] = np.empty(sample_count, dtype=np.int64)
    pos = 0
    for count, delta in entries:
        values[pos : pos + count] = delta
        pos += count
    if pos != sample_count:
        raise ValueError(f"stts describes {pos} samples, stsz describes {sample_count}")
    return values


def _parse_stsz(payload: bytes) -> NDArray[np.int64]:
    """Read each compressed sample's size in bytes."""
    sample_size, sample_count = struct.unpack_from(">II", payload, 4)
    if sample_size:
        return np.full(sample_count, sample_size, dtype=np.int64)
    offset = 12
    values = np.empty(sample_count, dtype=np.int64)
    for idx in range(sample_count):
        values[idx] = struct.unpack_from(">I", payload, offset)[0]
        offset += 4
    return values


def _parse_stsc(payload: bytes) -> list[tuple[int, int, int]]:
    """Read chunk starts, samples per chunk and sample-description indices."""
    count = struct.unpack_from(">I", payload, 4)[0]
    out = []
    offset = 8
    for _ in range(count):
        out.append(struct.unpack_from(">III", payload, offset))
        offset += 12
    return out


def _parse_chunk_offsets(stbl: bytes) -> NDArray[np.int64]:
    """Read absolute chunk byte offsets from either 32-bit or 64-bit tables."""
    with_stco = None
    with_co64 = None
    for box in _children(stbl, 0, len(stbl)):
        if box.type == b"stco":
            with_stco = stbl[box.payload_start : box.end]
        elif box.type == b"co64":
            with_co64 = stbl[box.payload_start : box.end]
    if with_co64 is not None:
        count = struct.unpack_from(">I", with_co64, 4)[0]
        return np.array(
            [struct.unpack_from(">Q", with_co64, 8 + idx * 8)[0] for idx in range(count)], dtype=np.int64
        )
    if with_stco is None:
        raise ValueError("Missing stco/co64 chunk offsets")
    count = struct.unpack_from(">I", with_stco, 4)[0]
    return np.array(
        [struct.unpack_from(">I", with_stco, 8 + idx * 4)[0] for idx in range(count)], dtype=np.int64
    )


def _parse_stss(stbl: bytes, sample_count: int) -> NDArray[np.int64]:
    """Return zero-based keyframe indices, treating an absent table as all-sync."""
    for box in _children(stbl, 0, len(stbl)):
        if box.type == b"stss":
            payload = stbl[box.payload_start : box.end]
            count = struct.unpack_from(">I", payload, 4)[0]
            return np.array(
                [struct.unpack_from(">I", payload, 8 + idx * 4)[0] - 1 for idx in range(count)],
                dtype=np.int64,
            )
    return np.arange(sample_count, dtype=np.int64)


def _sample_offsets(
    stsc: list[tuple[int, int, int]], chunk_offsets: NDArray[np.int64], sample_sizes: NDArray[np.int64]
) -> NDArray[np.int64]:
    """Expand chunk tables into absolute byte offsets for individual samples."""
    if not stsc:
        raise ValueError("stsc is empty")
    offsets: NDArray[np.int64] = np.empty(len(sample_sizes), dtype=np.int64)
    sample_idx = 0
    for entry_idx, (first_chunk, samples_per_chunk, _desc_idx) in enumerate(stsc):
        next_first = stsc[entry_idx + 1][0] if entry_idx + 1 < len(stsc) else len(chunk_offsets) + 1
        for chunk_number in range(first_chunk, next_first):
            if chunk_number < 1 or chunk_number > len(chunk_offsets):
                raise ValueError("stsc references a chunk outside stco/co64")
            chunk_pos = int(chunk_offsets[chunk_number - 1])
            for _ in range(samples_per_chunk):
                if sample_idx >= len(sample_sizes):
                    return offsets
                offsets[sample_idx] = chunk_pos
                chunk_pos += int(sample_sizes[sample_idx])
                sample_idx += 1
    if sample_idx != len(sample_sizes):
        raise ValueError(f"stsc describes {sample_idx} samples, stsz describes {len(sample_sizes)}")
    return offsets


def _safe_sync_samples(
    durations: NDArray[np.int64],
    sync_samples: NDArray[np.int64],
    composition_offsets: NDArray[np.int64],
) -> NDArray[np.int64]:
    """Exclude internal open-GOP keyframes that would skip preceding pictures."""
    if np.any(composition_offsets) and len(sync_samples) > 1:
        # Approximate decoders seek using sync samples without scanning packets.
        # An internal open-GOP keyframe is later than its leading B-frames: do
        # not seek there and accidentally skip a requested preceding picture.
        presentation_times = np.cumsum(durations) - durations + composition_offsets
        safe_sync_samples = [int(sync_samples[0])]
        for position, sample in enumerate(sync_samples[1:], start=1):
            next_sample = (
                int(sync_samples[position + 1]) if position + 1 < len(sync_samples) else len(durations) + 1
            )
            if not np.any(presentation_times[sample : next_sample - 1] < presentation_times[sample - 1]):
                safe_sync_samples.append(int(sample))
        sync_samples = np.array(safe_sync_samples, dtype=np.int64)
    return sync_samples


def _make_moov(
    index: Mp4Index,
    durations: NDArray[np.int64],
    sizes: NDArray[np.int64],
    rel_offsets: NDArray[np.int64],
    sync_samples: NDArray[np.int64],
    composition_offsets: NDArray[np.int64],
    *,
    mdat_data_offset: int,
) -> bytes:
    """Build movie metadata for the slice, rebasing offsets and preserving presentation timing."""
    duration = int(durations.sum())
    sync_samples = _safe_sync_samples(durations, sync_samples, composition_offsets)
    stco_values = [int(mdat_data_offset + value) for value in rel_offsets]
    if any(value > 0xFFFFFFFF for value in stco_values):
        offset_box = _co64(stco_values)
    else:
        offset_box = _stco(stco_values)
    stbl = _box(
        b"stbl",
        _box(b"stsd", index.stsd_body)
        + _stts(durations)
        + (_ctts(composition_offsets) if np.any(composition_offsets) else b"")
        + _stsc_one_sample_per_chunk(len(sizes))
        + _stsz(sizes)
        + offset_box
        + (_stss(sync_samples) if len(sync_samples) else b""),
    )
    minf = _box(b"minf", _vmhd() + _dinf() + stbl)
    mdia = _box(b"mdia", _mdhd(index.timescale, duration) + _hdlr() + minf)
    # Preserve decode/presentation offsets, but put the first keyframe at local
    # time zero. The manifest stores its original presentation timestamp.
    media_time = int(composition_offsets[0])
    edit = (
        _box(b"edts", _full_box(b"elst", 1, 0, struct.pack(">IQqhh", 1, duration, media_time, 1, 0)))
        if media_time
        else b""
    )
    trak = _box(b"trak", _tkhd(index.track_id, duration, index.width, index.height) + edit + mdia)
    return _box(b"moov", _mvhd(index.timescale, duration, index.track_id + 1) + trak)


def _full_box(typ: bytes, version: int, flags: int, payload: bytes = b"") -> bytes:
    """Wrap a payload with an ISO BMFF version-and-flags header."""
    return _box(typ, bytes([version]) + flags.to_bytes(3, "big") + payload)


def _box(typ: bytes, payload: bytes) -> bytes:
    """Wrap a payload with a box header, using extended size when needed."""
    size = len(payload) + 8
    if size <= 0xFFFFFFFF:
        return struct.pack(">I4s", size, typ) + payload
    return struct.pack(">I4sQ", 1, typ, size + 8) + payload


def _mvhd(timescale: int, duration: int, next_track_id: int) -> bytes:
    """Encode the movie header using the supplied timescale and duration."""
    matrix = struct.pack(">9I", 0x00010000, 0, 0, 0, 0x00010000, 0, 0, 0, 0x40000000)
    payload = (
        struct.pack(">IIII", 0, 0, timescale, duration)
        + struct.pack(">IHH", 0x00010000, 0x0100, 0)
        + b"\0" * 8
        + matrix
        + b"\0" * 24
        + struct.pack(">I", next_track_id)
    )
    return _full_box(b"mvhd", 0, 0, payload)


def _tkhd(track_id: int, duration: int, width: int, height: int) -> bytes:
    """Encode an enabled video track header with fixed-point pixel dimensions."""
    matrix = struct.pack(">9I", 0x00010000, 0, 0, 0, 0x00010000, 0, 0, 0, 0x40000000)
    payload = (
        struct.pack(">IIIII", 0, 0, track_id, 0, duration)
        + b"\0" * 8
        + struct.pack(">hhhh", 0, 0, 0, 0)
        + matrix
        + struct.pack(">II", width << 16, height << 16)
    )
    return _full_box(b"tkhd", 0, 7, payload)


def _mdhd(timescale: int, duration: int) -> bytes:
    """Encode the media header with its timescale and duration."""
    return _full_box(b"mdhd", 0, 0, struct.pack(">IIIIH", 0, 0, timescale, duration, 0x55C4) + b"\0\0")


def _hdlr() -> bytes:
    """Encode a video handler descriptor."""
    return _full_box(b"hdlr", 0, 0, b"\0" * 4 + b"vide" + b"\0" * 12 + b"VideoHandler\0")


def _vmhd() -> bytes:
    """Encode the default video media header."""
    return _full_box(b"vmhd", 0, 1, struct.pack(">HHHH", 0, 0, 0, 0))


def _dinf() -> bytes:
    """Mark the synthesized media payload as self-contained."""
    url = _full_box(b"url ", 0, 1)
    dref = _full_box(b"dref", 0, 0, struct.pack(">I", 1) + url)
    return _box(b"dinf", dref)


def _stts(durations: NDArray[np.int64]) -> bytes:
    """Run-length encode sample durations into a decode-time table."""
    runs: list[list[int]] = []
    for duration in durations.tolist():
        if runs and runs[-1][1] == int(duration):
            runs[-1][0] += 1
        else:
            runs.append([1, int(duration)])
    payload = struct.pack(">I", len(runs)) + b"".join(
        struct.pack(">II", count, delta) for count, delta in runs
    )
    return _full_box(b"stts", 0, 0, payload)


def _ctts(offsets: NDArray[np.int64]) -> bytes:
    """Run-length encode composition offsets, using signed entries when needed."""
    runs: list[list[int]] = []
    for offset in offsets.tolist():
        if runs and runs[-1][1] == offset:
            runs[-1][0] += 1
        else:
            runs.append([1, offset])
    signed = bool(np.any(offsets < 0))
    payload = struct.pack(">I", len(runs)) + b"".join(
        struct.pack(">Ii" if signed else ">II", count, offset) for count, offset in runs
    )
    return _full_box(b"ctts", int(signed), 0, payload)


def _stsc_one_sample_per_chunk(sample_count: int) -> bytes:
    """Encode a sample-to-chunk table with one sample per chunk."""
    return _full_box(b"stsc", 0, 0, struct.pack(">IIII", 1, 1, 1, 1))


def _stsz(sizes: NDArray[np.int64]) -> bytes:
    """Encode variable compressed sample sizes."""
    return _full_box(
        b"stsz",
        0,
        0,
        struct.pack(">II", 0, len(sizes)) + b"".join(struct.pack(">I", int(size)) for size in sizes.tolist()),
    )


def _stco(values: list[int]) -> bytes:
    """Encode 32-bit absolute chunk byte offsets."""
    return _full_box(
        b"stco", 0, 0, struct.pack(">I", len(values)) + b"".join(struct.pack(">I", v) for v in values)
    )


def _co64(values: list[int]) -> bytes:
    """Encode 64-bit absolute chunk byte offsets."""
    return _full_box(
        b"co64", 0, 0, struct.pack(">I", len(values)) + b"".join(struct.pack(">Q", v) for v in values)
    )


def _stss(values: NDArray[np.int64]) -> bytes:
    """Encode one-based keyframe indices in a sync-sample table."""
    return _full_box(
        b"stss",
        0,
        0,
        struct.pack(">I", len(values)) + b"".join(struct.pack(">I", int(value)) for value in values.tolist()),
    )
