diff --git a/src/mcbebop/media/capture.py b/src/mcbebop/media/capture.py new file mode 100644 index 0000000..5e14af4 --- /dev/null +++ b/src/mcbebop/media/capture.py @@ -0,0 +1,280 @@ +"""Record the aircraft's RTP stream, and replay it packet for packet. + +Two halves of one idea. `record` binds the port we name in the ARSDK +handshake and writes every datagram the drone sends us, with the time it +arrived. `ReplaySource` feeds that file back to a client with the gaps it was +recorded with, which is the only way to put a viewer in front of the +aircraft's real pacing without the aircraft. + +That distinction is the point of having two paths at all. +`rtp.PacketisedSource` builds a stream from any video file and sends it at a +steady frame rate, which proves a decoder and a renderer work. A replay +reproduces bursts, reordering and jitter as they happened, which is what a +latency measurement needs. They are different instruments and a caller should +know which one it has. + +The file is deliberately dull: a magic string, a version, then a record per +packet. Nothing is parsed, nothing is dropped, and the payload bytes are kept +whole, so a later pass can answer questions the recorder never asked. Counting +the NAL types carried inside FU-A packets is the obvious one: the only +measurement we have of the drone's stream read the outer header byte of each +packet and stopped there, so what the fragments carried is still unknown. + +Recording needs the drone. Reading, writing and replaying a file do not, which +is what makes the format testable here. + + python -m mcbebop.media.capture out.rtpcap --seconds 10 +""" + +from __future__ import annotations + +import logging +import random +import socket +import statistics +import struct +import time +from collections.abc import Iterable, Iterator +from dataclasses import dataclass, field +from itertools import pairwise +from pathlib import Path + +from mcbebop.media.rtp import CLOCK_RATE, DEFAULT_FPS, parse_packet, resolve_start_offset, rewrite_header + +log = logging.getLogger(__name__) + +MAGIC = b"BEBOPRTP" +VERSION = 1 +_FILE_HEADER = struct.Struct("<8sHH") +_RECORD = struct.Struct(" int: + """Write `(arrival seconds, datagram)` pairs. Returns the packet count. + + Times are relative to the first packet, so a file is comparable with + itself regardless of when it was taken. + """ + out = Path(path) + out.parent.mkdir(parents=True, exist_ok=True) + count = 0 + with out.open("wb") as fh: + fh.write(_FILE_HEADER.pack(MAGIC, VERSION, 0)) + for at, data in packets: + fh.write(_RECORD.pack(at, len(data)) + data) + count += 1 + return count + + +def read_capture(path: str | Path) -> Iterator[tuple[float, bytes]]: + """Stream a capture back as `(arrival seconds, datagram)` pairs. + + A truncated tail stops the read rather than raising: a capture cut short + by Ctrl-C is still worth replaying, and the alternative is losing the + whole file to its last partial record. + """ + with Path(path).open("rb") as fh: + head = fh.read(_FILE_HEADER.size) + if len(head) < _FILE_HEADER.size: + raise CaptureFormatError(f"{path} is too short to be a capture") + magic, version, _flags = _FILE_HEADER.unpack(head) + if magic != MAGIC: + raise CaptureFormatError(f"{path} does not start with {MAGIC!r}") + if version != VERSION: + raise CaptureFormatError(f"{path} is capture version {version}, this reads {VERSION}") + while True: + raw = fh.read(_RECORD.size) + if len(raw) < _RECORD.size: + if raw: + log.warning("%s ends mid-record; stopping there", path) + return + at, length = _RECORD.unpack(raw) + data = fh.read(length) + if len(data) < length: + log.warning("%s ends mid-packet; stopping there", path) + return + yield at, data + + +def record( + path: str | Path, + *, + port: int = DEFAULT_CAPTURE_PORT, + host: str = "", + seconds: float | None = None, + max_packets: int | None = None, +) -> int: + """Bind `port` and write every datagram that arrives. Returns the count. + + Read-only with respect to the aircraft: it sends nothing, so it cannot be + the reason a stream stops. Something else has to hold the ARSDK session + and send VideoEnable, and it has to do it after this is listening, because + RTP is connectionless and whatever arrives before the bind is gone. + """ + sock = socket.socket(socket.AF_INET, socket.SOCK_DGRAM) + sock.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1) + sock.bind((host, port)) + sock.settimeout(_POLL) + deadline = None if seconds is None else time.monotonic() + seconds + first: float | None = None + packets: list[tuple[float, bytes]] = [] + log.info("recording RTP on %s:%d to %s", host or "0.0.0.0", port, path) + try: + while deadline is None or time.monotonic() < deadline: + try: + data, _addr = sock.recvfrom(_MAX_DATAGRAM) + except TimeoutError: + continue + now = time.monotonic() + first = now if first is None else first + packets.append((now - first, data)) + if max_packets is not None and len(packets) >= max_packets: + break + except KeyboardInterrupt: + log.info("interrupted; writing what arrived") + finally: + sock.close() + count = write_capture(path, packets) + span = packets[-1][0] if packets else 0.0 + log.info("wrote %d packets over %.2fs to %s", count, span, path) + return count + + +def _median_positive(values: Iterable[int | float], fallback: float) -> float: + positive = [v for v in values if v > 0] + return statistics.median(positive) if positive else fallback + + +@dataclass +class ReplaySource: + """A recorded RTP stream, sent again with its original inter-packet gaps. + + Only the SSRC and the sequence numbers are rewritten, so the client sees a + stream from us rather than one that looks like a replay of someone else's. + Timestamps keep their recorded deltas; what is added is a running offset, + so a loop carries the clock forward instead of jumping back to the start + of the capture and stalling every decoder downstream. + + The loop point is the one gap in the file that does not exist in it: the + interval between the last packet and a next one that was never recorded. + It is estimated as what would carry the stream to the start of the frame + after the last, from the capture's own mean frame period. The median gap + is the wrong estimator here and was the first thing tried: with a dozen + packets to a frame, most gaps are the ones *inside* a frame, so the median + is an intra-frame gap and the stream would loop a frame early every time. + """ + + packets_in: tuple[tuple[float, bytes], ...] + start_offset: float | str | None = None + seed: int | None = None + ssrc: int = field(default_factory=lambda: random.getrandbits(32)) + + @classmethod + def from_path( + cls, + path: str | Path, + *, + start_offset: float | str | None = None, + seed: int | None = None, + ) -> ReplaySource: + packets = tuple(read_capture(path)) + if not packets: + raise CaptureFormatError(f"{path} holds no packets") + return cls(packets_in=packets, start_offset=start_offset, seed=seed) + + def __post_init__(self) -> None: + # Both counters live on the instance, not in `packets()`, so a client + # that disables and re-enables video sees the stream continue rather + # than a clock that jumps backwards under an unchanged SSRC. + counters = random.Random(self.seed) + self._seq = counters.getrandbits(16) + self._ticks = counters.getrandbits(32) + # A separate generator, so the chosen offset does not move when the + # number of counter draws above changes. + offset = resolve_start_offset(self.start_offset, self.duration, random.Random(self.seed)) + self.start_index = next( + (i for i, (at, _) in enumerate(self.packets_in) if at >= offset), + 0, + ) + stamps = [parse_packet(data)[2] for _at, data in self.packets_in] + # Ticks step only between frames, so the median of the positive steps + # is one frame's worth however many packets a frame took. + ticks = [(b - a) % (1 << 32) for a, b in pairwise(stamps)] + self._loop_ticks = int(_median_positive(ticks, CLOCK_RATE / DEFAULT_FPS)) + self._loop_gap = self._gap_to_the_next_frame(stamps) + + def _gap_to_the_next_frame(self, stamps: list[int]) -> float: + """How long after the last recorded packet the next frame would start.""" + starts = [at for i, (at, _) in enumerate(self.packets_in) if i == 0 or stamps[i] != stamps[i - 1]] + if len(starts) < 2: + return 1.0 / DEFAULT_FPS + period = (starts[-1] - starts[0]) / (len(starts) - 1) + return max(0.0, period - (self.duration - starts[-1])) + + @property + def duration(self) -> float: + return self.packets_in[-1][0] if self.packets_in else 0.0 + + @property + def describe(self) -> str: + return ( + f"{len(self.packets_in)} recorded packets over {self.duration:.2f}s, " + f"replayed from index {self.start_index}" + ) + + def packets(self) -> Iterator[tuple[float, bytes]]: + total = len(self.packets_in) + index = self.start_index + prev_at, prev_ticks = None, None + while True: + at, data = self.packets_in[index] + stamp = parse_packet(data)[2] + if prev_at is None: + delay, step = 0.0, 0 + elif index == 0: + # The wrap: the file's own gaps say nothing about this one. + delay, step = self._loop_gap, self._loop_ticks + else: + delay, step = at - prev_at, (stamp - prev_ticks) % (1 << 32) + self._ticks = (self._ticks + step) % (1 << 32) + yield ( + max(0.0, delay), + rewrite_header(data, ssrc=self.ssrc, seq=self._seq, timestamp=self._ticks), + ) + self._seq = (self._seq + 1) & 0xFFFF + prev_at, prev_ticks = at, stamp + index = (index + 1) % total + + +def main(argv: list[str] | None = None) -> int: + import argparse + import sys + + parser = argparse.ArgumentParser( + prog="python -m mcbebop.media.capture", + description="Record the Bebop 2's RTP video stream for later replay.", + ) + parser.add_argument("out", help="where to write the capture") + parser.add_argument("--port", type=int, default=DEFAULT_CAPTURE_PORT) + parser.add_argument("--seconds", type=float, default=None, help="stop after this long") + parser.add_argument("--packets", type=int, default=None, help="stop after this many packets") + args = parser.parse_args(argv) + + # stderr, because this module is importable from the MCP server and + # stdout there is the JSON-RPC transport. + logging.basicConfig(level=logging.INFO, format="%(message)s", stream=sys.stderr) + count = record(args.out, port=args.port, seconds=args.seconds, max_packets=args.packets) + return 0 if count else 1 + + +if __name__ == "__main__": # pragma: no cover - a hand-run tool + raise SystemExit(main()) diff --git a/src/mcbebop/media/rtp.py b/src/mcbebop/media/rtp.py new file mode 100644 index 0000000..d6981df --- /dev/null +++ b/src/mcbebop/media/rtp.py @@ -0,0 +1,397 @@ +"""RTP/H.264 packetisation, so the simulator can emit a video stream. + +The Bebop 2 does not serve RTSP. It pushes plain RTP to the port the +controller named in its handshake, which means a simulator can imitate the +whole video path with a UDP socket and nothing else. This module is the +packetiser half of that: Annex-B H.264 in, RFC 6184 datagrams out. Pure +stdlib, because the simulator ships in the package and must not drag a media +library in behind it. + +What is implemented, and what is not: + +Single-NAL packets when a NAL fits the MTU, FU-A fragmentation when it does +not, and STAP-A aggregation for runs of parameter sets and SEI. Nothing else. +In particular a VCL NAL is never aggregated, so an access unit that fits in +one datagram still goes out as one single-NAL packet rather than a STAP-A. +Interleaved mode (types 25 to 29 other than 28), MTAP and the RTCP side are +absent; a viewer that needs RTCP from the simulator needs a capture replay. + +The live aircraft's packet-type mix is not reproduced and cannot be from +first principles. Six seconds off the drone carried 1461 packets whose outer +NAL byte was STAP-A, 740 FU-A and 200 non-IDR, which is around eight STAP-A +per frame at 30 fps and far more aggregation than parameter sets alone can +explain. Those are *outer* header bytes: for a FU-A packet the type lives in +the next byte and was never read, so what the carried NALs were is unknown. +Closing that gap needs a real capture, which is what `media/capture.py` and +`ReplaySource` exist for. + +Repeating the parameter sets is the one behaviour here that is not optional. +The aircraft's SPS and PPS recur throughout the stream rather than appearing +once at the start, which is why a viewer can join a flight already in +progress and recover in a second or two. `PacketisedSource` does the same. +""" + +from __future__ import annotations + +import logging +import random +import struct +from collections.abc import Iterator, Sequence +from dataclasses import dataclass, field +from pathlib import Path +from typing import Protocol + +log = logging.getLogger(__name__) + +PAYLOAD_TYPE = 96 # dynamic, bound to H264/90000 by the SDP +CLOCK_RATE = 90_000 +DEFAULT_MTU = 1400 # payload bytes, so a 1500-byte link does not fragment +DEFAULT_FPS = 30.0 + +# version 2, no padding, no extension, no CSRC. +_V2 = 0x80 +_MARKER = 0x80 +_RTP_HEADER = struct.Struct(">BBHII") +_STAP_A_LENGTH = struct.Struct(">H") + +NAL_STAP_A = 24 +NAL_FU_A = 28 + +_VCL = frozenset(range(1, 6)) # 1 non-IDR .. 5 IDR +_SPS, _PPS = 7, 8 +# What is worth aggregating: parameter sets, SEI and the access unit +# delimiter. All small, all repeated, none of them a picture. +_AGGREGATABLE = frozenset({6, _SPS, _PPS, 9}) + +_START_CODE = b"\x00\x00\x01" + + +def nal_type(nal: bytes) -> int: + return nal[0] & 0x1F + + +def _nri(nal: bytes) -> int: + return nal[0] & 0x60 + + +def iter_nals(data: bytes) -> Iterator[bytes]: + """Split an Annex-B elementary stream into NAL units. + + Both start codes are accepted. The three-byte form is the one actually + searched for, and a fourth leading zero is simply a trailing zero of the + previous NAL, which is why trailing zeros are stripped: a decoder ignores + them but they would otherwise inflate every NAL's length and so change + where the MTU falls. + """ + start = data.find(_START_CODE) + if start < 0: + return + pos = start + len(_START_CODE) + while True: + nxt = data.find(_START_CODE, pos) + end = len(data) if nxt < 0 else nxt + nal = data[pos:end].rstrip(b"\x00") + if nal: + yield nal + if nxt < 0: + return + pos = nxt + len(_START_CODE) + + +def _opens_picture(nal: bytes) -> bool: + """Is this a VCL NAL whose slice starts at macroblock zero? + + `first_mb_in_slice` is the first ue(v) in the slice header, and ue(v) is + zero exactly when the first bit is set. That makes the boundary between + access units readable without a bitstream parser, and it keeps a + multi-slice picture in one access unit instead of splitting it into one + per slice, which would put the marker bit in the wrong places. + """ + return len(nal) > 1 and nal_type(nal) in _VCL and bool(nal[1] & 0x80) + + +@dataclass(frozen=True) +class AccessUnit: + """One decodable picture, with whatever headers precede it.""" + + nals: tuple[bytes, ...] + + @property + def carries_parameter_sets(self) -> bool: + return any(nal_type(n) in (_SPS, _PPS) for n in self.nals) + + @property + def is_idr(self) -> bool: + return any(nal_type(n) == 5 for n in self.nals) + + +def access_units(nals: Sequence[bytes]) -> list[AccessUnit]: + """Group NALs into access units. + + A picture-opening VCL NAL closes the previous unit, so parameter sets and + SEI attach to the picture they precede rather than to the one before. + """ + units: list[AccessUnit] = [] + current: list[bytes] = [] + have_picture = False + for nal in nals: + if _opens_picture(nal) and have_picture: + units.append(AccessUnit(tuple(current))) + current, have_picture = [], False + current.append(nal) + have_picture = have_picture or nal_type(nal) in _VCL + if current: + units.append(AccessUnit(tuple(current))) + return units + + +@dataclass(frozen=True) +class AnnexBStream: + """A parsed Annex-B file, ready to be looped.""" + + units: tuple[AccessUnit, ...] + parameter_sets: tuple[bytes, ...] + + @classmethod + def from_bytes(cls, data: bytes) -> AnnexBStream: + nals = list(iter_nals(data)) + if not nals: + raise ValueError("no NAL units found; is this an Annex-B elementary stream?") + # Keep the last of each kind: a stream whose resolution changes + # mid-file should be repeated with the sets that are in force, and + # x264 writes the same pair every time anyway. + latest: dict[int, bytes] = {} + for nal in nals: + kind = nal_type(nal) + if kind in (_SPS, _PPS): + latest[kind] = nal + sets = tuple(latest[k] for k in (_SPS, _PPS) if k in latest) + if not sets: + log.warning("the stream carries no SPS/PPS, so a viewer joining late cannot start") + return cls(units=tuple(access_units(nals)), parameter_sets=sets) + + @classmethod + def from_path(cls, path: str | Path) -> AnnexBStream: + return cls.from_bytes(Path(path).read_bytes()) + + +@dataclass +class Packetiser: + """Annex-B access units in, RTP datagrams out. + + Stateful on purpose. The sequence number, the timestamp and the SSRC + belong to a stream rather than to a frame, and a looping source must not + reset any of them: a decoder handed a timestamp that jumps backwards + treats the whole stream as corrupt. + """ + + mtu: int = DEFAULT_MTU + payload_type: int = PAYLOAD_TYPE + ssrc: int = field(default_factory=lambda: random.getrandbits(32)) + seq: int = field(default_factory=lambda: random.getrandbits(16)) + timestamp: int = field(default_factory=lambda: random.getrandbits(32)) + aggregate: bool = True + + def __post_init__(self) -> None: + # Two bytes of FU-A header, so an MTU below three could never carry a + # payload byte and would loop forever. + if self.mtu < 3: + raise ValueError(f"mtu must be at least 3 payload bytes, got {self.mtu}") + + def packetise(self, nals: Sequence[bytes], *, advance: int = 0) -> list[bytes]: + """Datagrams for one access unit, then move the clock on by `advance`. + + The marker bit lands on the last packet of the unit, which is how a + receiver knows the picture is complete without parsing it. + """ + payloads: list[bytes] = [] + for group in self._groups(nals): + if len(group) > 1: + payloads.append(self._stap_a(group)) + elif len(group[0]) <= self.mtu: + payloads.append(group[0]) + else: + payloads.extend(self._fu_a(group[0])) + + packets = [self._frame(p, marker=i == len(payloads) - 1) for i, p in enumerate(payloads)] + self.timestamp = (self.timestamp + advance) % (1 << 32) + return packets + + # -- grouping -------------------------------------------------------- + def _groups(self, nals: Sequence[bytes]) -> list[list[bytes]]: + """Runs of NALs that travel together. Singletons unless aggregating. + + A STAP-A of one NAL is legal but pointless, so a run of length one is + emitted as a single-NAL packet instead. + """ + if not self.aggregate: + return [[nal] for nal in nals] + groups: list[list[bytes]] = [] + run: list[bytes] = [] + used = 1 # the STAP-A header byte + for nal in nals: + cost = _STAP_A_LENGTH.size + len(nal) + if nal_type(nal) in _AGGREGATABLE and used + cost <= self.mtu: + run.append(nal) + used += cost + continue + if run: + groups.append(run) + run, used = [], 1 + if nal_type(nal) in _AGGREGATABLE and 1 + cost <= self.mtu: + run, used = [nal], 1 + cost + else: + groups.append([nal]) + if run: + groups.append(run) + return groups + + # -- packet shapes --------------------------------------------------- + def _frame(self, payload: bytes, *, marker: bool) -> bytes: + seq = self.seq + self.seq = (seq + 1) & 0xFFFF + second = (_MARKER if marker else 0) | self.payload_type + return _RTP_HEADER.pack(_V2, second, seq, self.timestamp, self.ssrc) + payload + + def _stap_a(self, nals: Sequence[bytes]) -> bytes: + # The aggregate's NRI is the highest of what it carries, so dropping + # it costs a receiver no more than dropping the most important NAL in + # it would have. + out = bytearray([max(_nri(n) for n in nals) | NAL_STAP_A]) + for nal in nals: + out += _STAP_A_LENGTH.pack(len(nal)) + nal + return bytes(out) + + def _fu_a(self, nal: bytes) -> list[bytes]: + indicator = _nri(nal) | NAL_FU_A + kind = nal_type(nal) + body = nal[1:] # the original header is rebuilt by the receiver + budget = self.mtu - 2 + chunks = [body[i : i + budget] for i in range(0, len(body), budget)] + out = [] + for i, chunk in enumerate(chunks): + flags = (0x80 if i == 0 else 0) | (0x40 if i == len(chunks) - 1 else 0) + out.append(bytes([indicator, flags | kind]) + chunk) + return out + + +def parse_packet(packet: bytes) -> tuple[int, int, int, int, bool, bytes]: + """Read an RTP packet back: payload type, seq, timestamp, ssrc, marker, payload. + + Here rather than in the tests because the capture tools and the replay + path need it too, and a second implementation would be a second chance to + get the field order wrong. + """ + if len(packet) < _RTP_HEADER.size: + raise ValueError(f"an RTP packet is at least {_RTP_HEADER.size} bytes, got {len(packet)}") + first, second, seq, timestamp, ssrc = _RTP_HEADER.unpack_from(packet) + if first >> 6 != 2: + raise ValueError(f"not RTP version 2: first byte {first:#04x}") + csrc = first & 0x0F + offset = _RTP_HEADER.size + 4 * csrc + return second & 0x7F, seq, timestamp, ssrc, bool(second & _MARKER), packet[offset:] + + +def rewrite_header(packet: bytes, *, ssrc: int, seq: int, timestamp: int) -> bytes: + """Restamp a recorded packet with our SSRC, sequence number and clock. + + The caller supplies the timestamp rather than an offset because a replay + that loops has to carry the recorded gaps forward past the loop point + instead of jumping back to where the capture started. `ReplaySource` + accumulates the recorded deltas to do that. + """ + first, second, _seq, _timestamp, _ssrc = _RTP_HEADER.unpack_from(packet) + head = _RTP_HEADER.pack(first, second, seq & 0xFFFF, timestamp & 0xFFFFFFFF, ssrc & 0xFFFFFFFF) + return head + packet[_RTP_HEADER.size :] + + +class VideoSource(Protocol): + """Where the simulator gets datagrams and when to send them. + + One pair per packet: how long to wait before sending it, and the bytes. + The generator is expected to be endless, so a viewer can be left running. + """ + + def packets(self) -> Iterator[tuple[float, bytes]]: ... + + @property + def describe(self) -> str: ... + + +def resolve_start_offset(offset: float | str | None, duration: float, rng: random.Random) -> float: + """Where in the stream to begin, in seconds. + + `"random"` exists to reproduce a viewer switched on while the drone is + already flying, which hands a stateless decoder a stream that begins + mid-GOP. The seed is the caller's so a failure can be run again. + """ + if offset is None: + return 0.0 + if isinstance(offset, str): + if offset != "random": + raise ValueError(f"start_offset must be a number or 'random', got {offset!r}") + return rng.uniform(0.0, duration) if duration > 0 else 0.0 + if offset < 0: + raise ValueError(f"start_offset must not be negative, got {offset}") + return offset % duration if duration > 0 else 0.0 + + +@dataclass +class PacketisedSource: + """An Annex-B file, packetised and paced at a nominal frame rate. + + Steady pacing: every frame goes out one period after the last and the + packets within a frame go out back to back. That is a synthetic + instrument. It is the right one for checking that a decoder and a + renderer work, and the wrong one for measuring latency or jitter against + the aircraft, which sends in bursts this does not imitate. Use + `ReplaySource` over a real capture for that. + """ + + stream: AnnexBStream + fps: float = DEFAULT_FPS + mtu: int = DEFAULT_MTU + parameter_set_period: int = 30 + start_offset: float | str | None = None + seed: int | None = None + + def __post_init__(self) -> None: + if self.fps <= 0: + raise ValueError(f"fps must be positive, got {self.fps}") + if not self.stream.units: + raise ValueError("the stream has no access units") + self.packetiser = Packetiser(mtu=self.mtu) + self._rng = random.Random(self.seed) + self.start_index = round(self._offset_seconds() * self.fps) % len(self.stream.units) + + def _offset_seconds(self) -> float: + return resolve_start_offset(self.start_offset, len(self.stream.units) / self.fps, self._rng) + + @property + def describe(self) -> str: + return ( + f"{len(self.stream.units)} access units paced at {self.fps:g} fps from index {self.start_index}" + ) + + def packets(self) -> Iterator[tuple[float, bytes]]: + units = self.stream.units + ticks = round(CLOCK_RATE / self.fps) + period = 1.0 / self.fps + index = self.start_index + frame = 0 + while True: + unit = units[index] + nals = list(unit.nals) + # Nothing else here lets a viewer join a stream in progress: it + # cannot decode a picture whose parameter sets it never saw. + if ( + self.parameter_set_period > 0 + and frame % self.parameter_set_period == 0 + and not unit.carries_parameter_sets + ): + nals = [*self.stream.parameter_sets, *nals] + datagrams = self.packetiser.packetise(nals, advance=ticks) + for i, datagram in enumerate(datagrams): + yield (period if i == 0 else 0.0), datagram + index = (index + 1) % len(units) + frame += 1 diff --git a/tests/test_rtp.py b/tests/test_rtp.py new file mode 100644 index 0000000..79727e8 --- /dev/null +++ b/tests/test_rtp.py @@ -0,0 +1,475 @@ +"""The RTP packetiser and the capture format, against bytes. + +Synthetic NALs throughout: a NAL is a header byte and a payload as far as +RFC 6184 is concerned, so real H.264 would only make the assertions harder to +read. The one test that needs real video is the ffmpeg decode in +`test_sim_video.py`, which is the only one that can tell whether any of this +is actually correct. +""" + +import random +import struct +from itertools import pairwise + +import pytest + +from mcbebop.media import capture, rtp + +# header bytes: NRI 3 is "most important", which parameter sets carry. +SPS = bytes([0x67]) + b"\x42\xc0\x1e" +PPS = bytes([0x68]) + b"\xce\x3c\x80" +SEI = bytes([0x06]) + b"\x05\x02\x00\x00\x80" +# A slice whose first bit is set, so first_mb_in_slice is 0 and it opens a +# picture. 0x65 is an IDR at NRI 3, 0x41 a non-IDR at NRI 2. +IDR = bytes([0x65, 0x88]) + b"\xaa" * 40 +SLICE = bytes([0x41, 0x9A]) + b"\xbb" * 40 +# first bit clear, so first_mb_in_slice is not 0: a continuation slice. +SLICE_2 = bytes([0x41, 0x1A]) + b"\xcc" * 40 + + +def annex_b(*nals: bytes, four_byte: bool = False) -> bytes: + code = b"\x00\x00\x00\x01" if four_byte else b"\x00\x00\x01" + return b"".join(code + nal for nal in nals) + + +# -- Annex-B parsing ----------------------------------------------------- +def test_three_and_four_byte_start_codes_both_split(): + assert list(rtp.iter_nals(annex_b(SPS, PPS))) == [SPS, PPS] + assert list(rtp.iter_nals(annex_b(SPS, PPS, four_byte=True))) == [SPS, PPS] + + +def test_trailing_zeros_belong_to_the_start_code_not_the_nal(): + # A fourth zero before 00 00 01 is a trailing byte of the previous NAL. + # Counting it would change the length and so change where the MTU falls. + data = b"\x00\x00\x01" + SPS + b"\x00\x00\x00\x00\x01" + PPS + assert list(rtp.iter_nals(data)) == [SPS, PPS] + + +def test_a_stream_with_no_start_code_yields_nothing(): + assert list(rtp.iter_nals(b"\xde\xad\xbe\xef")) == [] + with pytest.raises(ValueError, match="Annex-B"): + rtp.AnnexBStream.from_bytes(b"\xde\xad\xbe\xef") + + +def test_leading_garbage_before_the_first_start_code_is_skipped(): + assert list(rtp.iter_nals(b"junk" + annex_b(SPS))) == [SPS] + + +def test_nal_type_reads_the_low_five_bits(): + assert rtp.nal_type(SPS) == 7 + assert rtp.nal_type(PPS) == 8 + assert rtp.nal_type(IDR) == 5 + assert rtp.nal_type(SLICE) == 1 + + +# -- access units -------------------------------------------------------- +def test_parameter_sets_attach_to_the_picture_that_follows_them(): + units = rtp.access_units([SPS, PPS, IDR, SLICE]) + assert [u.nals for u in units] == [(SPS, PPS, IDR), (SLICE,)] + assert units[0].is_idr and units[0].carries_parameter_sets + assert not units[1].is_idr and not units[1].carries_parameter_sets + + +def test_a_continuation_slice_stays_in_the_same_access_unit(): + # Splitting a multi-slice picture per slice would put the marker bit and + # the timestamp step in the wrong places. + units = rtp.access_units([IDR, SLICE_2, SLICE]) + assert [u.nals for u in units] == [(IDR, SLICE_2), (SLICE,)] + + +def test_the_stream_keeps_the_parameter_sets_for_reuse(): + stream = rtp.AnnexBStream.from_bytes(annex_b(SPS, PPS, IDR, SLICE, SLICE)) + assert stream.parameter_sets == (SPS, PPS) + assert len(stream.units) == 3 + + +def test_a_stream_without_parameter_sets_warns(caplog): + rtp.AnnexBStream.from_bytes(annex_b(IDR)) + assert "joining late" in caplog.text + + +# -- RTP headers --------------------------------------------------------- +def unpack(packet): + return rtp.parse_packet(packet) + + +def test_a_small_nal_becomes_exactly_one_single_nal_packet(): + p = rtp.Packetiser(mtu=1400, ssrc=0x11223344, seq=7, timestamp=9000) + packets = p.packetise([IDR]) + assert len(packets) == 1 + pt, seq, ts, ssrc, marker, payload = unpack(packets[0]) + assert (pt, seq, ts, ssrc, marker) == (96, 7, 9000, 0x11223344, True) + assert payload == IDR # the NAL travels whole, header byte included + + +def test_the_header_is_twelve_bytes_of_version_two(): + packets = rtp.Packetiser().packetise([IDR]) + assert packets[0][0] == 0x80 # V=2, no padding, no extension, no CSRC + assert len(packets[0]) == 12 + len(IDR) + + +def test_parse_packet_rejects_anything_that_is_not_version_two(): + with pytest.raises(ValueError, match="version 2"): + rtp.parse_packet(b"\x40" + b"\x00" * 15) + with pytest.raises(ValueError, match="at least"): + rtp.parse_packet(b"\x80\x60") + + +# -- FU-A ---------------------------------------------------------------- +def big_nal(size: int, header: int = 0x65) -> bytes: + return bytes([header]) + bytes((i % 251) + 1 for i in range(size - 1)) + + +def test_an_oversized_nal_fragments_into_well_formed_fu_a(): + nal = big_nal(4000) + p = rtp.Packetiser(mtu=1400) + packets = p.packetise([nal]) + payloads = [unpack(pkt)[5] for pkt in packets] + assert len(payloads) == 3 # 1398 payload bytes per fragment, 3999 to carry + + for payload in payloads: + assert payload[0] & 0x1F == rtp.NAL_FU_A + assert payload[0] & 0x60 == nal[0] & 0x60 # NRI is preserved + assert payload[1] & 0x1F == rtp.nal_type(nal) # the carried type + assert payload[1] & 0x20 == 0 # the reserved bit must be zero + + starts = [bool(p[1] & 0x80) for p in payloads] + ends = [bool(p[1] & 0x40) for p in payloads] + assert starts == [True, False, False] + assert ends == [False, False, True] + + +def test_fu_a_fragments_reassemble_into_the_original_nal(): + nal = big_nal(4000) + payloads = [unpack(p)[5] for p in rtp.Packetiser(mtu=1400).packetise([nal])] + rebuilt = bytes([(payloads[0][0] & 0xE0) | (payloads[0][1] & 0x1F)]) + rebuilt += b"".join(p[2:] for p in payloads) + assert rebuilt == nal + + +def test_no_fragment_exceeds_the_mtu(): + for mtu in (3, 50, 1400): + packets = rtp.Packetiser(mtu=mtu).packetise([big_nal(4000)]) + assert max(len(unpack(p)[5]) for p in packets) <= mtu + + +def test_a_nal_exactly_at_the_mtu_is_not_fragmented(): + nal = big_nal(100) + assert len(rtp.Packetiser(mtu=100, aggregate=False).packetise([nal])) == 1 + assert len(rtp.Packetiser(mtu=99, aggregate=False).packetise([nal])) == 2 + + +def test_an_mtu_too_small_to_carry_a_payload_byte_is_refused(): + with pytest.raises(ValueError, match="at least 3"): + rtp.Packetiser(mtu=2) + + +# -- STAP-A -------------------------------------------------------------- +def test_parameter_sets_aggregate_into_one_stap_a(): + packets = rtp.Packetiser(mtu=1400).packetise([SPS, PPS, SEI, IDR]) + assert len(packets) == 2 # one STAP-A for the headers, one for the picture + payload = unpack(packets[0])[5] + assert payload[0] & 0x1F == rtp.NAL_STAP_A + assert payload[0] & 0x60 == 0x60 # the highest NRI of what it carries + + carried, pos = [], 1 + while pos < len(payload): + (length,) = struct.unpack_from(">H", payload, pos) + carried.append(payload[pos + 2 : pos + 2 + length]) + pos += 2 + length + assert carried == [SPS, PPS, SEI] + + +def test_a_picture_is_never_aggregated(): + # A STAP-A carrying a slice would be legal, but it is not what this does, + # and the docstring says so. Two small slices stay two packets. + packets = rtp.Packetiser(mtu=1400).packetise([IDR, SLICE_2]) + assert len(packets) == 2 + assert all(unpack(p)[5][0] & 0x1F != rtp.NAL_STAP_A for p in packets) + + +def test_a_lone_parameter_set_goes_out_as_a_single_nal_not_a_stap_a(): + (packet,) = rtp.Packetiser(mtu=1400).packetise([SPS]) + assert unpack(packet)[5] == SPS + + +def test_aggregation_can_be_switched_off(): + packets = rtp.Packetiser(mtu=1400, aggregate=False).packetise([SPS, PPS, SEI, IDR]) + assert len(packets) == 4 + assert [unpack(p)[5] for p in packets] == [SPS, PPS, SEI, IDR] + + +def test_aggregation_stops_at_the_mtu_rather_than_overflowing(): + sets = [big_nal(300, header=0x68) for _ in range(5)] + packets = rtp.Packetiser(mtu=700).packetise(sets) + assert len(packets) == 3 # 2 + 2 + 1 at 302 bytes of cost each + assert all(len(unpack(p)[5]) <= 700 for p in packets) + + +# -- sequence, marker, clock --------------------------------------------- +def test_sequence_numbers_advance_by_one_per_packet(): + p = rtp.Packetiser(mtu=1400, seq=100) + packets = p.packetise([SPS, PPS, IDR]) + p.packetise([SLICE]) + assert [unpack(pkt)[1] for pkt in packets] == [100, 101, 102] + assert p.seq == 103 + + +def test_sequence_numbers_wrap_at_sixteen_bits(): + p = rtp.Packetiser(mtu=1400, seq=0xFFFE) + packets = p.packetise([IDR]) + p.packetise([IDR]) + p.packetise([IDR]) + assert [unpack(pkt)[1] for pkt in packets] == [0xFFFE, 0xFFFF, 0] + + +def test_the_marker_bit_lands_only_on_the_last_packet_of_an_access_unit(): + p = rtp.Packetiser(mtu=1400) + packets = p.packetise([SPS, PPS, big_nal(4000)]) + markers = [unpack(pkt)[4] for pkt in packets] + assert markers == [False, False, False, True] + + +def test_the_timestamp_is_one_per_access_unit_and_advances_between_them(): + p = rtp.Packetiser(mtu=1400, timestamp=1000) + first = p.packetise([SPS, PPS, big_nal(4000)], advance=3000) + second = p.packetise([SLICE], advance=3000) + assert {unpack(pkt)[2] for pkt in first} == {1000} + assert {unpack(pkt)[2] for pkt in second} == {4000} + + +def test_the_timestamp_wraps_at_thirty_two_bits(): + p = rtp.Packetiser(timestamp=(1 << 32) - 1000) + p.packetise([IDR], advance=3000) + assert p.timestamp == 2000 + + +def test_the_ssrc_is_the_same_on_every_packet(): + p = rtp.Packetiser(mtu=1400) + packets = p.packetise([SPS, PPS, big_nal(4000)]) + p.packetise([SLICE]) + assert len({unpack(pkt)[3] for pkt in packets}) == 1 + + +# -- start offsets ------------------------------------------------------- +def test_no_offset_starts_at_the_beginning(): + assert rtp.resolve_start_offset(None, 10.0, random.Random(0)) == 0.0 + + +def test_a_random_offset_lands_inside_the_stream_and_repeats_with_a_seed(): + first = rtp.resolve_start_offset("random", 10.0, random.Random(5)) + second = rtp.resolve_start_offset("random", 10.0, random.Random(5)) + assert first == second + assert 0.0 <= first < 10.0 + + +def test_an_offset_past_the_end_wraps_rather_than_falling_off(): + assert rtp.resolve_start_offset(12.5, 10.0, random.Random(0)) == 2.5 + + +def test_a_nonsense_offset_is_an_error(): + with pytest.raises(ValueError, match="'random'"): + rtp.resolve_start_offset("middle", 10.0, random.Random(0)) + with pytest.raises(ValueError, match="negative"): + rtp.resolve_start_offset(-1.0, 10.0, random.Random(0)) + + +# -- the packetised source ---------------------------------------------- +@pytest.fixture +def stream(): + # Ten pictures, parameter sets only at the front, which is the case that + # makes repetition matter. + nals = [SPS, PPS, IDR] + [SLICE] * 9 + return rtp.AnnexBStream.from_bytes(annex_b(*nals)) + + +def take(source, n): + out = [] + for item in source.packets(): + out.append(item) + if len(out) >= n: + break + return out + + +def test_the_source_paces_one_period_per_frame_and_nothing_between_packets(stream): + source = rtp.PacketisedSource(stream, fps=30.0) + delays = [delay for delay, _ in take(source, 6)] + assert delays[0] == pytest.approx(1 / 30) + assert delays[1] == 0.0 # the second packet of the same frame + assert sum(1 for d in delays if d > 0) == len([d for d in delays if d > 0]) + + +def test_the_source_loops_without_the_clock_going_backwards(stream): + source = rtp.PacketisedSource(stream, fps=30.0, parameter_set_period=0) + packets = [data for _delay, data in take(source, 40)] + stamps = [unpack(p)[2] for p in packets] + assert stamps == sorted(stamps), "a timestamp that goes back makes a decoder give up" + seqs = [unpack(p)[1] for p in packets] + assert seqs == list(range(seqs[0], seqs[0] + len(seqs))) + assert len(packets) > len(stream.units), "it must have passed the loop point" + + +def test_parameter_sets_recur_so_a_viewer_can_join_late(stream): + source = rtp.PacketisedSource(stream, fps=30.0, parameter_set_period=4) + payloads = [unpack(data)[5] for _delay, data in take(source, 30)] + stap_a = [p for p in payloads if p[0] & 0x1F == rtp.NAL_STAP_A] + # One at the file's own head plus one every fourth frame after it. + assert len(stap_a) >= 5 + assert all(SPS in p and PPS in p for p in stap_a) + + +def test_turning_repetition_off_sends_the_parameter_sets_once_per_pass(stream): + source = rtp.PacketisedSource(stream, fps=30.0, parameter_set_period=0) + payloads = [unpack(data)[5] for _delay, data in take(source, 25)] + assert sum(1 for p in payloads if p[0] & 0x1F == rtp.NAL_STAP_A) == 3 + + +def test_a_random_start_offset_begins_mid_stream_and_is_reproducible(stream): + a = rtp.PacketisedSource(stream, fps=30.0, start_offset="random", seed=11) + b = rtp.PacketisedSource(stream, fps=30.0, start_offset="random", seed=11) + assert a.start_index == b.start_index + assert 0 <= a.start_index < len(stream.units) + assert rtp.PacketisedSource(stream, fps=30.0, start_offset=0.1).start_index == 3 + + +def test_a_source_needs_a_positive_frame_rate(stream): + with pytest.raises(ValueError, match="fps"): + rtp.PacketisedSource(stream, fps=0) + + +# -- the capture format -------------------------------------------------- +def rtp_packet(seq, timestamp, ssrc=0xAABBCCDD, payload=b"\x41\x9a\xff"): + return struct.pack(">BBHII", 0x80, 96, seq, timestamp, ssrc) + payload + + +def test_a_capture_round_trips(tmp_path): + path = tmp_path / "c.rtpcap" + packets = [(0.0, rtp_packet(1, 9000)), (0.033, rtp_packet(2, 12000))] + assert capture.write_capture(path, packets) == 2 + assert list(capture.read_capture(path)) == packets + + +def test_a_capture_creates_its_parent_directory(tmp_path): + path = tmp_path / "deep" / "c.rtpcap" + capture.write_capture(path, [(0.0, rtp_packet(1, 0))]) + assert path.is_file() + + +def test_something_that_is_not_a_capture_is_rejected(tmp_path): + path = tmp_path / "nope.rtpcap" + path.write_bytes(b"not a capture at all") + with pytest.raises(capture.CaptureFormatError, match="does not start"): + list(capture.read_capture(path)) + short = tmp_path / "short.rtpcap" + short.write_bytes(b"BEBO") + with pytest.raises(capture.CaptureFormatError, match="too short"): + list(capture.read_capture(short)) + + +def test_a_future_capture_version_is_refused_rather_than_misread(tmp_path): + path = tmp_path / "future.rtpcap" + path.write_bytes(struct.pack("<8sHH", capture.MAGIC, 99, 0)) + with pytest.raises(capture.CaptureFormatError, match="version 99"): + list(capture.read_capture(path)) + + +def test_a_capture_cut_short_still_reads_up_to_the_break(tmp_path, caplog): + path = tmp_path / "cut.rtpcap" + capture.write_capture(path, [(0.0, rtp_packet(1, 0)), (0.033, rtp_packet(2, 3000))]) + path.write_bytes(path.read_bytes()[:-5]) # a Ctrl-C mid-write + got = list(capture.read_capture(path)) + assert len(got) == 1 + assert "mid-packet" in caplog.text + + +# -- replay -------------------------------------------------------------- +@pytest.fixture +def capture_file(tmp_path): + """Three frames of two packets each, with a gap only between frames.""" + packets = [] + for frame in range(3): + at = frame * 0.04 + packets.append((at, rtp_packet(frame * 2, 9000 + frame * 3000))) + packets.append((at + 0.001, rtp_packet(frame * 2 + 1, 9000 + frame * 3000))) + path = tmp_path / "flight.rtpcap" + capture.write_capture(path, packets) + return path + + +def test_a_replay_reproduces_the_recorded_gaps(capture_file): + source = capture.ReplaySource.from_path(capture_file) + delays = [delay for delay, _ in take(source, 6)] + assert delays[0] == 0.0 + assert delays[1] == pytest.approx(0.001) + assert delays[2] == pytest.approx(0.039) + assert delays[3] == pytest.approx(0.001) + + +def test_a_replay_restamps_the_ssrc_and_renumbers_continuously(capture_file): + source = capture.ReplaySource.from_path(capture_file, seed=3) + packets = [data for _delay, data in take(source, 10)] + assert {unpack(p)[3] for p in packets} == {source.ssrc} + assert source.ssrc != 0xAABBCCDD + seqs = [unpack(p)[1] for p in packets] + assert seqs == list(range(seqs[0], seqs[0] + len(seqs))) + + +def test_a_replay_keeps_the_recorded_timestamp_deltas(capture_file): + source = capture.ReplaySource.from_path(capture_file) + stamps = [unpack(data)[2] for _delay, data in take(source, 6)] + # Two packets per frame share a timestamp; frames are 3000 ticks apart. + assert [b - a for a, b in pairwise(stamps)] == [0, 3000, 0, 3000, 0] + + +def test_a_replay_carries_the_clock_past_the_loop_point(capture_file): + source = capture.ReplaySource.from_path(capture_file) + items = take(source, 14) # six recorded packets, then round again + stamps = [unpack(data)[2] for _delay, data in items] + assert stamps == sorted(stamps), "looping must not rewind the clock" + # The gap the file cannot know: filled with its own median, one frame. + assert items[6][0] == pytest.approx(0.039) + assert stamps[6] - stamps[5] == 3000 + + +def test_a_replay_can_start_part_way_in(capture_file): + source = capture.ReplaySource.from_path(capture_file, start_offset=0.04) + assert source.start_index == 2 + assert capture.ReplaySource.from_path(capture_file, start_offset="random", seed=2).start_index >= 0 + + +def test_an_empty_capture_is_an_error_not_a_silent_stall(tmp_path): + path = tmp_path / "empty.rtpcap" + capture.write_capture(path, []) + with pytest.raises(capture.CaptureFormatError, match="no packets"): + capture.ReplaySource.from_path(path) + + +def test_the_recorder_writes_what_arrives_on_the_port_it_binds(tmp_path): + import socket + import threading + + path = tmp_path / "live.rtpcap" + port = 0 + probe = socket.socket(socket.AF_INET, socket.SOCK_DGRAM) + probe.bind(("127.0.0.1", 0)) + port = probe.getsockname()[1] + probe.close() + + done = threading.Event() + count = [] + + def run(): + count.append(capture.record(path, port=port, host="127.0.0.1", max_packets=2)) + done.set() + + thread = threading.Thread(target=run, daemon=True) + thread.start() + sender = socket.socket(socket.AF_INET, socket.SOCK_DGRAM) + for i in range(2): + for _attempt in range(20): + sender.sendto(rtp_packet(i, i * 3000), ("127.0.0.1", port)) + if done.wait(0.05) or len(count) or i == 0: + break + done.wait(5) + sender.close() + assert count == [2] + got = list(capture.read_capture(path)) + assert len(got) == 2 + assert got[0][0] == 0.0 # times are relative to the first packet