diff --git a/benchmarks/__main__.py b/benchmarks/__main__.py index 5647794..89d38a8 100644 --- a/benchmarks/__main__.py +++ b/benchmarks/__main__.py @@ -12,16 +12,23 @@ python -m benchmarks --only loopback # one group: copy, loopback, audio, slow, soak """ +from __future__ import annotations + import argparse import asyncio import datetime import pathlib +import shutil import subprocess import sys +from typing import TYPE_CHECKING from benchmarks import media from benchmarks.measure import machine +if TYPE_CHECKING: + from collections.abc import Collection + RESOLUTIONS = [(320, 240), (640, 480), (1280, 720), (1920, 1080)] GROUPS = ('copy', 'loopback', 'audio', 'slow', 'soak') @@ -31,9 +38,12 @@ def _mb(value: int) -> str: def _commit() -> str: + git = shutil.which('git') + if git is None: + return 'unknown' try: - commit = subprocess.check_output(['git', 'rev-parse', '--short', 'HEAD'], text=True).strip() - dirty = subprocess.check_output(['git', 'status', '--porcelain'], text=True).strip() + commit = subprocess.check_output([git, 'rev-parse', '--short', 'HEAD'], text=True).strip() + dirty = subprocess.check_output([git, 'status', '--porcelain'], text=True).strip() return commit + (' with uncommitted changes' if dirty else '') except (OSError, subprocess.CalledProcessError): return 'unknown' @@ -43,112 +53,149 @@ def _log(message: str) -> None: print(message, file=sys.stderr, flush=True) -async def run(groups, quick: bool) -> str: - seconds = 5 if quick else 30 - budget = 0.3 if quick else 1.0 - lines = [ - '# Benchmarks', - '', - f'{datetime.date.today().isoformat()}, {machine()}, commit {_commit()}. ' - f'Written by `python -m benchmarks{" --quick" if quick else ""}`.', - '', - 'Both peers run in one process on the machine, so encoding, decoding and Python share its cores. ' - 'CPU is of one core (100% is a core busy).', - '', +async def _copy(budget: float) -> list[str]: + _log('copy: VideoFrame.copy_to per format and size') + results = await media.copy_costs(RESOLUTIONS, budget=budget) + lines = ['## VideoFrame.copy_to', '', 'The cost of a copy of the planes (I420) or of a conversion to RGB.', ''] + lines += ['| Size | Format | ms per frame | Megapixels/s |', '| --- | --- | --- | --- |'] + lines += [ + f'| {r.width}x{r.height} | {r.format} | {r.milliseconds:.3f} | {r.megapixels_per_second:.0f} |' for r in results ] + construct = media.construct_cost(1920, 1080, budget=budget) + return [*lines, '', f'Creating a 1920x1080 I420 VideoFrame from bytes (a copy of them): {construct:.3f} ms.', ''] - if 'copy' in groups: - _log('copy: VideoFrame.copy_to per format and size') - results = media.copy_costs(RESOLUTIONS, budget=budget) - lines += ['## VideoFrame.copy_to', '', 'The cost of a copy of the planes (I420) or of a conversion to RGB.', ''] - lines += ['| Size | Format | ms per frame | Megapixels/s |', '| --- | --- | --- | --- |'] - for r in results: - lines.append( - f'| {r.width}x{r.height} | {r.format} | {r.milliseconds:.3f} | {r.megapixels_per_second:.0f} |' - ) - construct = media.construct_cost(1920, 1080, budget=budget) - lines += ['', f'Creating a 1920x1080 I420 VideoFrame from bytes (a copy of them): {construct:.3f} ms.', ''] - if 'loopback' in groups: - lines += [ - '## Video through a connection', - '', +async def _loopback(seconds: int) -> list[str]: + lines = [ + '## Video through a connection', + '', + ( f'Frames written to a VideoTrackGenerator at 30 fps for {seconds} s (after a warmup), sent over VP8, ' 'read from a MediaStreamTrackProcessor of the remote track (buffer of 1 frame). Latency is from the ' 'write to the read, matched by a frame number drawn in the pixels. Lag is how late the event loop runs ' - 'a 5 ms timer.', - '', + 'a 5 ms timer.' + ), + '', + ( '| Size | Delivered fps | Sent | Received | Received sizes | Dropped by the processor ' - '| Latency p50 / p95 (ms) | Loop lag p95 / max (ms) | CPU | RSS start / end (MB) |', - '| --- | --- | --- | --- | --- | --- | --- | --- | --- | --- |', - ] - for width, height in RESOLUTIONS: - _log(f'loopback: {width}x{height}') - r = await media.video_loopback(width, height, seconds=seconds) - lines.append( - f'| {width}x{height} | {r.delivered_fps:.1f} | {r.sent} | {r.received} | {r.received_sizes} ' - f'| {r.discarded} ' - f'| {r.latency_p50_ms:.0f} / {r.latency_p95_ms:.0f} | {r.lag_p95_ms:.1f} / {r.lag_max_ms:.1f} ' - f'| {r.usage.cpu_percent:.0f}% | {_mb(r.usage.rss_start)} / {_mb(r.usage.rss_end)} |' - ) - lines.append('') - - if 'audio' in groups: - audio_seconds = 5 if quick else 60 - _log(f'audio: 48 kHz stereo for {audio_seconds} s') - r = await media.audio_loopback(channels=2, seconds=audio_seconds) - lines += [ - '## Audio through a connection', - '', - f'10 ms chunks of 48 kHz stereo written to a MediaStreamTrackGenerator in real time for {audio_seconds} s, ' - 'sent over Opus, read from a MediaStreamTrackProcessor of the remote track.', - '', + '| Latency p50 / p95 (ms) | Loop lag p95 / max (ms) | CPU | RSS start / end (MB) |' + ), + '| --- | --- | --- | --- | --- | --- | --- | --- | --- | --- |', + ] + for width, height in RESOLUTIONS: + _log(f'loopback: {width}x{height}') + r = await media.VideoLoopback(width, height, seconds=seconds).run() + lines.append( + f'| {width}x{height} | {r.delivered_fps:.1f} | {r.sent} | {r.received} | {r.received_sizes} ' + f'| {r.discarded} ' + f'| {r.latency_p50_ms:.0f} / {r.latency_p95_ms:.0f} | {r.lag_p95_ms:.1f} / {r.lag_max_ms:.1f} ' + f'| {r.usage.cpu_percent:.0f}% | {_mb(r.usage.rss_start)} / {_mb(r.usage.rss_end)} |' + ) + return [*lines, ''] + + +async def _audio(seconds: int) -> list[str]: + _log(f'audio: 48 kHz stereo for {seconds} s') + r = await media.audio_loopback(channels=2, seconds=seconds) + return [ + '## Audio through a connection', + '', + ( + f'10 ms chunks of 48 kHz stereo written to a MediaStreamTrackGenerator in real time for {seconds} s, ' + 'sent over Opus, read from a MediaStreamTrackProcessor of the remote track.' + ), + '', + ( '| Written | Received | Chunks/s | Frames received | Dropped | Loop lag p95 (ms) | CPU ' - '| RSS start / end (MB) |', - '| --- | --- | --- | --- | --- | --- | --- | --- |', + '| RSS start / end (MB) |' + ), + '| --- | --- | --- | --- | --- | --- | --- | --- |', + ( f'| {r.written} | {r.received} | {r.chunks_per_second:.1f} | {r.received_frames} | {r.discarded} ' - f'| {r.lag_p95_ms:.1f} | {r.usage.cpu_percent:.0f}% | {_mb(r.usage.rss_start)} / {_mb(r.usage.rss_end)} |', - '', - ] + f'| {r.lag_p95_ms:.1f} | {r.usage.cpu_percent:.0f}% ' + f'| {_mb(r.usage.rss_start)} / {_mb(r.usage.rss_end)} |' + ), + '', + ] - if 'slow' in groups: - _log('slow: a consumer taking 100 ms per frame, 720p') - r = await media.video_loopback(1280, 720, seconds=seconds, consumer_delay=0.1, rss_every=1) - lines += [ - '## Slow consumer', - '', + +async def _slow(seconds: int) -> list[str]: + _log('slow: a consumer taking 100 ms per frame, 720p') + r = await media.VideoLoopback(1280, 720, seconds=seconds, consumer_delay=0.1, rss_every=1).run() + return [ + '## Slow consumer', + '', + ( f'720p at 30 fps read by a consumer that takes 100 ms per frame, for {seconds} s: the processor drops ' - 'the frames it can\'t deliver, so memory stays flat.', - '', - '| Delivered fps | Dropped by the processor | RSS start / end (MB) | RSS slope (MB/min) |', - '| --- | --- | --- | --- |', + "the frames it can't deliver, so memory stays flat." + ), + '', + '| Delivered fps | Dropped by the processor | RSS start / end (MB) | RSS slope (MB/min) |', + '| --- | --- | --- | --- |', + ( f'| {r.delivered_fps:.1f} | {r.discarded} | {_mb(r.usage.rss_start)} / {_mb(r.usage.rss_end)} ' - f'| {r.rss_slope:+.2f} |', - '', - ] + f'| {r.rss_slope:+.2f} |' + ), + '', + ] - if 'soak' in groups and not quick: - _log('soak: 720p for 10 minutes') - r = await media.video_loopback(1280, 720, seconds=600, rss_every=10) - lines += [ - '## Soak', - '', - '720p at 30 fps through a connection for 10 minutes, the resident memory sampled every 10 s.', - '', + +async def _soak() -> list[str]: + _log('soak: 720p for 10 minutes') + r = await media.VideoLoopback(1280, 720, seconds=600, rss_every=10).run() + return [ + '## Soak', + '', + '720p at 30 fps through a connection for 10 minutes, the resident memory sampled every 10 s.', + '', + ( '| Delivered fps | Latency p95 (ms) | CPU | RSS start / end (MB) | RSS slope (MB/min) ' - '| RSS slope, second half (MB/min) |', - '| --- | --- | --- | --- | --- | --- |', + '| RSS slope, second half (MB/min) |' + ), + '| --- | --- | --- | --- | --- | --- |', + ( f'| {r.delivered_fps:.1f} | {r.latency_p95_ms:.0f} | {r.usage.cpu_percent:.0f}% ' f'| {_mb(r.usage.rss_start)} / {_mb(r.usage.rss_end)} | {r.rss_slope:+.2f} ' - f'| {r.rss_slope_second_half:+.2f} |', - '', - 'RSS samples (s, MB): ' + ', '.join(f'{t:.0f}: {b / 1e6:.0f}' for t, b in r.rss), - '', - ] + f'| {r.rss_slope_second_half:+.2f} |' + ), + '', + 'RSS samples (s, MB): ' + ', '.join(f'{t:.0f}: {b / 1e6:.0f}' for t, b in r.rss), + '', + ] + + +async def run(groups: Collection[str], *, quick: bool) -> str: + """Runs the benchmark groups and returns their results as Markdown.""" + seconds = 5 if quick else 30 + lines = [ + '# Benchmarks', + '', + ( + f'{datetime.datetime.now().astimezone().date().isoformat()}, {machine()}, commit {_commit()}. ' + f'Written by `python -m benchmarks{" --quick" if quick else ""}`.' + ), + '', + ( + 'Both peers run in one process on the machine, so encoding, decoding and Python share its cores. ' + 'CPU is of one core (100% is a core busy).' + ), + '', + ] + if 'copy' in groups: + lines += await _copy(0.3 if quick else 1.0) + if 'loopback' in groups: + lines += await _loopback(seconds) + if 'audio' in groups: + lines += await _audio(5 if quick else 60) + if 'slow' in groups: + lines += await _slow(seconds) + if 'soak' in groups and not quick: + lines += await _soak() return '\n'.join(lines) -def main(): +def main() -> None: + """Runs the benchmarks of the command line and writes the results.""" parser = argparse.ArgumentParser(prog='python -m benchmarks', description=__doc__.split('\n')[0]) parser.add_argument('--quick', action='store_true', help='a few seconds of each benchmark, without the soak') parser.add_argument('--only', choices=GROUPS, action='append', help='run a group only (repeatable)') @@ -156,7 +203,7 @@ def main(): '--output', type=pathlib.Path, default=pathlib.Path(__file__).with_name('RESULTS.md'), help='the results file' ) args = parser.parse_args() - report = asyncio.run(run(args.only or GROUPS, args.quick)) + report = asyncio.run(run(args.only or GROUPS, quick=args.quick)) args.output.write_text(report + '\n') print(report) diff --git a/benchmarks/measure.py b/benchmarks/measure.py index ef1545b..3d9f5c8 100644 --- a/benchmarks/measure.py +++ b/benchmarks/measure.py @@ -7,21 +7,32 @@ """What the benchmarks measure besides media: event loop lag, CPU, memory, and the machine.""" +from __future__ import annotations + import array import asyncio +import contextlib import os import platform +import shutil import statistics import subprocess import sys import time from dataclasses import dataclass, field -from typing import List, Optional +from typing import TYPE_CHECKING from tests.helpers import rss_bytes +if TYPE_CHECKING: + from collections.abc import Sequence + from types import TracebackType + + from typing_extensions import Self + -def percentile(values: List[float], fraction: float) -> float: +def percentile(values: Sequence[float], fraction: float) -> float: + """The value below which the fraction of the values are, or NaN without values.""" if not values: return float('nan') ordered = sorted(values) @@ -29,53 +40,60 @@ def percentile(values: List[float], fraction: float) -> float: class LoopLag: - """How late the event loop runs a timer: a busy loop (like one blocked on the GIL) delays everything on it""" + """How late the event loop runs a timer: a busy loop (like one blocked on the GIL) delays everything on it.""" - def __init__(self, interval: float = 0.005): + def __init__(self, interval: float = 0.005) -> None: self.interval = interval # compact: a 5 ms probe collects 12000 a minute self.lags = array.array('d') - self._task: Optional[asyncio.Task] = None + self._task: asyncio.Task | None = None - async def _probe(self): + async def _probe(self) -> None: loop = asyncio.get_running_loop() while True: start = loop.time() await asyncio.sleep(self.interval) self.lags.append(max(0.0, loop.time() - start - self.interval)) - def __enter__(self): + def __enter__(self) -> Self: self._task = asyncio.ensure_future(self._probe()) return self - def __exit__(self, *exc_info): - self._task.cancel() + def __exit__( + self, exc_type: type[BaseException] | None, exc: BaseException | None, traceback: TracebackType | None + ) -> None: + if self._task: + self._task.cancel() @property def p95_ms(self) -> float: + """The 95th percentile lag.""" return percentile(self.lags, 0.95) * 1000 @property def max_ms(self) -> float: + """The largest lag.""" return max(self.lags, default=float('nan')) * 1000 @dataclass class Usage: - """CPU and memory over a span of time""" + """CPU and memory over a span of time.""" wall: float = 0 cpu: float = 0 rss_start: int = 0 rss_end: int = 0 - _started: tuple = field(default=(0.0, 0.0), repr=False) + _started: tuple[float, float] = field(default=(0.0, 0.0), repr=False) - def start(self) -> 'Usage': + def start(self) -> Usage: + """Starts the span.""" self.rss_start = rss_bytes() self._started = (time.perf_counter(), time.process_time()) return self - def stop(self) -> 'Usage': + def stop(self) -> Usage: + """Ends the span.""" wall, cpu = self._started self.wall = time.perf_counter() - wall self.cpu = time.process_time() - cpu @@ -84,13 +102,13 @@ def stop(self) -> 'Usage': @property def cpu_percent(self) -> float: - """Of one core: the process uses several threads (encoders, decoders, network)""" + """Of one core: the process uses several threads (encoders, decoders, network).""" return self.cpu / self.wall * 100 if self.wall else float('nan') -def slope_mb_per_minute(samples: List[tuple]) -> float: - """The trend of (seconds, bytes) samples, by least squares""" - if len(samples) < 2: +def slope_mb_per_minute(samples: Sequence[tuple[float, int]]) -> float: + """The trend of (seconds, bytes) samples, by least squares: NaN for less than two.""" + if not samples: return float('nan') xs = [t for t, _ in samples] ys = [b / 1e6 for _, b in samples] @@ -102,12 +120,12 @@ def slope_mb_per_minute(samples: List[tuple]) -> float: def machine() -> str: + """The CPU, the OS and Python of the machine.""" cpu = platform.processor() or platform.machine() - if sys.platform == 'darwin': - try: - cpu = subprocess.check_output(['sysctl', '-n', 'machdep.cpu.brand_string'], text=True).strip() - except (OSError, subprocess.CalledProcessError): - pass + sysctl = shutil.which('sysctl') if sys.platform == 'darwin' else None + if sysctl: + with contextlib.suppress(OSError, subprocess.CalledProcessError): + cpu = subprocess.check_output([sysctl, '-n', 'machdep.cpu.brand_string'], text=True).strip() return ( f'{cpu}, {os.cpu_count()} cores, {platform.system()} {platform.release()} ({platform.machine()}), ' f'Python {platform.python_version()}' diff --git a/benchmarks/media.py b/benchmarks/media.py index 7625f91..0d78dab 100644 --- a/benchmarks/media.py +++ b/benchmarks/media.py @@ -7,6 +7,8 @@ """Media through a connection in the same process: generated on one end, read with a processor on the other.""" +from __future__ import annotations + import array import asyncio import contextlib @@ -14,19 +16,25 @@ import math import time from dataclasses import dataclass, field -from typing import Dict, List, Optional +from typing import TYPE_CHECKING, Callable import webrtc from benchmarks.measure import LoopLag, Usage, percentile, slope_mb_per_minute from tests.helpers import connect, rss_bytes, wait_for_event +if TYPE_CHECKING: + from collections.abc import AsyncIterator, Awaitable, Iterable, Sequence + # the frame number, drawn as bits in blocks of luma at the top left of each frame BITS = 16 BLOCK = 16 +LUMA_ONE, LUMA_ZERO, LUMA_THRESHOLD = 235, 16, 125 @dataclass class VideoResult: + """The frames sent and received, their latency, and what it cost.""" + width: int height: int fps: float @@ -35,43 +43,48 @@ class VideoResult: received: int = 0 discarded: int = 0 # compact, as a soak collects thousands - latencies: array.array = field(default_factory=lambda: array.array('d')) + latencies: array.array[float] = field(default_factory=lambda: array.array('d')) lag_p95_ms: float = 0 lag_max_ms: float = 0 usage: Usage = field(default_factory=Usage) # (seconds, bytes) of the resident memory during the run - rss: List[tuple] = field(default_factory=list) + rss: list[tuple[float, int]] = field(default_factory=list) # the sizes of the frames received, as the encoder may scale them down - sizes: Dict[tuple, int] = field(default_factory=dict) + sizes: dict[tuple[int, int], int] = field(default_factory=dict) @property def delivered_fps(self) -> float: + """The frames received per second while measuring.""" return self.received / self.seconds @property def latency_p50_ms(self) -> float: + """The median latency of a frame.""" return percentile(self.latencies, 0.5) * 1000 @property def latency_p95_ms(self) -> float: + """The 95th percentile latency of a frame.""" return percentile(self.latencies, 0.95) * 1000 @property def rss_slope(self) -> float: + """The trend of the resident memory, in MB per minute.""" return slope_mb_per_minute(self.rss) @property def rss_slope_second_half(self) -> float: - """Once caches and pools of the process reached their size""" + """Once caches and pools of the process reached their size.""" return slope_mb_per_minute(self.rss[len(self.rss) // 2 :]) @property def received_sizes(self) -> str: + """The sizes of the frames received, the most frequent first.""" return ', '.join(f'{w}x{h}' for (w, h), _ in sorted(self.sizes.items(), key=lambda item: -item[1])) -def _frames(width: int, height: int, count: int = 10) -> List[bytearray]: - """I420 frames of a gradient with a bar moving across, so the encoder has motion to encode""" +def _frames(width: int, height: int, count: int = 10) -> list[bytearray]: + """I420 frames of a gradient with a bar moving across, so the encoder has motion to encode.""" frames = [] chroma = (width // 2) * (height // 2) row = bytes((x * 255 // width) for x in range(width)) @@ -88,7 +101,7 @@ def _frames(width: int, height: int, count: int = 10) -> List[bytearray]: def _draw_number(frame: bytearray, width: int, number: int) -> None: for bit in range(BITS): - value = 235 if number >> bit & 1 else 16 + value = LUMA_ONE if number >> bit & 1 else LUMA_ZERO block = bytes([value]) * BLOCK for y in range(BLOCK): start = y * width + bit * BLOCK @@ -100,13 +113,13 @@ def _read_number(luma: bytes, stride: int) -> int: for bit in range(BITS): # the center of the block, away from the blur of compression at its edges samples = [luma[y * stride + bit * BLOCK + x] for y in range(5, 11) for x in range(5, 11)] - if sum(samples) / len(samples) > 125: + if sum(samples) / len(samples) > LUMA_THRESHOLD: number |= 1 << bit return number @contextlib.asynccontextmanager -async def _connection(): +async def _connection() -> AsyncIterator[tuple[webrtc.RTCPeerConnection, webrtc.RTCPeerConnection]]: caller, callee = webrtc.RTCPeerConnection(), webrtc.RTCPeerConnection() try: yield caller, callee @@ -115,7 +128,13 @@ async def _connection(): callee.close() -async def _remote_track(caller, callee, track, max_bitrate: Optional[int] = None): +async def _remote_track( + caller: webrtc.RTCPeerConnection, + callee: webrtc.RTCPeerConnection, + track: webrtc.MediaStreamTrack, + *, + max_bitrate: int | None = None, +) -> webrtc.MediaStreamTrack: sender = caller.add_track(track) track_event = wait_for_event(callee, 'track', 30) await connect(caller, callee, 30) @@ -127,107 +146,143 @@ async def _remote_track(caller, callee, track, max_bitrate: Optional[int] = None return (await track_event).track -async def _measure(result, processor, write, read, measuring, done, warmup, seconds, sample=None) -> LoopLag: - """Runs the writer and the reader: warms up, measures for the seconds, then stops both""" - writing, reading = asyncio.ensure_future(write()), asyncio.ensure_future(read()) - await asyncio.sleep(warmup) - discarded = processor.discarded_frames - gc.collect() - sampling = asyncio.ensure_future(sample()) if sample else None - with LoopLag() as lag: - result.usage.start() - measuring.set() - await asyncio.sleep(seconds) - measuring.clear() - result.usage.stop() - done.set() - result.discarded = processor.discarded_frames - discarded - await asyncio.wait_for(writing, 10) - reading.cancel() - with contextlib.suppress(asyncio.CancelledError): - await reading - if sampling: - await sampling - return lag - - -async def video_loopback( - width: int, - height: int, - fps: float = 30, - seconds: float = 30, - warmup: float = 3, - consumer_delay: float = 0, - max_buffer_size: int = 1, - rss_every: Optional[float] = None, -) -> VideoResult: - """Frames written to a generator at a steady rate, read from the processor of the remote track""" - result = VideoResult(width, height, fps, seconds) - frames = _frames(width, height) - loop = asyncio.get_running_loop() - sent_at: Dict[int, float] = {} - measuring = asyncio.Event() - done = asyncio.Event() +@dataclass +class _Phases: + """A writer and a reader warm up, then count while measuring, then stop once done.""" - async with _connection() as (caller, callee): - generator = webrtc.VideoTrackGenerator() - # enough for the size, so the encoder doesn't drop frames for bitrate - remote = await _remote_track(caller, callee, generator.track, max_bitrate=int(width * height * fps * 0.2)) - processor = webrtc.MediaStreamTrackProcessor(remote, max_buffer_size=max_buffer_size) + warmup: float + seconds: float + measuring: asyncio.Event = field(default_factory=asyncio.Event) + done: asyncio.Event = field(default_factory=asyncio.Event) + + async def measure( + self, + result: VideoResult | AudioResult, + processor: webrtc.MediaStreamTrackProcessor, + *, + write: Callable[[], Awaitable[None]], + read: Callable[[], Awaitable[None]], + ) -> LoopLag: + """Runs the writer and the reader: warms up, measures for the seconds, then stops both.""" + writing, reading = asyncio.ensure_future(write()), asyncio.ensure_future(read()) + await asyncio.sleep(self.warmup) + discarded = processor.discarded_frames + gc.collect() + with LoopLag() as lag: + result.usage.start() + self.measuring.set() + await asyncio.sleep(self.seconds) + self.measuring.clear() + result.usage.stop() + self.done.set() + result.discarded = processor.discarded_frames - discarded + await asyncio.wait_for(writing, 10) + reading.cancel() + with contextlib.suppress(asyncio.CancelledError): + await reading + return lag - async def write(): - writer = generator.writable.get_writer() - start = loop.time() - number = 0 - while not done.is_set(): - frame = frames[number % len(frames)] - _draw_number(frame, width, number) - sent_at[number] = loop.time() - # frames lost on the way are forgotten - sent_at.pop(number - 300, None) - if measuring.is_set(): - result.sent += 1 - await writer.write( - webrtc.VideoFrame(frame, format='I420', coded_width=width, coded_height=height, timestamp=number) - ) - number += 1 - await asyncio.sleep(max(0.0, start + number / fps - loop.time())) - - async def read(): - header = {'rect': {'x': 0, 'y': 0, 'width': BITS * BLOCK, 'height': BLOCK}} - async for frame in processor.readable: - now = loop.time() - if measuring.is_set(): - size = (frame.coded_width, frame.coded_height) - result.sizes[size] = result.sizes.get(size, 0) + 1 - luma = bytearray(frame.allocation_size(header)) - # synchronous, without the future of copy_to - frame._copy_to(luma, header) - frame.close() - number = _read_number(luma, BITS * BLOCK) - if measuring.is_set() and number in sent_at: - result.received += 1 - result.latencies.append(now - sent_at.pop(number)) - if consumer_delay: - await asyncio.sleep(consumer_delay) - if done.is_set(): - break - async def sample_rss(): - start = loop.time() - while not done.is_set(): - result.rss.append((loop.time() - start, rss_bytes())) - await asyncio.sleep(rss_every) - - sample = sample_rss if rss_every else None - lag = await _measure(result, processor, write, read, measuring, done, warmup, seconds, sample) - result.lag_p95_ms, result.lag_max_ms = lag.p95_ms, lag.max_ms - generator.track.stop() - return result +@dataclass +class VideoLoopback: + """Frames written to a generator at a steady rate, read from the processor of the remote track.""" + + width: int + height: int + fps: float = 30 + seconds: float = 30 + warmup: float = 3 + consumer_delay: float = 0 + max_buffer_size: int = 1 + rss_every: float | None = None + + async def run(self) -> VideoResult: + """Runs the benchmark.""" + run = _VideoRun(self, VideoResult(self.width, self.height, self.fps, self.seconds)) + async with _connection() as (caller, callee): + generator = webrtc.VideoTrackGenerator() + # enough for the size, so the encoder doesn't drop frames for bitrate + bitrate = int(self.width * self.height * self.fps * 0.2) + remote = await _remote_track(caller, callee, generator.track, max_bitrate=bitrate) + processor = webrtc.MediaStreamTrackProcessor(remote, max_buffer_size=self.max_buffer_size) + sampling = asyncio.ensure_future(run.sample_rss(self.rss_every)) if self.rss_every else None + lag = await run.phases.measure( + run.result, processor, write=lambda: run.write(generator), read=lambda: run.read(processor) + ) + if sampling: + await sampling + run.result.lag_p95_ms, run.result.lag_max_ms = lag.p95_ms, lag.max_ms + generator.track.stop() + return run.result + + +@dataclass +class _VideoRun: + options: VideoLoopback + result: VideoResult + phases: _Phases = field(init=False) + # when each frame number was written + sent_at: dict[int, float] = field(default_factory=dict) + + def __post_init__(self) -> None: + self.phases = _Phases(self.options.warmup, self.options.seconds) + + async def write(self, generator: webrtc.VideoTrackGenerator) -> None: + width, height, fps = self.options.width, self.options.height, self.options.fps + frames = _frames(width, height) + writer = generator.writable.get_writer() + loop = asyncio.get_running_loop() + start = loop.time() + number = 0 + while not self.phases.done.is_set(): + frame = frames[number % len(frames)] + _draw_number(frame, width, number) + self.sent_at[number] = loop.time() + # frames lost on the way are forgotten + self.sent_at.pop(number - 300, None) + if self.phases.measuring.is_set(): + self.result.sent += 1 + await writer.write( + webrtc.VideoFrame(frame, format='I420', coded_width=width, coded_height=height, timestamp=number) + ) + number += 1 + await asyncio.sleep(max(0.0, start + number / fps - loop.time())) + + async def read(self, processor: webrtc.MediaStreamTrackProcessor) -> None: + header = {'rect': {'x': 0, 'y': 0, 'width': BITS * BLOCK, 'height': BLOCK}} + loop = asyncio.get_running_loop() + async for frame in processor.readable: + now = loop.time() + luma = bytearray(frame.allocation_size(header)) + await frame.copy_to(luma, header) + size = (frame.coded_width, frame.coded_height) + frame.close() + if self.phases.measuring.is_set(): + self._count(size, _read_number(luma, BITS * BLOCK), now) + if self.options.consumer_delay: + await asyncio.sleep(self.options.consumer_delay) + if self.phases.done.is_set(): + break + + def _count(self, size: tuple[int, int], number: int, received_at: float) -> None: + self.result.sizes[size] = self.result.sizes.get(size, 0) + 1 + if number in self.sent_at: + self.result.received += 1 + self.result.latencies.append(received_at - self.sent_at.pop(number)) + + async def sample_rss(self, every: float) -> None: + await self.phases.measuring.wait() + loop = asyncio.get_running_loop() + start = loop.time() + while not self.phases.done.is_set(): + self.result.rss.append((loop.time() - start, rss_bytes())) + await asyncio.sleep(every) @dataclass class AudioResult: + """The chunks written and received, and what it cost.""" + channels: int seconds: float written: int = 0 @@ -239,15 +294,15 @@ class AudioResult: @property def chunks_per_second(self) -> float: + """The chunks received per second while measuring.""" return self.received / self.seconds async def audio_loopback(channels: int = 2, seconds: float = 60, warmup: float = 2) -> AudioResult: - """10 ms chunks of a 48 kHz sine written to a generator in real time, read from the remote processor""" + """10 ms chunks of a 48 kHz sine written to a generator in real time, read from the remote processor.""" result = AudioResult(channels, seconds) loop = asyncio.get_running_loop() - measuring = asyncio.Event() - done = asyncio.Event() + phases = _Phases(warmup, seconds) chunk = array.array( 'h', (int(8000 * math.sin(2 * math.pi * 440 * (i // channels) / 48000)) for i in range(480 * channels)) ).tobytes() @@ -257,11 +312,11 @@ async def audio_loopback(channels: int = 2, seconds: float = 60, warmup: float = remote = await _remote_track(caller, callee, generator) processor = webrtc.MediaStreamTrackProcessor(remote, max_buffer_size=50) - async def write(): + async def write() -> None: writer = generator.writable.get_writer() start = loop.time() written = 0 - while not done.is_set(): + while not phases.done.is_set(): await writer.write( webrtc.AudioData( format='s16', @@ -273,20 +328,20 @@ async def write(): ) ) written += 1 - if measuring.is_set(): + if phases.measuring.is_set(): result.written += 1 await asyncio.sleep(max(0.0, start + written / 100 - loop.time())) - async def read(): + async def read() -> None: async for audio in processor.readable: - if measuring.is_set(): + if phases.measuring.is_set(): result.received += 1 result.received_frames += audio.number_of_frames audio.close() - if done.is_set(): + if phases.done.is_set(): break - lag = await _measure(result, processor, write, read, measuring, done, warmup, seconds) + lag = await phases.measure(result, processor, write=write, read=read) result.lag_p95_ms = lag.p95_ms generator.stop() return result @@ -294,6 +349,8 @@ async def read(): @dataclass class CopyResult: + """The cost of a copy of a frame to a format.""" + width: int height: int format: str @@ -301,11 +358,14 @@ class CopyResult: @property def megapixels_per_second(self) -> float: + """The throughput of the copy.""" return self.width * self.height / (self.milliseconds / 1000) / 1e6 -def copy_costs(sizes, formats=('I420', 'RGBA', 'BGRA'), budget: float = 1.0) -> List[CopyResult]: - """How long VideoFrame.copy_to takes: a copy of the planes, or a conversion to RGB""" +async def copy_costs( + sizes: Iterable[tuple[int, int]], formats: Sequence[str] = ('I420', 'RGBA', 'BGRA'), budget: float = 1.0 +) -> list[CopyResult]: + """How long VideoFrame.copy_to takes: a copy of the planes, or a conversion to RGB.""" results = [] for width, height in sizes: chroma = (width // 2) * (height // 2) @@ -318,8 +378,7 @@ def copy_costs(sizes, formats=('I420', 'RGBA', 'BGRA'), budget: float = 1.0) -> runs = 0 start = time.perf_counter() while time.perf_counter() - start < budget: - # the copy alone, without the future of copy_to - frame._copy_to(destination, options) + await frame.copy_to(destination, options) runs += 1 results.append(CopyResult(width, height, format, (time.perf_counter() - start) / runs * 1000)) frame.close() @@ -327,7 +386,7 @@ def copy_costs(sizes, formats=('I420', 'RGBA', 'BGRA'), budget: float = 1.0) -> def construct_cost(width: int, height: int, budget: float = 1.0) -> float: - """How long creating a VideoFrame from a buffer takes, in milliseconds (it copies the pixels)""" + """How long creating a VideoFrame from a buffer takes, in milliseconds (it copies the pixels).""" data = bytes(width * height * 3 // 2) runs = 0 start = time.perf_counter() diff --git a/cmake/libcxx/update.py b/cmake/libcxx/update.py old mode 100644 new mode 100755 index 0329005..749e786 --- a/cmake/libcxx/update.py +++ b/cmake/libcxx/update.py @@ -1,4 +1,12 @@ -"""Regenerates the libc++ pins used by Linux builds (maintainer tool, needs `git` and the `gh` CLI). +#!/usr/bin/env python3 +# +# Copyright 2026 Ilya (Marshal) . All rights reserved. +# +# Use of this source code is governed by a BSD-style license +# that can be found in the LICENSE.md file in the root of the project. +# + +r"""Regenerates the libc++ pins used by Linux builds (maintainer tool, needs `git` and the `gh` CLI). The Linux libwebrtc prebuilt ships only the *.h headers of Chromium's libc++. This script - resolves the libc++, libc++abi and llvm-libc revisions of WebRTC's DEPS to llvm-project commits, and checks @@ -14,18 +22,22 @@ --webrtc-branch 7977 --chromium-tag 152.0.7977.0 """ +from __future__ import annotations + import argparse import base64 import datetime +import functools import hashlib import json import posixpath import re +import shutil import subprocess import tempfile from concurrent.futures import ThreadPoolExecutor from pathlib import Path -from typing import Any, Dict, List, Set +from typing import TypedDict REPO = 'repos/llvm/llvm-project' CHROMIUM_CONFIG = ('__config_site', '__assertion_handler') @@ -40,65 +52,118 @@ INCLUDE = re.compile(rb'^\s*#\s*include\s*[<"]([^>"]+)[>"]', re.MULTILINE) -def gh(path: str) -> Any: - return json.loads(subprocess.check_output(['gh', 'api', path])) +class TreeEntry(TypedDict): + """An entry of a git tree of the GitHub API.""" + + path: str + sha: str + type: str + + +class Commit(TypedDict): + """A commit of the GitHub API.""" + + sha: str + + +class Tree(TypedDict): + """A git tree of the GitHub API.""" + + tree: list[TreeEntry] + truncated: bool + + +@functools.cache +def tool(name: str) -> str: + """The full path of an executable on PATH.""" + path = shutil.which(name) + if path is None: + msg = f'{name} is not on PATH' + raise SystemExit(msg) + return path + + +def gh(path: str) -> bytes: + """A response of the GitHub API.""" + return subprocess.check_output([tool('gh'), 'api', path]) + + +def git_tree(sha: str, *, recursive: bool = False) -> Tree: + """A git tree of llvm-project.""" + tree: Tree = json.loads(gh(f'{REPO}/git/trees/{sha}' + ('?recursive=1' if recursive else ''))) + return tree def blob_sha(data: bytes) -> str: - return hashlib.sha1(b'blob %d\0' % len(data) + data).hexdigest() + """The git object hash of a blob.""" + return hashlib.sha1(b'blob %d\0' % len(data) + data, usedforsecurity=False).hexdigest() def tree_sha(commit: str, path: str) -> str: - tree: str = gh(f'{REPO}/git/commits/{commit}')['tree']['sha'] + """The git tree hash of a directory of llvm-project at a commit.""" + tree: str = json.loads(gh(f'{REPO}/git/commits/{commit}'))['tree']['sha'] for part in path.split('/'): - tree = next(e['sha'] for e in gh(f'{REPO}/git/trees/{tree}')['tree'] if e['path'] == part) + tree = next(e['sha'] for e in git_tree(tree)['tree'] if e['path'] == part) return tree -def subtree(commit: str, path: str) -> List[Dict[str, Any]]: - listing = gh(f'{REPO}/git/trees/{tree_sha(commit, path)}?recursive=1') +def subtree(commit: str, path: str) -> list[TreeEntry]: + """The files of a directory of llvm-project at a commit, recursively.""" + listing = git_tree(tree_sha(commit, path), recursive=True) if listing['truncated']: - raise SystemExit(f'{path} tree listing is truncated') + msg = f'{path} tree listing is truncated' + raise SystemExit(msg) return [e for e in listing['tree'] if e['type'] == 'blob'] def blob(sha: str) -> bytes: - return base64.b64decode(gh(f'{REPO}/git/blobs/{sha}')['content']) + """The content of a blob of llvm-project.""" + content: str = json.loads(gh(f'{REPO}/git/blobs/{sha}'))['content'] + return base64.b64decode(content) def chromium(path: str, tag: str) -> bytes: - return subprocess.check_output( - ['gh', 'api', '-H', 'Accept: application/vnd.github.raw', f'repos/chromium/chromium/contents/{path}?ref={tag}'] - ) + """A file of Chromium at a tag.""" + return subprocess.check_output([ + tool('gh'), + 'api', + '-H', + 'Accept: application/vnd.github.raw', + f'repos/chromium/chromium/contents/{path}?ref={tag}', + ]) def git_fetch(url: str, ref: str, repo: str) -> str: """Fetches only the commit object of , returns its hash.""" - subprocess.run(['git', 'init', '-q', repo], check=True) - subprocess.run(['git', '-C', repo, 'fetch', '-q', '--depth=1', '--filter=tree:0', url, ref], check=True) - return subprocess.check_output(['git', '-C', repo, 'rev-parse', 'FETCH_HEAD'], text=True).strip() + git = tool('git') + subprocess.run([git, 'init', '-q', repo], check=True) + subprocess.run([git, '-C', repo, 'fetch', '-q', '--depth=1', '--filter=tree:0', url, ref], check=True) + return subprocess.check_output([git, '-C', repo, 'rev-parse', 'FETCH_HEAD'], text=True).strip() def resolve_llvm(deps: str, name: str) -> str: """Chromium mirrors each llvm-project dir as its own repo: finds the llvm-project commit with the same tree.""" match = re.search(rf"'{re.escape(LLVM_DIRS[name][0])}':\s*'([^'@]+)@([0-9a-f]+)'", deps) if not match: - raise SystemExit(f'{LLVM_DIRS[name][0]} is not in the WebRTC DEPS') + msg = f'{LLVM_DIRS[name][0]} is not in the WebRTC DEPS' + raise SystemExit(msg) url, revision = match.groups() with tempfile.TemporaryDirectory() as repo: git_fetch(url, revision, repo) tree, committed = subprocess.check_output( - ['git', '-C', repo, 'log', '-1', '--format=%T %cI', 'FETCH_HEAD'], text=True + [tool('git'), '-C', repo, 'log', '-1', '--format=%T %cI', 'FETCH_HEAD'], text=True ).split() date = datetime.datetime.fromisoformat(committed).astimezone(datetime.timezone.utc) since, until = ((date + datetime.timedelta(days=d)).strftime('%Y-%m-%dT%H:%M:%SZ') for d in (-1, 1)) - for commit in gh(f'{REPO}/commits?path={name}&since={since}&until={until}&per_page=100'): + commits: list[Commit] = json.loads(gh(f'{REPO}/commits?path={name}&since={since}&until={until}&per_page=100')) + for commit in commits: if tree_sha(commit['sha'], name) == tree: return str(commit['sha']) - raise SystemExit(f'no llvm-project commit has the {name} tree of {url}@{revision}') + msg = f'no llvm-project commit has the {name} tree of {url}@{revision}' + raise SystemExit(msg) -def runtime_sources(tag: str) -> List[str]: +def runtime_sources(tag: str) -> list[str]: """The libc++ and libc++abi sources of Chromium's Linux build, as llvm-project paths.""" sources = [] for gn, llvm in (('libc%2B%2B', 'libcxx'), ('libc%2B%2Babi', 'libcxxabi')): @@ -109,58 +174,61 @@ def runtime_sources(tag: str) -> List[str]: return sorted(set(sources)) -def pin_runtime(commits: Dict[str, str], tag: str) -> None: +def includes(path: str, data: bytes) -> set[str]: + """The llvm-project paths a file may include: next to it, or from libcxx/src and llvm-libc.""" + found = set() + for include in INCLUDE.findall(data): + name = include.decode() + found.add(posixpath.normpath(posixpath.join(posixpath.dirname(path), name))) + found.update(f'{root}/{name}' for root in ('libcxx/src', 'libc')) + return found + + +def pin_runtime(commits: dict[str, str], tag: str) -> None: """Pins the runtime sources and, transitively, every file they include from RUNTIME_TREES.""" - tree: Dict[str, str] = {} + tree: dict[str, str] = {} for root in RUNTIME_TREES: tree.update({f'{root}/{e["path"]}': e['sha'] for e in subtree(commits[root.split('/')[0]], root)}) sources = runtime_sources(tag) - pinned: Dict[str, bytes] = {} + pinned: dict[str, bytes] = {} queue = list(sources) with ThreadPoolExecutor(8) as pool: while queue: - for path, data in zip(queue, pool.map(blob, [tree[p] for p in queue])): - pinned[path] = data - found: Set[str] = set() - for path in queue: - for include in INCLUDE.findall(pinned[path]): - name = include.decode() - candidates = [posixpath.normpath(posixpath.join(posixpath.dirname(path), name))] - candidates += [f'{root}/{name}' for root in ('libcxx/src', 'libc')] - found.update(c for c in candidates if c in tree and c not in pinned) - queue = sorted(found) + pinned.update(zip(queue, pool.map(blob, [tree[p] for p in queue]))) + found = set().union(*(includes(path, pinned[path]) for path in queue)) + queue = sorted(c for c in found if c in tree and c not in pinned) # CMake compiles every pinned libcxx/libcxxabi .cpp, so only the sources may be among them stray = [p for p in pinned if p.endswith('.cpp') and p not in sources and not p.startswith('libc/')] if stray: - raise SystemExit(f'sources include other .cpp files: {stray}') + msg = f'sources include other .cpp files: {stray}' + raise SystemExit(msg) lines = [f'{hashlib.sha256(data).hexdigest()} {path}\n' for path, data in sorted(pinned.items())] (Path(__file__).parent / 'runtime.sha256').write_text(''.join(lines)) print(f'{len(sources)} runtime sources, {len(lines)} files pinned') -def main() -> None: - parser = argparse.ArgumentParser() - parser.add_argument('headers', type=Path, help='libc++ include dir of the unpacked Linux prebuilt') - parser.add_argument('--webrtc-branch', required=True, help='WebRTC branch number of the prebuilt, e.g. 7977') - parser.add_argument('--chromium-tag', required=True, help='Chromium tag matching the WebRTC branch') - args = parser.parse_args() - +def resolve_commits(webrtc_branch: str) -> dict[str, str]: + """The llvm-project commits of the LLVM_DIRS the WebRTC branch pins.""" with tempfile.TemporaryDirectory() as repo: - webrtc = git_fetch('https://webrtc.googlesource.com/src', f'refs/branch-heads/{args.webrtc_branch}', repo) - deps = subprocess.check_output(['git', '-C', repo, 'show', 'FETCH_HEAD:DEPS'], text=True) - print(f'WebRTC branch-heads/{args.webrtc_branch} at {webrtc}') - commits = {name: resolve_llvm(deps, name) for name in LLVM_DIRS} + webrtc = git_fetch('https://webrtc.googlesource.com/src', f'refs/branch-heads/{webrtc_branch}', repo) + deps = subprocess.check_output([tool('git'), '-C', repo, 'show', 'FETCH_HEAD:DEPS'], text=True) + print(f'WebRTC branch-heads/{webrtc_branch} at {webrtc}') + return {name: resolve_llvm(deps, name) for name in LLVM_DIRS} - tree = subtree(commits['libcxx'], 'libcxx/include') + +def pin_headers(commit: str, headers: Path) -> None: + """Checks the prebuilt *.h headers against libc++ and pins the extensionless ones it lacks.""" + tree = subtree(commit, 'libcxx/include') remote = {e['path']: e['sha'] for e in tree} - local = {str(p.relative_to(args.headers)): blob_sha(p.read_bytes()) for p in args.headers.rglob('*') if p.is_file()} + local = {str(p.relative_to(headers)): blob_sha(p.read_bytes()) for p in headers.rglob('*') if p.is_file()} mismatches = [path for path, sha in local.items() if remote.get(path) != sha] if mismatches: - raise SystemExit(f'the prebuilt headers differ from llvm-project@{commits["libcxx"]}: {mismatches[:10]}') + msg = f'the prebuilt headers differ from llvm-project@{commit}: {mismatches[:10]}' + raise SystemExit(msg) missing = [e for e in tree if e['path'] not in local and '.' not in e['path'].rsplit('/', 1)[-1]] - def sha256(entry: Dict[str, Any]) -> str: + def sha256(entry: TreeEntry) -> str: data = blob(entry['sha']) return f'{hashlib.sha256(data).hexdigest()} {entry["path"]}\n' @@ -169,13 +237,28 @@ def sha256(entry: Dict[str, Any]) -> str: (Path(__file__).parent / 'headers.sha256').write_text(''.join(lines)) print(f'{len(lines)} headers pinned') + +def pin_config(tag: str) -> None: + """Pins Chromium's build-generated libc++ config.""" config = [] for name in CHROMIUM_CONFIG: - data = chromium(f'buildtools/third_party/libc%2B%2B/{name}', args.chromium_tag) + data = chromium(f'buildtools/third_party/libc%2B%2B/{name}', tag) config.append(f'{hashlib.sha256(data).hexdigest()} {name}\n') (Path(__file__).parent / 'config.sha256').write_text(''.join(config)) print(f'{len(config)} config headers pinned') + +def main() -> None: + """Regenerates the pins and prints the CMake variables.""" + parser = argparse.ArgumentParser() + parser.add_argument('headers', type=Path, help='libc++ include dir of the unpacked Linux prebuilt') + parser.add_argument('--webrtc-branch', required=True, help='WebRTC branch number of the prebuilt, e.g. 7977') + parser.add_argument('--chromium-tag', required=True, help='Chromium tag matching the WebRTC branch') + args = parser.parse_args() + + commits = resolve_commits(args.webrtc_branch) + pin_headers(commits['libcxx'], args.headers) + pin_config(args.chromium_tag) pin_runtime(commits, args.chromium_tag) for name, (_, variable) in LLVM_DIRS.items(): diff --git a/examples/echo.py b/examples/echo.py old mode 100644 new mode 100755 index dba0fa9..d00b201 --- a/examples/echo.py +++ b/examples/echo.py @@ -1,10 +1,20 @@ -"""An echo peer: it sends back the video it receives, in grayscale, with a processor piped through a transform -stream into a generator, as in a browser. +#!/usr/bin/env python3 +# +# Copyright 2026 Ilya (Marshal) . +# +# Dedicated to the public domain under CC0, see the LICENSE file of the examples. +# + +"""An echo peer: it sends back the video it receives, in grayscale. + +The video goes through a processor piped through a transform stream into a generator, as in a browser. Two connections in this process stand for the two peers: the caller sends its synthetic camera, the echo peer sends the frames back, and the caller checks that they have no color. """ +from __future__ import annotations + import asyncio import webrtc @@ -12,28 +22,26 @@ SECONDS = 3 -class Grayscale: - """A transformer of I420 frames: U and V at 128 leave only the luma""" - - async def transform(self, frame, controller): - data = bytearray(frame.allocation_size({'format': 'I420'})) - await frame.copy_to(data, {'format': 'I420'}) - luma = frame.coded_width * frame.coded_height - data[luma:] = b'\x80' * (len(data) - luma) - controller.enqueue( - webrtc.VideoFrame( - data, - format='I420', - coded_width=frame.coded_width, - coded_height=frame.coded_height, - timestamp=frame.timestamp, - ) +async def grayscale(frame: webrtc.VideoFrame, controller: webrtc.TransformStreamDefaultController) -> None: + """Transforms an I420 frame: U and V at 128 leave only the luma.""" + data = bytearray(frame.allocation_size({'format': 'I420'})) + await frame.copy_to(data, {'format': 'I420'}) + luma = frame.coded_width * frame.coded_height + data[luma:] = b'\x80' * (len(data) - luma) + controller.enqueue( + webrtc.VideoFrame( + data, + format='I420', + coded_width=frame.coded_width, + coded_height=frame.coded_height, + timestamp=frame.timestamp, ) - frame.close() + ) + frame.close() -async def watch(track): - """Reads the echoed frames for a while, then prints whether the last one is gray""" +async def watch(track: webrtc.MediaStreamTrack) -> None: + """Reads the echoed frames for a while, then prints whether the last one is gray.""" reader = webrtc.MediaStreamTrackProcessor(track).readable.get_reader() loop = asyncio.get_running_loop() end = loop.time() + SECONDS @@ -48,19 +56,21 @@ async def watch(track): print(f'{frames} frames echoed, the last one is gray: {red == green == blue} ({red}, {green}, {blue})') -def trickle(caller, callee): - """Passes the ICE candidates of each connection to the other one""" +def trickle(caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection) -> None: + """Passes the ICE candidates of each connection to the other one.""" for pc, other in ((caller, callee), (callee, caller)): - async def on_candidate(event, other=other): + async def on_candidate( + event: webrtc.RTCPeerConnectionIceEvent, other: webrtc.RTCPeerConnection = other + ) -> None: if event.candidate: await other.add_ice_candidate(event.candidate) pc.on('icecandidate', on_candidate) -async def negotiate(caller, callee): - """Exchanges an offer and an answer""" +async def negotiate(caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection) -> None: + """Exchanges an offer and an answer.""" offer = await caller.create_offer() await caller.set_local_description(offer) await callee.set_remote_description(offer) @@ -69,23 +79,24 @@ async def negotiate(caller, callee): await caller.set_remote_description(answer) -async def main(): +async def main() -> None: + """Echoes the camera of the caller back to it.""" caller, echo = webrtc.RTCPeerConnection(), webrtc.RTCPeerConnection() trickle(caller, echo) camera = webrtc.get_user_media(audio=False, video=True).get_video_tracks()[0] caller.add_track(camera) generator = webrtc.VideoTrackGenerator() echoed = asyncio.get_running_loop().create_future() - pipes = [] + pipes: list[asyncio.Future[None]] = [] @echo.on('track') - def on_echo_track(event): + def on_echo_track(event: webrtc.RTCTrackEvent) -> None: readable = webrtc.MediaStreamTrackProcessor(event.track).readable - pipe = readable.pipe_through(webrtc.TransformStream(Grayscale())).pipe_to(generator.writable) + pipe = readable.pipe_through(webrtc.TransformStream({'transform': grayscale})).pipe_to(generator.writable) pipes.append(asyncio.ensure_future(pipe)) @caller.on('track') - def on_caller_track(event): + def on_caller_track(event: webrtc.RTCTrackEvent) -> None: echoed.set_result(event.track) await negotiate(caller, echo) diff --git a/examples/janus_streaming.py b/examples/janus_streaming.py old mode 100644 new mode 100755 index c3012d4..8fba539 --- a/examples/janus_streaming.py +++ b/examples/janus_streaming.py @@ -3,6 +3,12 @@ # requires-python = ">=3.9" # dependencies = ["wrtc>=0.0.0.dev10", "sounddevice", "httpx"] # /// +# +# Copyright 2026 Ilya (Marshal) . +# +# Dedicated to the public domain under CC0, see the LICENSE file of the examples. +# + """Watches a stream of the public Janus demo server: the video in the terminal, the audio on your speakers. uv run janus_streaming.py @@ -11,6 +17,8 @@ The video is drawn with colored half blocks, so the terminal needs 24-bit color. Press Ctrl+C to stop. """ +from __future__ import annotations + import argparse import asyncio import contextlib @@ -18,73 +26,85 @@ import sys import threading import uuid +from typing import TYPE_CHECKING, Any import httpx import sounddevice import webrtc +if TYPE_CHECKING: + from collections.abc import Iterator + from types import TracebackType + + from typing_extensions import Self + JANUS = 'https://janus.conf.meetecho.com/janus' +Json = dict[str, Any] class Janus: - """A session with the streaming plugin of a Janus server, over its HTTP API""" + """A session with the streaming plugin of a Janus server, over its HTTP API.""" - def __init__(self): + def __init__(self) -> None: self.client = httpx.AsyncClient(timeout=60) - async def __aenter__(self): + async def __aenter__(self) -> Self: created = await self._post('', janus='create') self.session = f'/{created["data"]["id"]}' attached = await self._post(self.session, janus='attach', plugin='janus.plugin.streaming') self.handle = f'{self.session}/{attached["data"]["id"]}' return self - async def __aexit__(self, *exc_info): + async def __aexit__( + self, exc_type: type[BaseException] | None, exc: BaseException | None, traceback: TracebackType | None + ) -> None: await self._post(self.session, janus='destroy') await self.client.aclose() - async def request(self, body, **extra): + async def request(self, body: Json, **extra: Json) -> Json | None: + """Sends a message to the plugin, returns the data of its reply.""" reply = await self._post(self.handle, janus='message', body=body, **extra) return reply.get('plugindata', {}).get('data') - async def event(self): - """Waits up to 30 s for an event: requests reply right away and send their results, like offers, as events""" + async def event(self) -> Json: + """Waits up to 30 s for an event: requests reply right away and send their results, like offers, as events.""" return (await self.client.get(JANUS + self.session)).json() - async def _post(self, path, **message): - reply = (await self.client.post(JANUS + path, json={'transaction': uuid.uuid4().hex, **message})).json() + async def _post(self, path: str, **message: object) -> Json: + reply: Json = (await self.client.post(JANUS + path, json={'transaction': uuid.uuid4().hex, **message})).json() if reply['janus'] == 'error': raise RuntimeError(reply['error']['reason']) return reply class Speakers: - """Plays 16-bit audio from a buffer that keeps half a second at most""" + """Plays 16-bit audio from a buffer that keeps half a second at most.""" - def __init__(self, rate, channels): + def __init__(self, rate: int, channels: int) -> None: self.buffer, self.lock, self.limit = bytearray(), threading.Lock(), rate * channels self.stream = sounddevice.RawOutputStream(rate, channels=channels, dtype='int16', callback=self._on_need) self.stream.start() - def play(self, samples): + def play(self, samples: bytes) -> None: + """Queues samples, dropping the oldest ones over the limit.""" with self.lock: self.buffer += samples del self.buffer[: -self.limit] - def _on_need(self, out, *_): + def _on_need(self, out: memoryview, *_: object) -> None: with self.lock: chunk = self.buffer[: len(out)] del self.buffer[: len(out)] out[:] = chunk.ljust(len(out), b'\0') -def draw(rgbx, width, height): - """Draws a frame with ▀, whose foreground color is the upper pixel and background the lower one""" +def draw(rgbx: bytes, width: int, height: int) -> None: + """Draws a frame with ▀, whose foreground color is the upper pixel and background the lower one.""" columns, rows = shutil.get_terminal_size() scale = max(width / columns, height / rows / 2) - def color(x, y): + def color(x: int, y: int) -> str: i = (int(y * scale) * width + int(x * scale)) * 4 return '{};{};{}'.format(*rgbx[i : i + 3]) @@ -97,8 +117,8 @@ def color(x, y): @contextlib.contextmanager -def fullscreen(): - """Switches to the alternate screen, without the cursor""" +def fullscreen() -> Iterator[None]: + """Switches to the alternate screen, without the cursor.""" sys.stdout.write('\033[?1049h\033[?25l') try: yield @@ -106,7 +126,8 @@ def fullscreen(): sys.stdout.write('\033[?25h\033[?1049l') -async def watch(track): +async def watch(track: webrtc.MediaStreamTrack) -> None: + """Draws the frames of the video track until it ends.""" # a buffer of one frame drops the frames the terminal is too slow for async for frame in webrtc.MediaStreamTrackProcessor(track, max_buffer_size=1).readable: with frame: @@ -116,7 +137,8 @@ async def watch(track): draw(rgbx, int(size.width), int(size.height)) -async def listen(track): +async def listen(track: webrtc.MediaStreamTrack) -> None: + """Plays the audio track until it ends.""" speakers = None try: async for data in webrtc.MediaStreamTrackProcessor(track, max_buffer_size=50).readable: @@ -131,8 +153,8 @@ async def listen(track): speakers.stream.close() -async def answer(pc, offer): - """Answers with all the ICE candidates in the SDP, since there is no trickling""" +async def answer(pc: webrtc.RTCPeerConnection, offer: dict[str, str]) -> dict[str, str]: + """Answers with all the ICE candidates in the SDP, since there is no trickling.""" gathered = asyncio.Event() pc.on('icegatheringstatechange', lambda _: pc.ice_gathering_state == 'complete' and gathered.set()) await pc.set_remote_description(offer) @@ -142,31 +164,42 @@ async def answer(pc, offer): return {'type': 'answer', 'sdp': pc.local_description.sdp} -async def main(stream_id): - pc, tasks = webrtc.RTCPeerConnection(), [] +async def keep_alive(janus: Janus) -> None: + """Polls for events, which keeps the session alive.""" + while True: + await janus.event() + + +async def pick_stream(janus: Janus) -> int: + """Lists the streams of the server, returns the first one.""" + streams = (await janus.request({'request': 'list'}))['list'] + for stream in streams: + print(f'{stream["id"]}: {stream.get("description")}') + return streams[0]['id'] + + +async def main(stream_id: int | None) -> None: + """Watches the stream until interrupted.""" + pc = webrtc.RTCPeerConnection() + tasks: list[asyncio.Future[None]] = [] @pc.on('track') - def on_track(event): + def on_track(event: webrtc.RTCTrackEvent) -> None: play = watch if event.track.kind == 'video' else listen tasks.append(asyncio.ensure_future(play(event.track))) async with Janus() as janus: if stream_id is None: - streams = (await janus.request({'request': 'list'}))['list'] - for stream in streams: - print(f'{stream["id"]}: {stream.get("description")}') - stream_id = streams[0]['id'] - + stream_id = await pick_stream(janus) await janus.request({'request': 'watch', 'id': stream_id}) - event = {} + event: Json = {} while 'jsep' not in event: event = await janus.event() await janus.request({'request': 'start'}, jsep=await answer(pc, event['jsep'])) with fullscreen(): try: - while True: # polling for events keeps the session alive - await janus.event() + await keep_alive(janus) finally: pc.close() # ends the tracks, so watch and listen return await asyncio.gather(*tasks) diff --git a/examples/openai_live.py b/examples/openai_live.py index b575987..b408677 100755 --- a/examples/openai_live.py +++ b/examples/openai_live.py @@ -3,6 +3,12 @@ # requires-python = ">=3.9" # dependencies = ["wrtc>=0.0.0.dev10", "sounddevice", "httpx"] # /// +# +# Copyright 2026 Ilya (Marshal) . +# +# Dedicated to the public domain under CC0, see the LICENSE file of the examples. +# + """Talk to OpenAI GPT-Live from the terminal: your microphone goes to the model, its voice to your speakers. No server needed: the script posts its SDP offer to the OpenAI API with your key, as the guide @@ -23,8 +29,11 @@ hear itself; with headphones pass --barge-in to be able to interrupt it. """ +from __future__ import annotations + import argparse import asyncio +import contextlib import json import os import signal @@ -32,6 +41,7 @@ import threading import time from array import array +from typing import Any, ClassVar import httpx import sounddevice @@ -44,49 +54,65 @@ VOICE_LEVEL = 0.02 # the peak level above which audio counts as speech ECHO_TAIL = 0.6 # seconds the microphone stays muted after the assistant stops MAX_PLAYBACK = 0.5 # seconds of assistant audio buffered before the oldest is dropped +SILENT_MIC_WARNING = 10 # seconds of silence before the microphone is suspected class Console: - """Human-readable output: timestamped status lines and live transcripts that share the terminal""" - - COLORS = {'dim': '2', 'red': '31', 'green': '32', 'yellow': '33', 'blue': '34', 'magenta': '35', 'cyan': '36'} + """Human-readable output: timestamped status lines and live transcripts that share the terminal.""" + + COLORS: ClassVar[dict[str, str]] = { + 'dim': '2', + 'red': '31', + 'green': '32', + 'yellow': '33', + 'blue': '34', + 'magenta': '35', + 'cyan': '36', + } - def __init__(self, verbose): + def __init__(self, *, verbose: bool) -> None: self.verbose = verbose self.color = sys.stdout.isatty() and not os.environ.get('NO_COLOR') - self.speaker = None # who the open transcript line belongs to + self.speaker: str | None = None # who the open transcript line belongs to - def paint(self, text, color): + def paint(self, text: str, color: str) -> str: + """The text in a color, if the terminal shows colors.""" return f'\033[{self.COLORS[color]}m{text}\033[0m' if self.color else text - def _end_transcript(self): + def _end_transcript(self) -> None: if self.speaker: print(flush=True) self.speaker = None - def log(self, icon, text, color=None): + def log(self, icon: str, text: str, color: str | None = None) -> None: + """Prints a timestamped status line.""" self._end_transcript() line = f'{self.paint(time.strftime("%H:%M:%S"), "dim")} {icon} {text}' print(self.paint(line, color) if color else line, flush=True) - def info(self, text): + def info(self, text: str) -> None: + """Prints a status line.""" self.log('•', text) - def ok(self, text): + def ok(self, text: str) -> None: + """Prints a line of a success.""" self.log('✓', text, 'green') - def warn(self, text): + def warn(self, text: str) -> None: + """Prints a line of a problem.""" self.log('!', text, 'yellow') - def error(self, text): + def error(self, text: str) -> None: + """Prints a line of a failure.""" self.log('✗', text, 'red') - def debug(self, text): + def debug(self, text: str) -> None: + """Prints a line in verbose mode only.""" if self.verbose: self.log('·', text, 'dim') - def transcript(self, speaker, delta): - """Appends a fragment to the speaker's line; fragments have no end marker, so a new speaker starts a line""" + def transcript(self, speaker: str, delta: str) -> None: + """Appends a fragment to the speaker's line; fragments have no end marker, so a new speaker starts a line.""" if speaker != self.speaker: self._end_transcript() label, color = ('You', 'cyan') if speaker == 'user' else ('Assistant', 'magenta') @@ -95,51 +121,54 @@ def transcript(self, speaker, delta): print(delta, end='', flush=True) -def peak(samples): - """The peak level of 16-bit samples, from 0 to 1""" +def peak(samples: bytes) -> float: + """The peak level of 16-bit samples, from 0 to 1.""" values = array('h', samples) - return max(max(values), -min(values)) / 32768 if values else 0 + return max(*values, -min(values)) / 32768 if values else 0 class Microphone: - """Captures 10 ms chunks of 16-bit mono audio from a device into an asyncio queue""" + """Captures 10 ms chunks of 16-bit mono audio from a device into an asyncio queue.""" - def __init__(self, device, loop): - self.queue = asyncio.Queue(maxsize=50) + def __init__(self, device: str | int | None, loop: asyncio.AbstractEventLoop) -> None: + self.queue: asyncio.Queue[bytes] = asyncio.Queue(maxsize=50) self._loop = loop self._stream = sounddevice.RawInputStream( samplerate=SAMPLE_RATE, channels=1, dtype='int16', blocksize=FRAME, device=device, callback=self._on_audio ) self.name = sounddevice.query_devices(self._stream.device)['name'] - def _on_audio(self, data, frames, time_info, status): + def _on_audio(self, data: memoryview, *_: object) -> None: self._loop.call_soon_threadsafe(self._put, bytes(data)) - def _put(self, chunk): + def _put(self, chunk: bytes) -> None: if self.queue.full(): self.queue.get_nowait() self.queue.put_nowait(chunk) - def start(self): + def start(self) -> None: + """Starts capturing.""" self._stream.start() - def close(self): + def close(self) -> None: + """Stops capturing and releases the device.""" self._stream.stop() self._stream.close() class Speakers: - """Plays 16-bit audio on a device from a small buffer, opened on the first chunk to match its format""" + """Plays 16-bit audio on a device from a small buffer, opened on the first chunk to match its format.""" - def __init__(self, device): + def __init__(self, device: str | int | None) -> None: self.device = device self.name = sounddevice.query_devices(device, 'output')['name'] - self._stream = None + self._stream: sounddevice.RawOutputStream | None = None self._buffer = bytearray() self._lock = threading.Lock() self._limit = 0 - def play(self, samples, sample_rate, channels): + def play(self, samples: bytes, sample_rate: int, channels: int) -> None: + """Queues samples, opening the device on the first ones.""" if self._stream is None: self._limit = int(sample_rate * MAX_PLAYBACK) * channels * 2 self._stream = sounddevice.RawOutputStream( @@ -151,7 +180,7 @@ def play(self, samples, sample_rate, channels): if len(self._buffer) > self._limit: del self._buffer[: len(self._buffer) - self._limit] - def _on_need(self, out, frames, time_info, status): + def _on_need(self, out: memoryview, *_: object) -> None: size = len(out) with self._lock: chunk = self._buffer[:size] @@ -159,25 +188,29 @@ def _on_need(self, out, frames, time_info, status): out[: len(chunk)] = chunk out[len(chunk) :] = bytes(size - len(chunk)) - def close(self): + def close(self) -> None: + """Releases the device.""" if self._stream is not None: self._stream.stop() self._stream.close() class LiveCall: - def __init__(self, args, console): + """A call with the model: the microphone and the speakers, the connection and its event channel.""" + + def __init__(self, args: argparse.Namespace, console: Console) -> None: self.args = args self.console = console - self.pc = None - self.events = None - self.microphone = None - self.speakers = None - self.tasks = [] + self.pc: webrtc.RTCPeerConnection | None = None + self.events: webrtc.RTCDataChannel | None = None + self.microphone: Microphone | None = None + self.speakers: Speakers | None = None + self.tasks: list[asyncio.Future[None]] = [] self.session_closed = asyncio.Event() self.assistant_spoke_at = 0.0 # when the assistant's audio was last above the speech level - async def run(self, hang_up): + async def run(self, hang_up: asyncio.Event) -> None: + """Connects, then talks until hung up.""" console, args = self.console, self.args loop = asyncio.get_running_loop() @@ -196,7 +229,7 @@ async def run(self, hang_up): self.pc.add_track(generator) # created before the offer, so the offer negotiates it self.events = self.pc.create_data_channel('oai-events') - self.events.on('open', lambda event: console.ok('Event channel open')) + self.events.on('open', lambda _event: console.ok('Event channel open')) self.events.on('message', self._on_message) offer = await self.pc.create_offer() @@ -212,18 +245,18 @@ async def run(self, hang_up): await hang_up.wait() await self.close() - async def _gathered(self): - """Waits for the local ICE candidates, which go in the offer since there is no trickling""" + async def _gathered(self) -> None: + """Waits for the local ICE candidates, which go in the offer since there is no trickling.""" done = asyncio.Event() - self.pc.on('icegatheringstatechange', lambda event: self.pc.ice_gathering_state == 'complete' and done.set()) + self.pc.on('icegatheringstatechange', lambda _event: self.pc.ice_gathering_state == 'complete' and done.set()) if self.pc.ice_gathering_state != 'complete': try: await asyncio.wait_for(done.wait(), 5) except asyncio.TimeoutError: self.console.warn('ICE gathering is slow, sending the candidates found so far') - async def _create_session(self, sdp): - session = {'model': self.args.model} + async def _create_session(self, sdp: str) -> str: + session: dict[str, object] = {'model': self.args.model} if self.args.instructions: session['instructions'] = self.args.instructions if self.args.voice: @@ -236,14 +269,15 @@ async def _create_session(self, sdp): json={'session': session, 'transport': {'type': 'webrtc', 'sdp': sdp}}, ) except httpx.HTTPError as e: - raise CallError(f'Could not reach the OpenAI API: {e or type(e).__name__}') from None + msg = f'Could not reach the OpenAI API: {e or type(e).__name__}' + raise CallError(msg) from None if response.is_error: raise CallError(_describe_http_error(response)) - reply = response.json() + reply: dict[str, Any] = response.json() self.console.ok(f'Session created: {reply.get("session", {}).get("id", "?")}') return reply['transport']['sdp'] - def _on_connection_state(self, event): + def _on_connection_state(self, _event: webrtc.Event) -> None: state = self.pc.connection_state messages = { 'connecting': ('info', 'Connecting audio...'), @@ -255,12 +289,12 @@ def _on_connection_state(self, event): level, text = messages[state] getattr(self.console, level)(text) - def _on_track(self, event): - self.console.debug(f'Receiving the assistant\'s {event.track.kind} track') + def _on_track(self, event: webrtc.RTCTrackEvent) -> None: + self.console.debug(f"Receiving the assistant's {event.track.kind} track") self.tasks.append(asyncio.ensure_future(self._play(event.track))) - async def _send_microphone(self, writer): - """Sends the microphone to the model, silence instead while the echo guard holds it""" + async def _send_microphone(self, writer: webrtc.WritableStreamDefaultWriter) -> None: + """Sends the microphone to the model, silence instead while the echo guard holds it.""" console, loop = self.console, asyncio.get_running_loop() silence = bytes(FRAME * 2) timestamp, heard, started = 0, False, loop.time() @@ -270,9 +304,12 @@ async def _send_microphone(self, writer): if peak(chunk) > VOICE_LEVEL: heard = True console.ok('Microphone is picking up sound') - elif loop.time() - started > 10: + elif loop.time() - started > SILENT_MIC_WARNING: heard = True - console.warn('The microphone has been silent for 10 s: check its permission and input level') + console.warn( + f'The microphone has been silent for {SILENT_MIC_WARNING} s: ' + 'check its permission and input level' + ) guarded = not self.args.barge_in and time.monotonic() - self.assistant_spoke_at < ECHO_TAIL data = webrtc.AudioData( format='s16', @@ -285,8 +322,8 @@ async def _send_microphone(self, writer): await writer.write(data) timestamp += 10_000 - async def _play(self, track): - """Plays the assistant's audio""" + async def _play(self, track: webrtc.MediaStreamTrack) -> None: + """Plays the assistant's audio.""" heard = False async for data in webrtc.MediaStreamTrackProcessor(track, max_buffer_size=50).readable: with data: @@ -297,13 +334,13 @@ async def _play(self, track): self.assistant_spoke_at = time.monotonic() if not heard: heard = True - self.console.ok(f'Receiving the assistant\'s voice ({rate} Hz, {channels} ch)') + self.console.ok(f"Receiving the assistant's voice ({rate} Hz, {channels} ch)") self.speakers.play(samples, rate, channels) - def _on_message(self, event): + def _on_message(self, event: webrtc.MessageEvent) -> None: console = self.console try: - message = json.loads(event.data) + message: dict[str, Any] = json.loads(event.data) except ValueError: console.debug(f'Not JSON: {event.data!r}') return @@ -330,8 +367,8 @@ def _on_message(self, event): else: console.debug(f'Event {kind}') - async def close(self): - """Ends the session gracefully, then releases the audio devices""" + async def close(self) -> None: + """Ends the session gracefully, then releases the audio devices.""" console = self.console if self.events is not None and self.events.ready_state == 'open': console.info('Hanging up...') @@ -339,7 +376,7 @@ async def close(self): try: await asyncio.wait_for(self.session_closed.wait(), 5) except asyncio.TimeoutError: - console.warn('The session didn\'t confirm closing') + console.warn("The session didn't confirm closing") for task in self.tasks: task.cancel() await asyncio.gather(*self.tasks, return_exceptions=True) @@ -353,10 +390,10 @@ async def close(self): class CallError(Exception): - pass + """The call couldn't be set up.""" -def _describe_http_error(response): +def _describe_http_error(response: httpx.Response) -> str: try: message = response.json().get('error', {}).get('message') except ValueError: @@ -371,7 +408,8 @@ def _describe_http_error(response): return f'OpenAI API returned {response.status_code}, {summary}' + (f': {message}' if message else '') -def parse_args(): +def parse_args() -> argparse.Namespace: + """The options of the command line.""" parser = argparse.ArgumentParser( description='Voice chat with OpenAI GPT-Live over WebRTC.', formatter_class=argparse.ArgumentDefaultsHelpFormatter, @@ -393,21 +431,21 @@ def parse_args(): return args -async def main(): +async def main() -> None: + """Runs a call until Ctrl+C.""" args = parse_args() if args.list_devices: print(sounddevice.query_devices()) return - console = Console(args.verbose) + console = Console(verbose=args.verbose) if not args.api_key: console.error('No API key: set OPENAI_API_KEY or pass --api-key') sys.exit(1) hang_up = asyncio.Event() - try: + # on Windows, Ctrl+C raises KeyboardInterrupt instead + with contextlib.suppress(NotImplementedError): asyncio.get_running_loop().add_signal_handler(signal.SIGINT, hang_up.set) - except NotImplementedError: - pass # Windows: Ctrl+C raises KeyboardInterrupt instead call = LiveCall(args, console) try: @@ -419,7 +457,5 @@ async def main(): if __name__ == '__main__': - try: + with contextlib.suppress(KeyboardInterrupt): asyncio.run(main()) - except KeyboardInterrupt: - pass diff --git a/examples/recorder.py b/examples/recorder.py old mode 100644 new mode 100755 index e7b13ec..d3440a5 --- a/examples/recorder.py +++ b/examples/recorder.py @@ -1,3 +1,10 @@ +#!/usr/bin/env python3 +# +# Copyright 2026 Ilya (Marshal) . +# +# Dedicated to the public domain under CC0, see the LICENSE file of the examples. +# + """Records the media a peer receives to raw files, with MediaStreamTrackProcessor. Two connections in this process stand for the two peers: one sends the synthetic camera and microphone of @@ -8,16 +15,21 @@ ffplay -f s16le -ar 48000 -ch_layout mono audio.pcm """ +from __future__ import annotations + import asyncio +import pathlib +from typing import BinaryIO import webrtc SECONDS = 5 -async def record(track, path): +async def record(track: webrtc.MediaStreamTrack, file: BinaryIO) -> None: + """Writes the frames of a track to a file until the track ends.""" frames = 0 - with open(path, 'wb') as file: + with file: async for media in webrtc.MediaStreamTrackProcessor(track, max_buffer_size=30).readable: if track.kind == 'audio': data = bytearray(media.allocation_size({'plane_index': 0})) @@ -28,22 +40,24 @@ async def record(track, path): media.close() file.write(data) frames += 1 - print(f'{path}: {frames} {track.kind} frames') + print(f'{file.name}: {frames} {track.kind} frames') -def trickle(caller, callee): - """Passes the ICE candidates of each connection to the other one""" +def trickle(caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection) -> None: + """Passes the ICE candidates of each connection to the other one.""" for pc, other in ((caller, callee), (callee, caller)): - async def on_candidate(event, other=other): + async def on_candidate( + event: webrtc.RTCPeerConnectionIceEvent, other: webrtc.RTCPeerConnection = other + ) -> None: if event.candidate: await other.add_ice_candidate(event.candidate) pc.on('icecandidate', on_candidate) -async def negotiate(caller, callee): - """Exchanges an offer and an answer""" +async def negotiate(caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection) -> None: + """Exchanges an offer and an answer.""" offer = await caller.create_offer() await caller.set_local_description(offer) await callee.set_remote_description(offer) @@ -52,19 +66,20 @@ async def negotiate(caller, callee): await caller.set_remote_description(answer) -async def main(): +async def main() -> None: + """Records the camera and the microphone for a few seconds.""" sender, receiver = webrtc.RTCPeerConnection(), webrtc.RTCPeerConnection() trickle(sender, receiver) stream = webrtc.get_user_media(audio=True, video=True) for track in stream.get_tracks(): sender.add_track(track, stream) - recordings = [] + recordings: list[asyncio.Future[None]] = [] @receiver.on('track') - def on_track(event): - path = 'audio.pcm' if event.track.kind == 'audio' else 'video.i420' - recordings.append(asyncio.ensure_future(record(event.track, path))) + def on_track(event: webrtc.RTCTrackEvent) -> None: + path = pathlib.Path('audio.pcm' if event.track.kind == 'audio' else 'video.i420') + recordings.append(asyncio.ensure_future(record(event.track, path.open('wb')))) await negotiate(sender, receiver) await asyncio.sleep(SECONDS) diff --git a/examples/telegram_group_calls.py b/examples/telegram_group_calls.py old mode 100644 new mode 100755 index 7e5856a..5374663 --- a/examples/telegram_group_calls.py +++ b/examples/telegram_group_calls.py @@ -1,61 +1,102 @@ +#!/usr/bin/env python3 +# +# Copyright 2022 Ilya (Marshal) . +# +# Dedicated to the public domain under CC0, see the LICENSE file of the examples. +# + +"""Plays a raw audio file to a Telegram group call, with MediaStreamTrackGenerator. + +The file is 48 kHz stereo 16-bit PCM. Pyrogram reads its session from the SESSION_NAME, API_ID and API_HASH +environment variables. +""" + +from __future__ import annotations + import asyncio import json import os +import pathlib import time +from typing import TYPE_CHECKING, BinaryIO, TypedDict # pip install pytgcalls[pyrogram]==3.0.0.dev21 import pyrogram +from pytgcalls.mtproto.data import GroupCallWrapper +from pytgcalls.mtproto.data.update import UpdateGroupCallWrapper from pytgcalls.mtproto.pyrogram_bridge import PyrogramBridge import webrtc -remote_sdp = None +if TYPE_CHECKING: + from pytgcalls.mtproto.data.update import UpdateGroupCallParticipantsWrapper -def parse_sdp(sdp): - lines = sdp.split('\r\n') +class Fingerprint(TypedDict): + """A DTLS fingerprint of the Telegram server.""" - def lookup(prefix): - for line in lines: - if line.startswith(prefix): - return line[len(prefix) :] + fingerprint: str - info = { - 'fingerprint': lookup('a=fingerprint:').split(' ')[1], - 'hash': lookup('a=fingerprint:').split(' ')[0], - 'setup': lookup('a=setup:'), - 'pwd': lookup('a=ice-pwd:'), - 'ufrag': lookup('a=ice-ufrag:'), - } - ssrc = lookup('a=ssrc:') - if ssrc: - info['source'] = int(ssrc.split(' ')[0]) - return info +class Candidate(TypedDict): + """An ICE candidate of the Telegram server.""" + + foundation: str + component: str + protocol: str + priority: str + ip: str + port: str + type: str + generation: str + + +class Transport(TypedDict): + """The transport of the Telegram server.""" + + ufrag: str + pwd: str + fingerprints: list[Fingerprint] + candidates: list[Candidate] + + +class CallParams(TypedDict): + """The parameters of a joined call.""" + + transport: Transport + +def sdp_attributes(sdp: str) -> dict[str, str]: + """The first value of each attribute (``a=name:value``) of an SDP.""" + attributes: dict[str, str] = {} + for line in sdp.split('\r\n'): + if line.startswith('a='): + name, _, value = line[2:].partition(':') + attributes.setdefault(name, value) + return attributes -def get_params_from_parsed_sdp(info): + +def join_params(offer: str) -> dict[str, object]: + """The transport of the offer, as Telegram takes it to join a call.""" + attributes = sdp_attributes(offer) + hash_, fingerprint = attributes['fingerprint'].split(' ') return { - 'fingerprints': [{'fingerprint': info['fingerprint'], 'hash': info['hash'], 'setup': 'active'}], - 'pwd': info['pwd'], - 'ssrc': info['source'], + 'fingerprints': [{'fingerprint': fingerprint, 'hash': hash_, 'setup': 'active'}], + 'pwd': attributes['ice-pwd'], + 'ssrc': int(attributes['ssrc'].split(' ')[0]), 'ssrc-groups': [], - 'ufrag': info['ufrag'], + 'ufrag': attributes['ice-ufrag'], } -def build_answer(sdp): - def add_candidates(): - candidates_sdp = [] - for cand in sdp['transport']['candidates']: - candidates_sdp.append( - f"a=candidate:{cand['foundation']} {cand['component']} {cand['protocol']} " - f"{cand['priority']} {cand['ip']} {cand['port']} typ {cand['type']} " - f"generation {cand['generation']}" - ) - - return '\n'.join(candidates_sdp) - +def build_answer(params: CallParams) -> str: + """The SDP answer of the Telegram server.""" + transport = params['transport'] + candidates = '\n'.join( + f'a=candidate:{c["foundation"]} {c["component"]} {c["protocol"]} {c["priority"]} {c["ip"]} {c["port"]} ' + f'typ {c["type"]} generation {c["generation"]}' + for c in transport['candidates'] + ) return f"""v=0 o=- {time.time()} 2 IN IP4 0.0.0.0 s=- @@ -65,11 +106,11 @@ def add_candidates(): m=audio 1 RTP/SAVPF 111 126 c=IN IP4 0.0.0.0 a=mid:0 -a=ice-ufrag:{sdp['transport']['ufrag']} -a=ice-pwd:{sdp['transport']['pwd']} -a=fingerprint:sha-256 {sdp['transport']['fingerprints'][0]['fingerprint']} +a=ice-ufrag:{transport['ufrag']} +a=ice-pwd:{transport['pwd']} +a=fingerprint:sha-256 {transport['fingerprints'][0]['fingerprint']} a=setup:passive -{add_candidates()} +{candidates} a=rtpmap:111 opus/48000/2 a=rtpmap:126 telephone-event/8000 a=fmtp:111 minptime=10; useinbandfec=1; usedtx=1 @@ -79,45 +120,33 @@ def add_candidates(): a=extmap:1 urn:ietf:params:rtp-hdrext:ssrc-audio-level a=recvonly """ - # a=sendrecv - - -async def group_call_participants_update_callback(_): - pass -async def group_call_update_callback(update): - global remote_sdp - - data = update.call.params.data - remote_sdp = build_answer(json.loads(data)) - - -async def send_audio_data(generator, input_filename): - """Writes raw 48 kHz stereo 16-bit audio to the track, 10 ms at a time, at the pace of real time""" +async def send_audio_data(generator: webrtc.MediaStreamTrackGenerator, file: BinaryIO) -> None: + """Writes raw 48 kHz stereo 16-bit audio to the track, 10 ms at a time, at the pace of real time.""" writer = generator.writable.get_writer() loop = asyncio.get_running_loop() start = loop.time() chunks = 0 - with open(input_filename, 'rb') as f: - while data := f.read(480 * 4): # 480 frames of 2 channels of 16 bits - frames = len(data) // 4 - await writer.write( - webrtc.AudioData( - format='s16', - sample_rate=48000, - number_of_frames=frames, - number_of_channels=2, - timestamp=chunks * 10_000, - data=data[: frames * 4], - ) + while data := file.read(480 * 4): # 480 frames of 2 channels of 16 bits + frames = len(data) // 4 + await writer.write( + webrtc.AudioData( + format='s16', + sample_rate=48000, + number_of_frames=frames, + number_of_channels=2, + timestamp=chunks * 10_000, + data=data[: frames * 4], ) - chunks += 1 - await asyncio.sleep(max(0.0, start + chunks / 100 - loop.time())) + ) + chunks += 1 + await asyncio.sleep(max(0.0, start + chunks / 100 - loop.time())) -async def main(input_peer, input_filename): +async def main(input_peer: str, audio: BinaryIO) -> None: + """Joins the group call of the peer and plays the audio to it.""" client = pyrogram.Client( os.environ.get('SESSION_NAME'), api_hash=os.environ.get('API_HASH'), api_id=os.environ.get('API_ID') ) @@ -127,31 +156,36 @@ async def main(input_peer, input_filename): generator = webrtc.MediaStreamTrackGenerator('audio') pc.add_track(generator) - local_sdp = await pc.create_offer() - await pc.set_local_description(local_sdp) + offer = await pc.create_offer() + await pc.set_local_description(offer) + answered = asyncio.Event() + + async def on_update(update: UpdateGroupCallWrapper | UpdateGroupCallParticipantsWrapper) -> None: + # only the parameters of the joined call matter, the first time they come + if answered.is_set() or not isinstance(update, UpdateGroupCallWrapper): + return + if isinstance(update.call, GroupCallWrapper): + answered.set() + answer = build_answer(json.loads(update.call.params.data)) + await pc.set_remote_description( + webrtc.RTCSessionDescription(webrtc.RTCSessionDescriptionInit(webrtc.RTCSdpType.answer, answer)) + ) app = PyrogramBridge(client) - app.register_group_call_native_callback(group_call_participants_update_callback, group_call_update_callback) + app.register_group_call_native_callback(on_update, on_update) await app.get_and_set_group_call(input_peer) await app.resolve_and_set_join_as(None) - def pre_update_processing(): + def pre_update_processing() -> None: pass - parsed_sdp = parse_sdp(local_sdp.sdp) - payload = get_params_from_parsed_sdp(parsed_sdp) - - await app.join_group_call(None, json.dumps(payload), False, False, pre_update_processing) - - while not remote_sdp: - await asyncio.sleep(0.1) - # await asyncio.wait_for(REMOTE_ANSWER_EVENT.wait(), 30) - - await pc.set_remote_description( - webrtc.RTCSessionDescription(webrtc.RTCSessionDescriptionInit(webrtc.RTCSdpType.answer, remote_sdp)) + params = json.dumps(join_params(offer.sdp)) + await app.join_group_call( + None, params, muted=False, video_stopped=False, pre_update_processing=pre_update_processing ) + await asyncio.wait_for(answered.wait(), 30) - sending = asyncio.ensure_future(send_audio_data(generator, input_filename)) + sending = asyncio.ensure_future(send_audio_data(generator, audio)) await pyrogram.idle() sending.cancel() @@ -162,4 +196,5 @@ def pre_update_processing(): peer = input('Input peer:') filename = input('Input filename to play:') - asyncio.run(main(peer, filename)) + with pathlib.Path(filename).open('rb') as file: + asyncio.run(main(peer, file)) diff --git a/pyproject.toml b/pyproject.toml index 4509a90..d8ef7df 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -57,7 +57,7 @@ dev = [ { include-group = "test" }, "scikit-build-core>=0.11", "pybind11>=3.0", - "ruff>=0.13", + "ruff>=0.16.9", "pybind11-stubgen>=2.5", ] wpt = [ @@ -110,22 +110,68 @@ timeout_method = "thread" faulthandler_timeout = 60 [tool.ruff] +required-version = ">=0.16.9" line-length = 120 target-version = "py39" +preview = true include = [ "python-webrtc/python/webrtc/**/*.py", "tests/**/*.py", "examples/**/*.py", "benchmarks/**/*.py", "cmake/**/*.py", ] extend-exclude = ["*.pyi"] [tool.ruff.format] -quote-style = "preserve" +quote-style = "single" +docstring-code-format = true [tool.ruff.lint] -select = ["E", "F", "W", "I", "UP", "B"] -ignore = ["UP006", "UP007", "UP035", "UP045", "B024"] +select = ["ALL"] +ignore = [ + "COM812", "Q", # the formatter owns trailing commas and quotes + "TRY003", # a message per raise, not an exception class per message + "D105", "D107", # class docstrings document constructors, dunders explain themselves + "DOC502", # exceptions raised by native calls and helpers are documented too +] [tool.ruff.lint.per-file-ignores] -"__init__.py" = ["I001", "E402"] +"__init__.py" = ["I001", "E402"] # the import order resolves the cycles +# camelCase aliases of the W3C names; spec internal slots shared between the classes of the package +"python-webrtc/python/webrtc/**" = ["N815", "N816", "SLF001"] +"python-webrtc/python/webrtc/interfaces/rtc_peer_connection.py" = ["PLR0904"] # the W3C interface +# user callbacks error the stream, as the specification says; handler errors go to the loop's exception handler +"python-webrtc/python/webrtc/{streams.py,utils/events.py}" = ["BLE001"] +"python-webrtc/python/webrtc/models/rtc_stats.py" = ["FURB189"] # a real dict, for json and isinstance +# native leak counters and queues have no public API; chaos replays a seeded random +"tests/**" = ["S101", "D1", "PLR2004", "S404", "S603", "SLF001", "S311"] +"examples/**" = ["T201"] +"benchmarks/**" = ["T201", "S404", "S603"] +"cmake/**" = ["T201", "S404", "S603"] + +[tool.ruff.lint.pydocstyle] +convention = "google" + +[tool.ruff.lint.flake8-builtins] +ignorelist = ["format", "id", "type"] # W3C member names + +[tool.ruff.lint.pydoclint] +ignore-one-line-docstrings = true + +[tool.ruff.lint.flake8-annotations] +mypy-init-return = false + +[tool.ruff.lint.mccabe] +max-complexity = 8 + +[tool.ruff.lint.pylint] +max-args = 5 +max-positional-args = 3 +max-branches = 8 +max-returns = 4 +max-statements = 30 +max-locals = 10 +max-nested-blocks = 3 +max-bool-expr = 3 +max-public-methods = 20 [tool.ruff.lint.isort] +required-imports = ["from __future__ import annotations"] known-first-party = ["webrtc", "wrtc", "tests", "benchmarks"] diff --git a/python-webrtc/cpp/src/exceptions.cpp b/python-webrtc/cpp/src/exceptions.cpp index c259187..8b42263 100644 --- a/python-webrtc/cpp/src/exceptions.cpp +++ b/python-webrtc/cpp/src/exceptions.cpp @@ -32,8 +32,9 @@ namespace python_webrtc { } // the exception classes are defined in Python, to be subclassed and constructed like any Python exception return pybind11::module_::import("webrtc.exceptions") - .attr("_from_native")(std::string(ToString(error.type())), std::string(error.message()), detail, sctpCauseCode, - lineNumber); + .attr("_from_native")(std::string(ToString(error.type())), std::string(error.message()), + pybind11::arg("detail") = detail, pybind11::arg("sctp_cause_code") = sctpCauseCode, + pybind11::arg("sdp_line_number") = lineNumber); } pybind11::object RTCCallbackException::ToPython() const { diff --git a/python-webrtc/python/webrtc/__init__.py b/python-webrtc/python/webrtc/__init__.py index 544b249..61adc6d 100644 --- a/python-webrtc/python/webrtc/__init__.py +++ b/python-webrtc/python/webrtc/__init__.py @@ -5,7 +5,11 @@ # that can be found in the LICENSE.md file in the root of the project. # -import wrtc # noqa: F401 (the modules import it from the package) +"""Python bindings to WebRTC, with the API of the browsers.""" + +from __future__ import annotations + +import wrtc as wrtc # re-exported, the modules import it from the package from .enums import ( RTCPeerConnectionState, @@ -60,6 +64,7 @@ InvalidCharacterError, OverconstrainedError, RTCError, + RTCErrorInit, ) from .utils.events import EventTarget from .models.events import ( @@ -142,7 +147,7 @@ from .interfaces.rtc_ice_transport import RTCIceTransport from .interfaces.rtc_dtls_transport import RTCDtlsTransport from .interfaces.rtc_sctp_transport import RTCSctpTransport -from .interfaces.rtc_data_channel import RTCDataChannel +from .interfaces.rtc_data_channel import RTCDataChannel, RTCDataChannelInit from .interfaces.media_stream_track_processor import MediaStreamTrackProcessorInit, MediaStreamTrackProcessor from .interfaces.track_generator import ( VideoTrackGenerator, @@ -153,137 +158,136 @@ from .functions.get_user_media import getUserMedia, get_user_media -#: Alias for :obj:`RTCRtpEncodingParameters` -RtpEncodingParameters = RTCRtpEncodingParameters __all__ = [ - 'PythonWebRTCExceptionBase', - 'PythonWebRTCException', - 'RTCException', - 'SdpParseException', - 'InvalidStateError', + 'AlphaOption', + 'AudioData', + 'AudioDataCopyToOptions', + 'AudioDataInit', + 'AudioSampleFormat', + 'BinaryType', + 'Blob', + 'CricketIceGatheringState', + 'DOMRectReadOnly', + 'DoubleRange', + 'DtlsTransportState', + 'Event', + 'EventTarget', 'InvalidAccessError', + 'InvalidCharacterError', 'InvalidModificationError', - 'OperationError', - 'NotSupportedError', - 'NetworkError', - 'InvalidSyntaxError', 'InvalidRangeError', - 'InvalidCharacterError', - 'OverconstrainedError', - 'RTCError', - 'RTCPeerConnectionState', - 'RTCSignalingState', - 'RTCIceConnectionState', - 'RTCIceGatheringState', - 'RTCSdpType', - 'MediaStreamTrackState', + 'InvalidStateError', + 'InvalidSyntaxError', + 'MediaStream', 'MediaStreamSourceState', - 'TransceiverDirection', - 'RTCIceComponent', - 'RTCIceRole', - 'RTCIceTransportState', - 'CricketIceGatheringState', - 'DtlsTransportState', - 'SctpTransportState', + 'MediaStreamTrack', + 'MediaStreamTrackEvent', + 'MediaStreamTrackGenerator', + 'MediaStreamTrackGeneratorInit', + 'MediaStreamTrackProcessor', + 'MediaStreamTrackProcessorInit', + 'MediaStreamTrackState', + 'MediaTrackCapabilities', + 'MediaTrackConstraints', + 'MediaTrackSettings', 'MediaType', + 'MessageEvent', + 'NetworkError', + 'NotSupportedError', + 'OperationError', + 'OverconstrainedError', + 'PlaneLayout', + 'PythonWebRTCException', + 'PythonWebRTCExceptionBase', + 'RTCBundlePolicy', + 'RTCCertificate', + 'RTCConfiguration', + 'RTCDTMFSender', + 'RTCDTMFToneChangeEvent', + 'RTCDataChannel', + 'RTCDataChannelEvent', + 'RTCDataChannelInit', + 'RTCDataChannelState', + 'RTCDegradationPreference', + 'RTCDtlsFingerprint', + 'RTCDtlsTransport', + 'RTCError', 'RTCErrorDetailType', - 'BinaryType', - 'VideoPixelFormat', - 'VideoColorPrimaries', - 'VideoTransferCharacteristics', - 'VideoMatrixCoefficients', - 'AlphaOption', - 'AudioSampleFormat', + 'RTCErrorEvent', + 'RTCErrorInit', + 'RTCException', + 'RTCIceCandidate', + 'RTCIceCandidatePair', 'RTCIceCandidateType', + 'RTCIceComponent', + 'RTCIceConnectionState', + 'RTCIceGatheringState', + 'RTCIceParameters', 'RTCIceProtocol', - 'RTCIceTcpCandidateType', + 'RTCIceRole', + 'RTCIceServer', 'RTCIceServerTransportProtocol', + 'RTCIceTcpCandidateType', + 'RTCIceTransport', 'RTCIceTransportPolicy', - 'RTCBundlePolicy', + 'RTCIceTransportState', + 'RTCOAuthCredential', + 'RTCPeerConnection', + 'RTCPeerConnectionIceErrorEvent', + 'RTCPeerConnectionIceEvent', + 'RTCPeerConnectionState', + 'RTCPriorityType', 'RTCRtcpMuxPolicy', + 'RTCRtcpParameters', + 'RTCRtpCapabilities', + 'RTCRtpCodec', + 'RTCRtpCodecParameters', + 'RTCRtpContributingSource', + 'RTCRtpEncodingParameters', 'RTCRtpHeaderEncryptionPolicy', - 'RTCDataChannelState', - 'RTCPriorityType', - 'RTCDegradationPreference', - 'WebRTCObject', - 'EventTarget', - 'Event', - 'RTCPeerConnectionIceEvent', - 'RTCPeerConnectionIceErrorEvent', - 'RTCTrackEvent', - 'RTCErrorEvent', - 'MessageEvent', - 'RTCDataChannelEvent', - 'MediaStreamTrackEvent', - 'RTCDTMFToneChangeEvent', - 'RTCPeerConnection', - 'MediaStreamTrack', - 'MediaStream', - 'RTCRtpSender', + 'RTCRtpHeaderExtensionCapability', + 'RTCRtpHeaderExtensionParameters', + 'RTCRtpReceiveParameters', 'RTCRtpReceiver', + 'RTCRtpSendParameters', + 'RTCRtpSender', + 'RTCRtpSynchronizationSource', 'RTCRtpTransceiver', - 'RTCIceTransport', - 'RTCDtlsTransport', 'RTCSctpTransport', - 'RTCDataChannel', - 'MediaStreamTrackProcessorInit', - 'MediaStreamTrackProcessor', - 'VideoTrackGenerator', - 'MediaStreamTrackGeneratorInit', - 'MediaStreamTrackGenerator', - 'RTCDTMFSender', - 'getUserMedia', - 'get_user_media', - 'RTCSessionDescriptionInit', + 'RTCSdpType', 'RTCSessionDescription', + 'RTCSessionDescriptionInit', + 'RTCSignalingState', + 'RTCStats', + 'RTCStatsReport', + 'RTCTrackEvent', + 'ReadableStream', + 'ReadableStreamDefaultController', + 'ReadableStreamDefaultReader', + 'ReadableStreamReadResult', + 'RtpTransceiverInit', + 'SctpTransportState', + 'SdpParseException', + 'TransceiverDirection', + 'TransformStream', + 'TransformStreamDefaultController', 'ULongRange', - 'DoubleRange', - 'MediaTrackSettings', - 'MediaTrackCapabilities', - 'MediaTrackConstraints', - 'Blob', - 'DOMRectReadOnly', - 'PlaneLayout', + 'VideoColorPrimaries', 'VideoColorSpace', - 'VideoFrameMetadata', + 'VideoFrame', 'VideoFrameBufferInit', - 'VideoFrameInit', 'VideoFrameCopyToOptions', - 'VideoFrame', - 'AudioDataInit', - 'AudioDataCopyToOptions', - 'AudioData', - 'ReadableStream', - 'ReadableStreamDefaultReader', - 'ReadableStreamDefaultController', - 'ReadableStreamReadResult', + 'VideoFrameInit', + 'VideoFrameMetadata', + 'VideoMatrixCoefficients', + 'VideoPixelFormat', + 'VideoTrackGenerator', + 'VideoTransferCharacteristics', + 'WebRTCObject', 'WritableStream', - 'WritableStreamDefaultWriter', 'WritableStreamDefaultController', - 'TransformStream', - 'TransformStreamDefaultController', - 'RtpEncodingParameters', - 'RTCRtpCodec', - 'RTCRtpCodecParameters', - 'RTCRtpHeaderExtensionParameters', - 'RTCRtcpParameters', - 'RTCRtpEncodingParameters', - 'RTCRtpReceiveParameters', - 'RTCRtpSendParameters', - 'RTCRtpHeaderExtensionCapability', - 'RTCRtpCapabilities', - 'RtpTransceiverInit', - 'RTCIceCandidate', - 'RTCIceCandidatePair', - 'RTCIceParameters', - 'RTCStats', - 'RTCRtpContributingSource', - 'RTCRtpSynchronizationSource', - 'RTCStatsReport', - 'RTCCertificate', - 'RTCDtlsFingerprint', - 'RTCConfiguration', - 'RTCIceServer', - 'RTCOAuthCredential', + 'WritableStreamDefaultWriter', + 'getUserMedia', + 'get_user_media', ] diff --git a/python-webrtc/python/webrtc/base.py b/python-webrtc/python/webrtc/base.py index 19f1657..c3d3b72 100644 --- a/python-webrtc/python/webrtc/base.py +++ b/python-webrtc/python/webrtc/base.py @@ -5,51 +5,67 @@ # that can be found in the LICENSE.md file in the root of the project. # -from abc import ABCMeta -from typing import List, Optional +"""The base class of the wrappers of native objects.""" + +from __future__ import annotations + +from typing import TYPE_CHECKING, Any, Callable, ClassVar, Generic, TypeVar from webrtc.utils.events import EventTarget +if TYPE_CHECKING: + from collections.abc import Iterable + + from typing_extensions import Self + +_NativeT = TypeVar('_NativeT') + -class WebRTCObject(metaclass=ABCMeta): - _class = None +class WebRTCObject(Generic[_NativeT]): + """The wrapper of a native object. Wrappers are equal when they wrap the same native object. - def __init__(self, native_obj=None): - self.__obj = native_obj - if not self.__obj: - self.__obj = self._class() + Args: + native_obj (optional): The native object, a new one of the native class if omitted. + """ + + #: The native class, created with no arguments when no native object is given + _class: ClassVar[Callable[[], Any] | None] = None + + def __init__(self, native_obj: _NativeT | None = None) -> None: + self._init_native(native_obj) + + def _init_native(self, native_obj: _NativeT | None) -> None: + self.__obj = native_obj or self._class() @property - def _native_obj(self): + def _native_obj(self) -> _NativeT: return self.__obj - def _set_native_obj(self, value): - self.__obj = value - @classmethod - def _wrap(cls, item) -> 'WebRTCObject': - """The wrapper of a native object. Constructors take the arguments of the public API, so it's not created - with the constructor of the class.""" + def _wrap(cls, item: _NativeT) -> Self: + """The wrapper of a native object, created without the constructor, which takes the public arguments.""" obj = cls.__new__(cls) - WebRTCObject.__init__(obj, item) + obj._init_native(item) if isinstance(obj, EventTarget): obj._attach() return obj @classmethod - def _wrap_optional(cls, item) -> Optional['WebRTCObject']: + def _wrap_optional(cls, item: _NativeT | None) -> Self | None: """The wrapper of a native object, or :obj:`None` for :obj:`None`.""" return cls._wrap(item) if item is not None else None @classmethod - def _wrap_many(cls, items) -> List['WebRTCObject']: + def _wrap_many(cls, items: Iterable[_NativeT]) -> list[Self]: return [cls._wrap(item) for item in items] - def __repr__(self): + def __repr__(self) -> str: return f' bool: if isinstance(other, WebRTCObject): - return id(self._native_obj) == id(other._native_obj) + return self._native_obj is other._native_obj + return NotImplemented - return super().__eq__(other) + def __hash__(self) -> int: + return id(self._native_obj) diff --git a/python-webrtc/python/webrtc/enums.py b/python-webrtc/python/webrtc/enums.py index bb9250f..26bcb0e 100644 --- a/python-webrtc/python/webrtc/enums.py +++ b/python-webrtc/python/webrtc/enums.py @@ -7,6 +7,8 @@ """Enums with the spec string values: ``RTCSignalingState.have_local_offer == 'have-local-offer'``.""" +from __future__ import annotations + import enum @@ -266,9 +268,11 @@ class BinaryType(_StrEnum): class VideoPixelFormat(_StrEnum): - """The layout of the pixels of a :obj:`webrtc.VideoFrame`: Y, U, V (and alpha) planes of 8-bit samples, or of - 10- and 12-bit samples stored in 16 bits (``P10``, ``P12``), NV12 with interleaved U and V, or 4 bytes per RGB - pixel.""" + """The layout of the pixels of a :obj:`webrtc.VideoFrame`. + + Y, U, V (and alpha) planes of 8-bit samples, or of 10- and 12-bit samples stored in 16 bits (``P10``, ``P12``), + NV12 with interleaved U and V, or 4 bytes per RGB pixel. + """ I420 = 'I420' I420P10 = 'I420P10' diff --git a/python-webrtc/python/webrtc/exceptions.py b/python-webrtc/python/webrtc/exceptions.py index ad9149e..8e9eb9a 100644 --- a/python-webrtc/python/webrtc/exceptions.py +++ b/python-webrtc/python/webrtc/exceptions.py @@ -7,10 +7,16 @@ """Exceptions raised by WebRTC operations, named after the ``DOMException`` of the specification.""" -from typing import Optional +from __future__ import annotations + +from dataclasses import dataclass, fields +from typing import TYPE_CHECKING, Any, ClassVar from webrtc import RTCErrorDetailType, wrtc -from webrtc.utils.names import alias +from webrtc.utils.names import Alias, alias, members + +if TYPE_CHECKING: + from collections.abc import Mapping PythonWebRTCExceptionBase = wrtc.PythonWebRTCExceptionBase PythonWebRTCException = wrtc.PythonWebRTCException @@ -65,17 +71,17 @@ class OverconstrainedError(RTCException): message (:obj:`str`, optional): A description of the error. """ - def __init__(self, constraint: str, message: str = ''): - super().__init__(message or f'The constraint {constraint} can\'t be satisfied') + def __init__(self, constraint: str, message: str = '') -> None: + super().__init__(message or f"The constraint {constraint} can't be satisfied") self.constraint = constraint -class RTCError(OperationError): - """An error carrying WebRTC-specific information. +@dataclass +class RTCErrorInit: + """The WebRTC-specific information of an :obj:`RTCError`. Args: error_detail (:obj:`RTCErrorDetailType`): The WebRTC-specific error code. - message (:obj:`str`, optional): A description of the error. sdp_line_number (:obj:`int`, optional): The line of the SDP where a syntax error occurred. sctp_cause_code (:obj:`int`, optional): The SCTP cause code of a failed SCTP negotiation. received_alert (:obj:`int`, optional): The DTLS alert received from the remote peer. @@ -83,28 +89,60 @@ class RTCError(OperationError): http_request_status_code (:obj:`int`, optional): The HTTP status code of a failed request. Raises: - :obj:`ValueError`: If ``error_detail`` isn't a member of :obj:`RTCErrorDetailType`. + ValueError: If ``error_detail`` isn't a member of :obj:`RTCErrorDetailType`. """ - def __init__( - self, - error_detail: RTCErrorDetailType, - message: str = '', - *, - sdp_line_number: Optional[int] = None, - sctp_cause_code: Optional[int] = None, - received_alert: Optional[int] = None, - sent_alert: Optional[int] = None, - http_request_status_code: Optional[int] = None, - ): + error_detail: RTCErrorDetailType + sdp_line_number: int | None = None + sctp_cause_code: int | None = None + received_alert: int | None = None + sent_alert: int | None = None + http_request_status_code: int | None = None + + def __post_init__(self) -> None: + self.error_detail = RTCErrorDetailType(self.error_detail) + + #: Alias for :attr:`error_detail` + errorDetail: ClassVar[Alias[RTCErrorDetailType]] = alias('error_detail') + #: Alias for :attr:`sdp_line_number` + sdpLineNumber: ClassVar[Alias[int | None]] = alias('sdp_line_number') + #: Alias for :attr:`sctp_cause_code` + sctpCauseCode: ClassVar[Alias[int | None]] = alias('sctp_cause_code') + #: Alias for :attr:`received_alert` + receivedAlert: ClassVar[Alias[int | None]] = alias('received_alert') + #: Alias for :attr:`sent_alert` + sentAlert: ClassVar[Alias[int | None]] = alias('sent_alert') + #: Alias for :attr:`http_request_status_code` + httpRequestStatusCode: ClassVar[Alias[int | None]] = alias('http_request_status_code') + + +class RTCError(OperationError): + """An error carrying WebRTC-specific information, the members of its :obj:`RTCErrorInit`. + + Args: + options (:obj:`RTCErrorInit` or :obj:`dict`): The WebRTC-specific information, or a dictionary of its + members. + message (:obj:`str`, optional): A description of the error. + + Raises: + TypeError: If a dictionary has no ``error_detail``. + ValueError: If ``error_detail`` isn't a member of :obj:`RTCErrorDetailType`. + """ + + def __init__(self, options: RTCErrorInit | Mapping[str, Any], message: str = '') -> None: + init = ( + options + if isinstance(options, RTCErrorInit) + else RTCErrorInit(**members(options, [field.name for field in fields(RTCErrorInit)])) + ) super().__init__(message) - self.error_detail = RTCErrorDetailType(error_detail) self.message = message - self.sdp_line_number = sdp_line_number - self.sctp_cause_code = sctp_cause_code - self.received_alert = received_alert - self.sent_alert = sent_alert - self.http_request_status_code = http_request_status_code + self.error_detail = init.error_detail + self.sdp_line_number = init.sdp_line_number + self.sctp_cause_code = init.sctp_cause_code + self.received_alert = init.received_alert + self.sent_alert = init.sent_alert + self.http_request_status_code = init.http_request_status_code #: Alias for :attr:`error_detail` errorDetail = alias('error_detail') @@ -138,17 +176,17 @@ def __init__( def _from_native( error_type: str, message: str, - detail: Optional[RTCErrorDetailType], - sctp_cause_code: Optional[int], - sdp_line_number: Optional[int] = None, + *, + detail: RTCErrorDetailType | None, + sctp_cause_code: int | None, + sdp_line_number: int | None, ) -> RTCException: - """Creates the exception for a webrtc::RTCError. Called by name, with positional arguments, from - cpp/src/exceptions.cpp.""" + """Creates the exception for a webrtc::RTCError, called from cpp/src/exceptions.cpp.""" if error_type == 'OPERATION_ERROR_WITH_DATA' or detail is not None: - return RTCError( + init = RTCErrorInit( detail or RTCErrorDetailType.data_channel_failure, - message, sctp_cause_code=sctp_cause_code, sdp_line_number=sdp_line_number, ) + return RTCError(init, message) return _BY_RTC_ERROR_TYPE.get(error_type, OperationError)(message) diff --git a/python-webrtc/python/webrtc/functions/__init__.py b/python-webrtc/python/webrtc/functions/__init__.py index 548a80a..a36721d 100644 --- a/python-webrtc/python/webrtc/functions/__init__.py +++ b/python-webrtc/python/webrtc/functions/__init__.py @@ -4,3 +4,5 @@ # Use of this source code is governed by a BSD-style license # that can be found in the LICENSE.md file in the root of the project. # + +"""The functions of the API.""" diff --git a/python-webrtc/python/webrtc/functions/get_user_media.py b/python-webrtc/python/webrtc/functions/get_user_media.py index 6521ce3..f907032 100644 --- a/python-webrtc/python/webrtc/functions/get_user_media.py +++ b/python-webrtc/python/webrtc/functions/get_user_media.py @@ -5,7 +5,11 @@ # that can be found in the LICENSE.md file in the root of the project. # -from typing import TYPE_CHECKING, Dict, Optional, Union +"""getUserMedia of Media Capture and Streams, with a synthetic microphone and camera.""" + +from __future__ import annotations + +from typing import TYPE_CHECKING, Union from webrtc import MediaStream, MediaTrackConstraints, MediaTrackSettings, OverconstrainedError, wrtc from webrtc.interfaces.media_stream_track import _CAMERA_CAPABILITIES, _check_numbers, _selected, _unsatisfied @@ -15,21 +19,22 @@ #: A value, or a constraint on it: a :obj:`dict` with any of ``exact``, ``ideal``, ``min`` and ``max`` -Constrain = Union[float, Dict[str, float]] +Constrain = Union[float, dict[str, float]] def get_user_media( + *, audio: bool = True, video: bool = False, - *, - width: Optional[Constrain] = None, - height: Optional[Constrain] = None, - frame_rate: Optional[Constrain] = None, -) -> 'webrtc.MediaStream': - """Returns a stream of local media, as requested: an audio track of a synthetic microphone (quiet noise), and/or - a video track of a synthetic camera, which draws a moving pattern (use :obj:`webrtc.VideoTrackGenerator` and - :obj:`webrtc.MediaStreamTrackGenerator` for real media). The constraints given are the ones of the tracks (see - :meth:`webrtc.MediaStreamTrack.get_constraints`). + width: Constrain | None = None, + height: Constrain | None = None, + frame_rate: Constrain | None = None, +) -> webrtc.MediaStream: + """Returns a stream of local media, as requested: a synthetic microphone and/or camera. + + The audio track is quiet noise, the video track draws a moving pattern (use :obj:`webrtc.VideoTrackGenerator` + and :obj:`webrtc.MediaStreamTrackGenerator` for real media). The constraints given are the ones of the tracks + (see :meth:`webrtc.MediaStreamTrack.get_constraints`). Args: audio (:obj:`bool`, optional): Whether the stream has an audio track. @@ -44,14 +49,15 @@ def get_user_media( :obj:`webrtc.MediaStream`: The stream. Raises: - :obj:`TypeError`: If neither audio nor video is requested, or a value isn't a finite number (negative for + TypeError: If neither audio nor video is requested, or a value isn't a finite number (negative for the size). - :obj:`webrtc.OverconstrainedError`: If a required value (``exact``, ``min``, ``max``) is beyond what the + webrtc.OverconstrainedError: If a required value (``exact``, ``min``, ``max``) is beyond what the camera can do: 1 to 4096 pixels wide and high, 1 to 120 frames per second. Other values are brought within that. """ if not audio and not video: - raise TypeError('audio or video must be requested') + msg = 'audio or video must be requested' + raise TypeError(msg) constraints = MediaTrackConstraints(width=width, height=height, frame_rate=frame_rate) _check_numbers(constraints) if video: diff --git a/python-webrtc/python/webrtc/interfaces/__init__.py b/python-webrtc/python/webrtc/interfaces/__init__.py index 548a80a..f9c1fce 100644 --- a/python-webrtc/python/webrtc/interfaces/__init__.py +++ b/python-webrtc/python/webrtc/interfaces/__init__.py @@ -4,3 +4,5 @@ # Use of this source code is governed by a BSD-style license # that can be found in the LICENSE.md file in the root of the project. # + +"""The interfaces of the WebRTC specifications, wrapping native objects.""" diff --git a/python-webrtc/python/webrtc/interfaces/media_stream.py b/python-webrtc/python/webrtc/interfaces/media_stream.py index d433ab1..8321cb9 100644 --- a/python-webrtc/python/webrtc/interfaces/media_stream.py +++ b/python-webrtc/python/webrtc/interfaces/media_stream.py @@ -5,18 +5,26 @@ # that can be found in the LICENSE.md file in the root of the project. # -from typing import TYPE_CHECKING, List, Optional, Union +"""MediaStream of Media Capture and Streams.""" + +from __future__ import annotations + +from typing import TYPE_CHECKING from webrtc import MediaStreamTrack, MediaStreamTrackEvent, MediaType, WebRTCObject, wrtc from webrtc.utils.events import EventTarget if TYPE_CHECKING: + from typing_extensions import Self + import webrtc -class MediaStream(WebRTCObject, EventTarget): - """The MediaStream interface represents a stream of media content. A stream consists of several tracks, - such as video or audio tracks. Each track is specified as an instance of :obj:`webrtc.MediaStreamTrack`. +class MediaStream(WebRTCObject[wrtc.MediaStream], EventTarget): + """The MediaStream interface represents a stream of media content. + + A stream consists of several tracks, such as video or audio tracks. Each track is specified as an instance of + :obj:`webrtc.MediaStreamTrack`. Events (see :meth:`on`): ``addtrack`` and ``removetrack`` (:obj:`webrtc.MediaStreamTrackEvent`): The remote peer added a track to @@ -30,35 +38,34 @@ class MediaStream(WebRTCObject, EventTarget): _class = wrtc.MediaStream _events = ('addtrack', 'removetrack') - def __init__(self, tracks: Optional[Union[List['webrtc.MediaStreamTrack'], 'webrtc.MediaStream']] = None): + def __init__(self, tracks: list[webrtc.MediaStreamTrack] | webrtc.MediaStream | None = None) -> None: if isinstance(tracks, MediaStream): tracks = tracks.get_tracks() super().__init__(self._class.create([track._native_obj for track in tracks or []])) self._keep_tracks() @classmethod - def _wrap(cls, item) -> 'MediaStream': + def _wrap(cls, item: wrtc.MediaStream) -> Self: stream = super()._wrap(item) stream._keep_tracks() return stream - def _keep_tracks(self) -> list: - """The native tracks, kept here: the native stream keeps them weakly""" + def _keep_tracks(self) -> list[wrtc.MediaStreamTrack]: + """The native tracks, kept here: the native stream keeps them weakly.""" self._tracks = self._native_obj.getTracks() return self._tracks - def _on_event(self, name: str, *args): - if name in ('addtrack', 'removetrack'): + def _on_event(self, name: str, *_args: object) -> None: + if name in {'addtrack', 'removetrack'}: self._keep_tracks() - def _create_event(self, name: str, *args): + def _create_event(self, name: str, *args: object) -> webrtc.Event | None: (track,) = args return MediaStreamTrackEvent(name, MediaStreamTrack._wrap(track), target=self) @property def id(self) -> str: - """:obj:`str`: A String containing 36 characters denoting a - universally unique identifier (UUID) for the object.""" + """:obj:`str`: The universally unique identifier (UUID) of the stream, 36 characters.""" return self._native_obj.id @property @@ -66,53 +73,67 @@ def active(self) -> bool: """:obj:`bool`: Whether the :obj:`webrtc.MediaStream` is active: whether a track of it isn't ended.""" return self._native_obj.active - def get_audio_tracks(self) -> List['webrtc.MediaStreamTrack']: - """Returns a :obj:`list` of the :obj:`webrtc.MediaStreamTrack` objects - stored in the :obj:`webrtc.MediaStream` object that have their kind attribute set to "audio". - The order is not defined, and may not only vary from one machine to another, but also from one call to another. + def get_audio_tracks(self) -> list[webrtc.MediaStreamTrack]: + """Returns the audio tracks of the stream, in no defined order. + + Returns: + :obj:`list` of :obj:`webrtc.MediaStreamTrack`: The tracks. """ return MediaStreamTrack._wrap_many([t for t in self._keep_tracks() if t.kind == MediaType.audio]) - def get_video_tracks(self) -> List['webrtc.MediaStreamTrack']: - """Returns a :obj:`list` of the :obj:`webrtc.MediaStreamTrack` objects stored in the :obj:`webrtc.MediaStream` - object that have their kind attribute set to "video". The order is not defined, - and may not only vary from one machine to another, but also from one call to another. + def get_video_tracks(self) -> list[webrtc.MediaStreamTrack]: + """Returns the video tracks of the stream, in no defined order. + + Returns: + :obj:`list` of :obj:`webrtc.MediaStreamTrack`: The tracks. """ return MediaStreamTrack._wrap_many([t for t in self._keep_tracks() if t.kind == MediaType.video]) - def get_tracks(self) -> List['webrtc.MediaStreamTrack']: - """Returns a :obj:`list` of all :obj:`webrtc.MediaStreamTrack` objects stored in the :obj:`webrtc.MediaStream` - object, regardless of the value of the kind attribute. The order is not defined, - and may not only vary from one machine to another, but also from one call to another. + def get_tracks(self) -> list[webrtc.MediaStreamTrack]: + """Returns all the tracks of the stream, in no defined order. + + Returns: + :obj:`list` of :obj:`webrtc.MediaStreamTrack`: The tracks. """ return MediaStreamTrack._wrap_many(self._keep_tracks()) - def get_track_by_id(self, track_id: str) -> Optional['webrtc.MediaStreamTrack']: - """Returns the track whose ID corresponds to the one given in parameters, :obj:`track_id`. - If no track with that ID does exist, it returns :obj:`None`. - If several tracks have the same ID, it returns the first one. + def get_track_by_id(self, track_id: str) -> webrtc.MediaStreamTrack | None: + """Returns the track of an ID, the first one if several tracks have it. + + Args: + track_id (:obj:`str`): The ID. + + Returns: + :obj:`webrtc.MediaStreamTrack`, optional: The track, :obj:`None` if no track has the ID. """ track = self._native_obj.getTrackById(track_id) self._keep_tracks() return MediaStreamTrack._wrap_optional(track) - def add_track(self, track: 'webrtc.MediaStreamTrack'): - """Stores a copy of the :obj:`webrtc.MediaStreamTrack` given as argument. If the track has already been added - to the :obj:`webrtc.MediaStream` object, nothing happens. + def add_track(self, track: webrtc.MediaStreamTrack) -> None: + """Adds a track to the stream, unless it's there already. + + Args: + track (:obj:`webrtc.MediaStreamTrack`): The track. """ self._native_obj.addTrack(track._native_obj) self._keep_tracks() - def remove_track(self, track: 'webrtc.MediaStreamTrack'): - """Removes the :obj:`webrtc.MediaStreamTrack` given as argument. If the track is not part of the - :obj:`webrtc.MediaStream` object, nothing happens. + def remove_track(self, track: webrtc.MediaStreamTrack) -> None: + """Removes a track from the stream, if it's there. + + Args: + track (:obj:`webrtc.MediaStreamTrack`): The track. """ self._native_obj.removeTrack(track._native_obj) self._keep_tracks() - def clone(self) -> 'webrtc.MediaStream': - """Returns a clone of the :obj:`webrtc.MediaStream` object. - The clone will, however, have a unique value for :obj:`id`.""" + def clone(self) -> webrtc.MediaStream: + """Returns a clone of the stream, with a new :attr:`id`. + + Returns: + :obj:`webrtc.MediaStream`: The clone. + """ return self._wrap(self._native_obj.clone()) #: Alias for :attr:`get_audio_tracks` diff --git a/python-webrtc/python/webrtc/interfaces/media_stream_track.py b/python-webrtc/python/webrtc/interfaces/media_stream_track.py index 70319a9..0a1bc7e 100644 --- a/python-webrtc/python/webrtc/interfaces/media_stream_track.py +++ b/python-webrtc/python/webrtc/interfaces/media_stream_track.py @@ -5,9 +5,13 @@ # that can be found in the LICENSE.md file in the root of the project. # +"""MediaStreamTrack of Media Capture and Streams, and the constraints of its synthetic sources.""" + +from __future__ import annotations + import asyncio import math -from typing import TYPE_CHECKING, Any, Dict, Optional +from typing import TYPE_CHECKING, Any from webrtc import ( DoubleRange, @@ -67,39 +71,50 @@ ) -def _satisfied(value: Any, capability: Any, current: Any) -> bool: - """Whether the required parts of a constraint (exact, min, max) are satisfiable: within the capability of the - source if it has one, by the current setting if it doesn't""" - if not isinstance(value, dict): +def _satisfied(value: object, capability: object, current: float | str | None) -> bool: + """Whether the required parts of a constraint (exact, min, max) are satisfiable. + + They're satisfiable within the capability of the source if it has one, by the current setting if it doesn't. + + Returns: + :obj:`bool`: Whether they are. + """ + if not isinstance(value, dict) or all(value.get(key) is None for key in ('exact', 'min', 'max')): return True + if isinstance(capability, (ULongRange, DoubleRange)): + return _within_range(value, capability) + if capability is None: + return _satisfied_by_setting(value, current) + return _matches(value.get('exact'), capability) + + +def _within_range(value: dict[str, float], capability: ULongRange | DoubleRange) -> bool: exact, low, high = value.get('exact'), value.get('min'), value.get('max') - if exact is None and low is None and high is None: + low_cap = capability.min if capability.min is not None else float('-inf') + high_cap = capability.max if capability.max is not None else float('inf') + exact_within = exact is None or low_cap <= exact <= high_cap + bounds_within = (low is None or low <= high_cap) and (high is None or high >= low_cap) + return exact_within and bounds_within and (low is None or high is None or low <= high) + + +def _matches(exact: object, capability: object) -> bool: + """Whether an exact value (or one of a list of them) is the capability, or one of a list of them.""" + if exact is None: return True - if isinstance(capability, (ULongRange, DoubleRange)): - low_cap = capability.min if capability.min is not None else float('-inf') - high_cap = capability.max if capability.max is not None else float('inf') - if exact is not None and not low_cap <= exact <= high_cap: - return False - if low is not None and low > high_cap: - return False - if high is not None and high < low_cap: - return False - return low is None or high is None or low <= high if isinstance(capability, list): - return exact is None or exact in capability - if capability is not None: - return exact is None or exact == capability or (isinstance(exact, list) and capability in exact) - if current is None: - return False - if exact is not None and exact != current and not (isinstance(exact, list) and current in exact): - return False - if low is not None and current < low: + return exact in capability + return exact == capability or (isinstance(exact, list) and capability in exact) + + +def _satisfied_by_setting(value: dict[str, float], current: float | str | None) -> bool: + if current is None or not _matches(value.get('exact'), current): return False - return high is None or current <= high + low, high = value.get('min'), value.get('max') + return (low is None or low <= current) and (high is None or current <= high) -def _selected(value: Any, current: float, capability: Any = None) -> float: - """The value a constraint selects (exact, ideal or current), the nearest within its range and the capability""" +def _selected(value: float | dict[str, float] | None, current: float, capability: object = None) -> float: + """The value a constraint selects (exact, ideal or current), the nearest within its range and the capability.""" low, high = float('-inf'), float('inf') if isinstance(capability, (ULongRange, DoubleRange)): low = capability.min if capability.min is not None else low @@ -123,28 +138,28 @@ def _selected(value: Any, current: float, capability: Any = None) -> float: def _check_numbers(constraint_set: MediaTrackConstraints) -> None: - """The WebIDL types of the numbers of a constraint set: finite, and not negative for unsigned longs""" + """The WebIDL types of the numbers of a constraint set: finite, and not negative for unsigned longs.""" for name in _ULONG_CONSTRAINTS + _DOUBLE_CONSTRAINTS: value = getattr(constraint_set, name) members = [value.get(key) for key in ('exact', 'ideal', 'min', 'max')] if isinstance(value, dict) else [value] + unsigned = name in _ULONG_CONSTRAINTS for member in members: - if member is None: - continue - unsigned = name in _ULONG_CONSTRAINTS - if ( - isinstance(member, bool) - or not isinstance(member, (int, float)) - or not math.isfinite(member) - or (unsigned and member < 0) - ): + if member is not None and not _valid_number(member, unsigned=unsigned): kind = 'a finite number that is not negative' if unsigned else 'a finite number' - raise TypeError(f'{name} must be {kind}, not {member!r}') + msg = f'{name} must be {kind}, not {member!r}' + raise TypeError(msg) + + +def _valid_number(member: object, *, unsigned: bool) -> bool: + if isinstance(member, bool) or not isinstance(member, (int, float)): + return False + return math.isfinite(member) and not (unsigned and member < 0) def _unsatisfied( constraint_set: MediaTrackConstraints, capabilities: MediaTrackCapabilities, settings: MediaTrackSettings -) -> Optional[str]: - """The name of the first constraint of the set that can't be satisfied, if any""" +) -> str | None: + """The name of the first constraint of the set that can't be satisfied, if any.""" for name in _CONSTRAINABLE: value = getattr(constraint_set, name) if value is not None and not _satisfied(value, getattr(capabilities, name), getattr(settings, name)): @@ -152,9 +167,8 @@ def _unsatisfied( return None -class MediaStreamTrack(WebRTCObject, EventTarget): - """The MediaStreamTrack interface represents a single media track within a stream; - typically, these are audio or video tracks, but other track types may exist as well. +class MediaStreamTrack(WebRTCObject[wrtc.MediaStreamTrack], EventTarget): + """A single audio or video track of media, within a stream. Events (see :meth:`on`): ``mute`` and ``unmute`` (:obj:`webrtc.Event`): :attr:`muted` changed: a remote track is muted until media @@ -166,9 +180,9 @@ class MediaStreamTrack(WebRTCObject, EventTarget): _class = wrtc.MediaStreamTrack _events = ('mute', 'unmute', 'ended') - def _on_event(self, name: str, *args): + def _on_event(self, name: str, *args: object) -> None: # muted changes along with the events - if name in ('mute', 'unmute'): + if name in {'mute', 'unmute'}: (muted,) = args self._native_obj._surfaceMuted(muted) elif name == 'ended': @@ -176,13 +190,14 @@ def _on_event(self, name: str, *args): @property def enabled(self) -> bool: - """:obj:`bool`: A Boolean whose value of true if the track is enabled, that is allowed to render - the media source stream; or false if it is disabled, that is not rendering the media source stream but silence - and blackness. If the track has been disconnected, this value can be changed but has no more effect.""" + """:obj:`bool`: Whether the track renders its source, rather than silence or blackness. + + Once the track is disconnected, it can still be changed, to no effect. + """ return self._native_obj.enabled @enabled.setter - def enabled(self, value: bool): + def enabled(self, value: bool) -> None: self._native_obj.enabled = value @property @@ -196,41 +211,43 @@ def label(self) -> str: return self._native_obj.label @property - def kind(self) -> 'webrtc.MediaType': - """:obj:`webrtc.MediaType`: Indicating type of media. Audio or video. It doesn't change if the track is - deassociated from its source.""" + def kind(self) -> webrtc.MediaType: + """:obj:`webrtc.MediaType`: The kind of media, audio or video, even once detached from the source.""" return self._native_obj.kind @property - def ready_state(self) -> 'webrtc.MediaStreamTrackState': + def ready_state(self) -> webrtc.MediaStreamTrackState: """:obj:`webrtc.MediaStreamTrackState`: Returns an enumerated value giving the status of the track.""" return self._native_obj.readyState @property def muted(self) -> bool: - """:obj:`bool`: A value indicating whether the track - is unable to provide media data due to a technical issue.""" + """:obj:`bool`: Whether the track can't provide media, due to a technical issue.""" return self._native_obj.muted @property def content_hint(self) -> str: - """:obj:`str`: What the track carries, which encoders optimize for: ``'speech'``, ``'speaking'`` or - ``'music'`` for audio, ``'motion'``, ``'detail'`` or ``'text'`` for video, empty if unknown (the default). - Other values, and the ones of the other kind, are ignored.""" + """:obj:`str`: What the track carries, which encoders optimize for, empty if unknown (the default). + + ``'speech'``, ``'speaking'`` or ``'music'`` for audio, ``'motion'``, ``'detail'`` or ``'text'`` for video. + Other values, and the ones of the other kind, are ignored. + """ return self._native_obj.contentHint @content_hint.setter - def content_hint(self, value: str): + def content_hint(self, value: str) -> None: self._native_obj.contentHint = str(value) def get_settings(self) -> MediaTrackSettings: - """Returns what the track carries: the size and frame rate of the frames last seen, the format of the audio, - and the device of the tracks of :func:`webrtc.get_user_media`. + """Returns what the track carries. + + The size and frame rate of the frames last seen, the format of the audio, and the device of the tracks of + :func:`webrtc.get_user_media`. Returns: :obj:`webrtc.MediaTrackSettings`: The settings. """ - native: Dict[str, Any] = self._native_obj._settings() + native: dict[str, Any] = self._native_obj._settings() settings = MediaTrackSettings() if 'width' in native: settings.width, settings.height = native['width'], native['height'] @@ -250,8 +267,10 @@ def get_settings(self) -> MediaTrackSettings: return settings def get_capabilities(self) -> MediaTrackCapabilities: - """Returns what the source of the track can do: the synthetic camera and microphone of - :func:`webrtc.get_user_media` have capabilities, other tracks (remote, generated) have none. + """Returns what the source of the track can do. + + The synthetic camera and microphone of :func:`webrtc.get_user_media` have capabilities, other tracks (remote, + generated) have none. Returns: :obj:`webrtc.MediaTrackCapabilities`: The capabilities. @@ -272,9 +291,13 @@ def get_constraints(self) -> MediaTrackConstraints: constraints = self._native_obj._constraints return constraints if constraints is not None else MediaTrackConstraints() - def apply_constraints(self, constraints: Optional[Any] = None) -> asyncio.Future: - """Applies constraints to the track: the synthetic camera of :func:`webrtc.get_user_media` changes its size - and frame rate, the source of other tracks stays as it is. + def apply_constraints( + self, constraints: MediaTrackConstraints | dict[str, Any] | None = None + ) -> asyncio.Future[None]: + """Applies constraints to the track. + + The synthetic camera of :func:`webrtc.get_user_media` changes its size and frame rate, the source of other + tracks stays as it is. Args: constraints (:obj:`webrtc.MediaTrackConstraints` or :obj:`dict`, optional): The constraints, none to @@ -317,14 +340,13 @@ def _apply_constraints(self, constraints: MediaTrackConstraints) -> None: self._native_obj._reconfigureCamera(int(width), int(height), float(frame_rate)) self._native_obj._constraints = constraints - def clone(self) -> 'webrtc.MediaStreamTrack': + def clone(self) -> webrtc.MediaStreamTrack: """Returns a duplicate of the :obj:`webrtc.MediaStreamTrack`.""" return self._wrap(self._native_obj.clone()) - def stop(self): - """Stops playing the source associated to the track, both the source and the track are deassociated. - The track state is set to ended.""" - return self._native_obj.stop() + def stop(self) -> None: + """Stops the track, detached from its source: its :attr:`ready_state` becomes ended.""" + self._native_obj.stop() #: Alias for :attr:`ready_state` readyState = ready_state diff --git a/python-webrtc/python/webrtc/interfaces/media_stream_track_processor.py b/python-webrtc/python/webrtc/interfaces/media_stream_track_processor.py index d32905d..7a93484 100644 --- a/python-webrtc/python/webrtc/interfaces/media_stream_track_processor.py +++ b/python-webrtc/python/webrtc/interfaces/media_stream_track_processor.py @@ -5,18 +5,26 @@ # that can be found in the LICENSE.md file in the root of the project. # +"""MediaStreamTrackProcessor of Chrome, the media of a track as a stream.""" + +from __future__ import annotations + from dataclasses import dataclass -from typing import Any, Dict, Optional, Tuple, Union +from typing import TYPE_CHECKING, Any, ClassVar from webrtc import AudioData, MediaStreamTrack, MediaType, VideoFrame, WebRTCObject, wrtc from webrtc.streams import ReadableStream from webrtc.utils.events import EventTarget -from webrtc.utils.names import alias +from webrtc.utils.names import Alias, alias + +if TYPE_CHECKING: + from webrtc.streams import ReadableStreamDefaultController #: How many frames of video are queued for reads, as the specification says DEFAULT_VIDEO_BUFFER_SIZE = 1 #: How many 10 ms chunks of audio are queued for reads, as Chrome does DEFAULT_AUDIO_BUFFER_SIZE = 10 +_MAX_BUFFER_SIZE = 65535 @dataclass @@ -29,16 +37,16 @@ class MediaStreamTrackProcessorInit: """ track: MediaStreamTrack - max_buffer_size: Optional[int] = None + max_buffer_size: int | None = None #: Alias for :attr:`max_buffer_size` - maxBufferSize = alias('max_buffer_size') + maxBufferSize: ClassVar[Alias[int | None]] = alias('max_buffer_size') def _parse_init( - track: Any, max_buffer_size: Optional[int], options: Dict[str, Any] -) -> Tuple[MediaStreamTrack, Optional[int]]: - """The track and the buffer size, from the init, its dict form, or the arguments""" + track: object, max_buffer_size: int | None, options: dict[str, object] +) -> tuple[MediaStreamTrack, int | None]: + """The track and the buffer size, from the init, its dict form, or the arguments.""" if isinstance(track, MediaStreamTrackProcessorInit): track, max_buffer_size = track.track, track.max_buffer_size elif isinstance(track, dict): @@ -46,28 +54,33 @@ def _parse_init( track = init.pop('track', None) max_buffer_size = init.pop('max_buffer_size', init.pop('maxBufferSize', max_buffer_size)) if init: - raise TypeError(f'MediaStreamTrackProcessorInit has no member {next(iter(init))!r}') + msg = f'MediaStreamTrackProcessorInit has no member {next(iter(init))!r}' + raise TypeError(msg) if 'maxBufferSize' in options: max_buffer_size = options.pop('maxBufferSize') if options: - raise TypeError(f'Unexpected arguments: {", ".join(options)}') + msg = f'Unexpected arguments: {", ".join(options)}' + raise TypeError(msg) if not isinstance(track, MediaStreamTrack): - raise TypeError(f'track must be a MediaStreamTrack, not {type(track).__name__}') + msg = f'track must be a MediaStreamTrack, not {type(track).__name__}' + raise TypeError(msg) return track, max_buffer_size class _TrackSource: - """The underlying source of :attr:`MediaStreamTrackProcessor.readable`: takes media from the native queue only - for pending reads, so the native queue drops the oldest items when full""" + """The underlying source of :attr:`MediaStreamTrackProcessor.readable`. - def __init__(self, processor: 'MediaStreamTrackProcessor'): + It takes media from the native queue only for pending reads, so the native queue drops the oldest items when full. + """ + + def __init__(self, processor: MediaStreamTrackProcessor) -> None: self._processor = processor - self._controller = None + self._controller: ReadableStreamDefaultController | None = None - def start(self, controller): + def start(self, controller: ReadableStreamDefaultController) -> None: self._controller = controller - def pull(self, controller): + def pull(self, _controller: ReadableStreamDefaultController) -> None: native = self._processor._native_obj created_outside_loop = native._listeners is None # the native events go to the loop reading @@ -77,11 +90,11 @@ def pull(self, controller): native._ackWakeup() self.deliver() - def cancel(self, reason): + def cancel(self, _reason: object) -> None: self._processor._native_obj.cancel() - def deliver(self): - """Fulfills the pending reads with the media queued, and closes the stream once the track ended""" + def deliver(self) -> None: + """Fulfills the pending reads with the media queued, and closes the stream once the track ended.""" native = self._processor._native_obj stream = self._processor._readable controller = self._controller @@ -94,10 +107,11 @@ def deliver(self): controller.close() -class MediaStreamTrackProcessor(WebRTCObject, EventTarget): - """Reads the media of a track as a stream: :obj:`webrtc.VideoFrame` objects for a video track, - :obj:`webrtc.AudioData` ones (10 ms each) for an audio track, as Chrome does - (https://developer.mozilla.org/en-US/docs/Web/API/MediaStreamTrackProcessor). +class MediaStreamTrackProcessor(WebRTCObject[wrtc.MediaStreamTrackProcessor], EventTarget): + """Reads the media of a track as a stream, as Chrome does. + + See https://developer.mozilla.org/en-US/docs/Web/API/MediaStreamTrackProcessor. It reads + :obj:`webrtc.VideoFrame` objects for a video track, :obj:`webrtc.AudioData` ones (10 ms each) for an audio track. Media is queued as it arrives, up to ``max_buffer_size`` items: when the queue is full, the oldest item is dropped (and counted in :attr:`discarded_frames`), so a slow reader never makes memory grow. The stream closes @@ -110,7 +124,7 @@ class MediaStreamTrackProcessor(WebRTCObject, EventTarget): of audio. Raises: - :obj:`TypeError`: If the track isn't a :obj:`webrtc.MediaStreamTrack`, or the size isn't from 0 to 65535. + TypeError: If the track isn't a :obj:`webrtc.MediaStreamTrack`, or the size isn't from 0 to 65535. Example:: @@ -124,18 +138,20 @@ class MediaStreamTrackProcessor(WebRTCObject, EventTarget): def __init__( self, - track: Union[MediaStreamTrack, MediaStreamTrackProcessorInit, Dict[str, Any], None] = None, - max_buffer_size: Optional[int] = None, - **options, - ): + track: MediaStreamTrack | MediaStreamTrackProcessorInit | dict[str, Any] | None = None, + max_buffer_size: int | None = None, + **options: object, + ) -> None: track, max_buffer_size = _parse_init(track, max_buffer_size, options) video = track.kind == MediaType.video if max_buffer_size is None: max_buffer_size = DEFAULT_VIDEO_BUFFER_SIZE if video else DEFAULT_AUDIO_BUFFER_SIZE if isinstance(max_buffer_size, bool) or not isinstance(max_buffer_size, int): - raise TypeError(f'max_buffer_size must be an integer, not {max_buffer_size!r}') - if not 0 <= max_buffer_size <= 65535: - raise TypeError(f'max_buffer_size must be from 0 to 65535, not {max_buffer_size}') + msg = f'max_buffer_size must be an integer, not {max_buffer_size!r}' + raise TypeError(msg) + if not 0 <= max_buffer_size <= _MAX_BUFFER_SIZE: + msg = f'max_buffer_size must be from 0 to {_MAX_BUFFER_SIZE}, not {max_buffer_size}' + raise TypeError(msg) super().__init__(self._class(track._native_obj, max(1, max_buffer_size))) # the native processor doesn't keep the track, Python does @@ -145,17 +161,15 @@ def __init__( self._readable = ReadableStream(self._source, high_water_mark=0) self._attach() - def _on_event(self, name: str, *args): + def _on_event(self, name: str, *_args: object) -> None: if name == '_ready': self._native_obj._ackWakeup() self._source.deliver() - def _wrap_media(self, item: tuple) -> Union[VideoFrame, AudioData]: + def _wrap_media(self, item: tuple[Any, ...]) -> VideoFrame | AudioData: if self._video: - buffer, timestamp, rotation, rtp_timestamp = item - return VideoFrame._from_native(buffer, timestamp, rotation, rtp_timestamp or None) - data, bits_per_sample, sample_rate, channels, frames, timestamp = item - return AudioData._from_native(data, bits_per_sample, sample_rate, channels, frames, timestamp) + return VideoFrame._from_native(item) + return AudioData._from_native(item) @property def readable(self) -> ReadableStream: diff --git a/python-webrtc/python/webrtc/interfaces/rtc_data_channel.py b/python-webrtc/python/webrtc/interfaces/rtc_data_channel.py index 3953bbb..1350ca9 100644 --- a/python-webrtc/python/webrtc/interfaces/rtc_data_channel.py +++ b/python-webrtc/python/webrtc/interfaces/rtc_data_channel.py @@ -5,18 +5,107 @@ # that can be found in the LICENSE.md file in the root of the project. # -from typing import TYPE_CHECKING, Optional, Union +"""RTCDataChannel of WebRTC.""" -from webrtc import BinaryType, Blob, MessageEvent, RTCDataChannelState, RTCErrorEvent, WebRTCObject, wrtc +from __future__ import annotations + +from dataclasses import dataclass +from typing import TYPE_CHECKING, ClassVar + +from webrtc import ( + BinaryType, + Blob, + MessageEvent, + RTCDataChannelState, + RTCErrorEvent, + RTCPriorityType, + WebRTCObject, + wrtc, +) from webrtc.utils.events import EventTarget +from webrtc.utils.names import Alias, alias if TYPE_CHECKING: import webrtc +#: The maximum of an unsigned short, which limits the members of the init +MAX_UNSIGNED_SHORT = 65535 + + +def check_utf8_length(name: str, value: str) -> None: + """Checks a string of the init fits in an unsigned short number of bytes, as the specification requires. + + Raises: + ValueError: If it doesn't. + """ + if len(value.encode()) > MAX_UNSIGNED_SHORT: + msg = f'{name} is longer than {MAX_UNSIGNED_SHORT} bytes' + raise ValueError(msg) + + +@dataclass +class RTCDataChannelInit: + """How :meth:`webrtc.RTCPeerConnection.create_data_channel` creates a channel. + + Args: + ordered (:obj:`bool`, optional): Whether messages are delivered in order. + max_packet_life_time (:obj:`int`, optional): Makes the channel unreliable: how long in milliseconds + a message is retransmitted. + max_retransmits (:obj:`int`, optional): Makes the channel unreliable: how many times a message is + retransmitted. Can't be set along with ``max_packet_life_time``. + protocol (:obj:`str`, optional): The subprotocol name, up to 65535 bytes in UTF-8. + negotiated (:obj:`bool`, optional): Whether the application creates the channel on both ends + with the same ``id``, instead of announcing it to the remote peer. + id (:obj:`int`, optional): The SCTP stream id (0 to 65534) of a ``negotiated`` channel, which requires it. + Ignored otherwise, as the connection picks the id. + priority (:obj:`webrtc.RTCPriorityType`, optional): The priority of the channel. + """ + + ordered: bool = True + max_packet_life_time: int | None = None + max_retransmits: int | None = None + protocol: str = '' + negotiated: bool = False + id: int | None = None + priority: RTCPriorityType | str = RTCPriorityType.low -class RTCDataChannel(WebRTCObject, EventTarget): - """A bidirectional channel of messages between the peers, created with - :meth:`webrtc.RTCPeerConnection.create_data_channel` or received with its ``datachannel`` event. + def _check(self) -> None: + """Checks the members, as the specification requires. + + Raises: + ValueError: If a member is out of range, or both ``max_packet_life_time`` and ``max_retransmits`` are + set, or ``negotiated`` is set without ``id``. + """ + check_utf8_length('protocol', self.protocol) + for name in ('max_packet_life_time', 'max_retransmits'): + value = getattr(self, name) + if value is not None and not 0 <= value <= MAX_UNSIGNED_SHORT: + msg = f'{name} must be from 0 to {MAX_UNSIGNED_SHORT}, not {value}' + raise ValueError(msg) + if self.max_packet_life_time is not None and self.max_retransmits is not None: + msg = 'max_packet_life_time and max_retransmits can not both be set' + raise ValueError(msg) + if not self.negotiated: + return + if self.id is None: + msg = 'a negotiated channel needs an id' + raise ValueError(msg) + # the last stream id is reserved + if not 0 <= self.id < MAX_UNSIGNED_SHORT: + msg = f'id must be from 0 to {MAX_UNSIGNED_SHORT - 1}, not {self.id}' + raise ValueError(msg) + + #: Alias for :attr:`max_packet_life_time` + maxPacketLifeTime: ClassVar[Alias[int | None]] = alias('max_packet_life_time') + #: Alias for :attr:`max_retransmits` + maxRetransmits: ClassVar[Alias[int | None]] = alias('max_retransmits') + + +class RTCDataChannel(WebRTCObject[wrtc.RTCDataChannel], EventTarget): + """A bidirectional channel of messages between the peers. + + It's created with :meth:`webrtc.RTCPeerConnection.create_data_channel` or received with its ``datachannel`` + event. Events (see :meth:`on`): ``open`` (:obj:`webrtc.Event`): The channel can be used to send messages. @@ -32,9 +121,9 @@ class RTCDataChannel(WebRTCObject, EventTarget): _class = wrtc.RTCDataChannel _events = ('open', 'message', 'bufferedamountlow', 'error', 'closing', 'close') - def _on_event(self, name: str, *args): + def _on_event(self, name: str, *args: object) -> None: # readyState changes along with the events - if name in ('open', 'closing', 'close'): + if name in {'open', 'closing', 'close'}: (state,) = args self._native_obj._surfaceState(state) elif name == '_sent': @@ -43,7 +132,7 @@ def _on_event(self, name: str, *args): # in the same task as the decrease, before anything that arrived meanwhile self._dispatch('bufferedamountlow') - def _create_event(self, name: str, *args): + def _create_event(self, name: str, *args: object) -> webrtc.Event | None: if name == 'open' and self.ready_state != RTCDataChannelState.open: # closed before it opened return None @@ -70,12 +159,12 @@ def ordered(self) -> bool: return self._native_obj.ordered @property - def max_packet_life_time(self) -> Optional[int]: + def max_packet_life_time(self) -> int | None: """:obj:`int`, optional: How long in milliseconds a message is retransmitted in unreliable mode.""" return self._native_obj.maxPacketLifeTime @property - def max_retransmits(self) -> Optional[int]: + def max_retransmits(self) -> int | None: """:obj:`int`, optional: How many times a message is retransmitted in unreliable mode.""" return self._native_obj.maxRetransmits @@ -90,17 +179,17 @@ def negotiated(self) -> bool: return self._native_obj.negotiated @property - def id(self) -> Optional[int]: + def id(self) -> int | None: """:obj:`int`, optional: The SCTP stream id of the channel, :obj:`None` until it's known.""" return self._native_obj.id @property - def priority(self) -> 'webrtc.RTCPriorityType': + def priority(self) -> webrtc.RTCPriorityType: """:obj:`webrtc.RTCPriorityType`: The priority of the channel.""" return self._native_obj.priority @property - def ready_state(self) -> 'webrtc.RTCDataChannelState': + def ready_state(self) -> webrtc.RTCDataChannelState: """:obj:`webrtc.RTCDataChannelState`: The state of the channel.""" return self._native_obj.readyState @@ -115,38 +204,42 @@ def buffered_amount_low_threshold(self) -> int: return self._native_obj.bufferedAmountLowThreshold @buffered_amount_low_threshold.setter - def buffered_amount_low_threshold(self, value: int): + def buffered_amount_low_threshold(self, value: int) -> None: if not 0 <= value < 2**64: - raise ValueError(f'buffered_amount_low_threshold must be from 0 to 2**64-1, not {value}') + msg = f'buffered_amount_low_threshold must be from 0 to 2**64-1, not {value}' + raise ValueError(msg) self._native_obj.bufferedAmountLowThreshold = value @property def binary_type(self) -> BinaryType: - """:obj:`webrtc.BinaryType`: What binary messages are delivered as: :obj:`bytes` (``arraybuffer``, the - default) or :obj:`webrtc.Blob` (``blob``).""" + """:obj:`webrtc.BinaryType`: What binary messages are delivered as, :obj:`bytes` by default. + + :obj:`bytes` for ``arraybuffer``, :obj:`webrtc.Blob` for ``blob``. + """ return BinaryType(self._native_obj.binaryType) @binary_type.setter - def binary_type(self, value: Union[BinaryType, str]): + def binary_type(self, value: BinaryType | str) -> None: self._native_obj.binaryType = BinaryType(value).value - def send(self, data: Union[str, bytes, bytearray, memoryview, Blob]) -> None: + def send(self, data: str | bytes | bytearray | memoryview | Blob) -> None: """Sends a message to the remote peer. Args: data (:obj:`str`, bytes-like or :obj:`webrtc.Blob`): A text message, or a binary one. Raises: - :obj:`TypeError`: If the data is neither text, bytes nor a :obj:`webrtc.Blob`. - :obj:`webrtc.InvalidStateError`: If the channel isn't open. - :obj:`webrtc.OperationError`: If the message can't be queued, like when the queue is full. + TypeError: If the data is neither text, bytes nor a :obj:`webrtc.Blob`. + webrtc.InvalidStateError: If the channel isn't open. + webrtc.OperationError: If the message can't be queued, like when the queue is full. """ if isinstance(data, str): - self._native_obj.send(data.encode(), False) + self._native_obj.send(data.encode(), binary=False) elif isinstance(data, (bytes, bytearray, memoryview, Blob)): - self._native_obj.send(bytes(data), True) + self._native_obj.send(bytes(data), binary=True) else: - raise TypeError(f'data must be str, bytes-like or Blob, not {type(data).__name__}') + msg = f'data must be str, bytes-like or Blob, not {type(data).__name__}' + raise TypeError(msg) def close(self) -> None: """Closes the channel. Messages queued before are still sent.""" diff --git a/python-webrtc/python/webrtc/interfaces/rtc_dtls_transport.py b/python-webrtc/python/webrtc/interfaces/rtc_dtls_transport.py index d612e61..62e8b8a 100644 --- a/python-webrtc/python/webrtc/interfaces/rtc_dtls_transport.py +++ b/python-webrtc/python/webrtc/interfaces/rtc_dtls_transport.py @@ -5,19 +5,20 @@ # that can be found in the LICENSE.md file in the root of the project. # -from typing import TYPE_CHECKING, List +"""RTCDtlsTransport of WebRTC.""" +from __future__ import annotations + +import webrtc from webrtc import RTCErrorEvent, WebRTCObject, wrtc from webrtc.utils.events import EventTarget -if TYPE_CHECKING: - import webrtc +class RTCDtlsTransport(WebRTCObject[wrtc.RTCDtlsTransport], EventTarget): + """The Datagram Transport Layer Security (DTLS) transport of a :obj:`webrtc.RTCPeerConnection`. -class RTCDtlsTransport(WebRTCObject, EventTarget): - """The :obj:`webrtc.RTCDtlsTransport` interface provides access to information about the Datagram Transport - Layer Security (DTLS) transport over which a :obj:`webrtc.RTCPeerConnection`'s RTP and RTCP packets are sent and - received by its :obj:`webrtc.RTCRtpSender` and :obj:`webrtc.RTCRtpReceiver` objects. + The RTP and RTCP packets of its :obj:`webrtc.RTCRtpSender` and :obj:`webrtc.RTCRtpReceiver` objects are sent and + received over it. Events (see :meth:`on`): ``statechange`` (:obj:`webrtc.Event`): :attr:`state` changed. @@ -27,32 +28,29 @@ class RTCDtlsTransport(WebRTCObject, EventTarget): _class = wrtc.RTCDtlsTransport _events = ('statechange', 'error') - def _on_event(self, name: str, *args): + def _on_event(self, name: str, *args: object) -> None: # the state changes along with its event if name == 'statechange': (state,) = args self._native_obj._surfaceState(state) - def _create_event(self, name: str, *args): + def _create_event(self, name: str, *args: object) -> webrtc.Event | None: if name == 'error': (error,) = args return RTCErrorEvent(name, error.toPython(), target=self) return super()._create_event(name, *args) @property - def ice_transport(self) -> 'webrtc.RTCIceTransport': + def ice_transport(self) -> webrtc.RTCIceTransport: """:obj:`webrtc.RTCIceTransport`: Returns a reference to the underlying :obj:`webrtc.RTCIceTransport` object.""" - from webrtc import RTCIceTransport - - return RTCIceTransport._wrap(self._native_obj.iceTransport) + return webrtc.RTCIceTransport._wrap(self._native_obj.iceTransport) @property - def state(self) -> 'webrtc.DtlsTransportState': - """:obj:`webrtc.DtlsTransportState`: Returns a member of :obj:`webrtc.DtlsTransportState` which describes the - underlying Datagram Transport Layer Security (DTLS) transport state.""" + def state(self) -> webrtc.DtlsTransportState: + """:obj:`webrtc.DtlsTransportState`: The state of the DTLS transport.""" return self._native_obj.state - def get_remote_certificates(self) -> List[bytes]: + def get_remote_certificates(self) -> list[bytes]: """Returns the certificates of the remote peer, once the DTLS handshake is done. Returns: diff --git a/python-webrtc/python/webrtc/interfaces/rtc_dtmf_sender.py b/python-webrtc/python/webrtc/interfaces/rtc_dtmf_sender.py index e86b224..ff294da 100644 --- a/python-webrtc/python/webrtc/interfaces/rtc_dtmf_sender.py +++ b/python-webrtc/python/webrtc/interfaces/rtc_dtmf_sender.py @@ -5,15 +5,23 @@ # that can be found in the LICENSE.md file in the root of the project. # +"""RTCDTMFSender of WebRTC.""" + +from __future__ import annotations + import re +from typing import TYPE_CHECKING from webrtc import InvalidCharacterError, RTCDTMFToneChangeEvent, WebRTCObject, wrtc from webrtc.utils.events import EventTarget +if TYPE_CHECKING: + import webrtc + _TONES = re.compile(r'[0-9A-Da-d#*,]*') -class RTCDTMFSender(WebRTCObject, EventTarget): +class RTCDTMFSender(WebRTCObject[wrtc.RTCDTMFSender], EventTarget): """Sends DTMF tones on an audio sender (:attr:`webrtc.RTCRtpSender.dtmf`). Events (see :meth:`on`): @@ -24,12 +32,12 @@ class RTCDTMFSender(WebRTCObject, EventTarget): _class = wrtc.RTCDTMFSender _events = ('tonechange',) - def _on_event(self, name: str, *args): + def _on_event(self, _name: str, *args: object) -> None: _, tone_buffer, insertion = args # the tone buffer is shortened along with the event self._native_obj._surfaceBuffer(tone_buffer, insertion) - def _create_event(self, name: str, *args): + def _create_event(self, name: str, *args: object) -> webrtc.Event | None: tone, _, _ = args return RTCDTMFToneChangeEvent(name, tone, target=self) @@ -43,11 +51,12 @@ def insert_dtmf(self, tones: str, duration: int = 100, inter_tone_gap: int = 70) inter_tone_gap (:obj:`int`, optional): The pause between tones in milliseconds, at least 30. Raises: - :obj:`webrtc.InvalidCharacterError`: If ``tones`` has another character. - :obj:`webrtc.InvalidStateError`: If the transceiver of the sender is stopped or doesn't send. + webrtc.InvalidCharacterError: If ``tones`` has another character. + webrtc.InvalidStateError: If the transceiver of the sender is stopped or doesn't send. """ if not _TONES.fullmatch(tones): - raise InvalidCharacterError(f'{tones!r} has characters that are not DTMF tones') + msg = f'{tones!r} has characters that are not DTMF tones' + raise InvalidCharacterError(msg) duration = min(max(int(duration), 40), 6000) inter_tone_gap = min(max(int(inter_tone_gap), 30), 6000) self._native_obj.insertDTMF(tones.upper(), duration, inter_tone_gap) diff --git a/python-webrtc/python/webrtc/interfaces/rtc_ice_transport.py b/python-webrtc/python/webrtc/interfaces/rtc_ice_transport.py index 385c762..bf37c94 100644 --- a/python-webrtc/python/webrtc/interfaces/rtc_ice_transport.py +++ b/python-webrtc/python/webrtc/interfaces/rtc_ice_transport.py @@ -5,9 +5,13 @@ # that can be found in the LICENSE.md file in the root of the project. # +"""RTCIceTransport of WebRTC, standalone too, as in WebRTC Extensions.""" + +from __future__ import annotations + import re import weakref -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Sequence, Union +from typing import TYPE_CHECKING, Any from webrtc import ( CricketIceGatheringState, @@ -26,6 +30,8 @@ from webrtc.utils.events import EventTarget if TYPE_CHECKING: + from collections.abc import Sequence + import webrtc @@ -34,12 +40,12 @@ _PASSWORD = re.compile(r'[A-Za-z0-9+/]{22,256}') # the candidates of every transport, by their candidate-attribute: a candidate is the same object each time -_candidates: 'weakref.WeakKeyDictionary[wrtc.RTCIceTransport, Dict[str, webrtc.RTCIceCandidate]]' = ( +_candidates: weakref.WeakKeyDictionary[wrtc.RTCIceTransport, dict[str, webrtc.RTCIceCandidate]] = ( weakref.WeakKeyDictionary() ) -class RTCIceTransport(WebRTCObject, EventTarget): +class RTCIceTransport(WebRTCObject[wrtc.RTCIceTransport], EventTarget): """The ICE transport a :obj:`webrtc.RTCDtlsTransport` of a connection runs over, or a standalone transport. A standalone one (``RTCIceTransport()``) connects after :meth:`gather`, :meth:`start` and @@ -56,12 +62,12 @@ class RTCIceTransport(WebRTCObject, EventTarget): _class = wrtc.RTCIceTransport _events = ('statechange', 'gatheringstatechange', 'selectedcandidatepairchange', 'icecandidate') - def __init__(self): + def __init__(self) -> None: super().__init__() # a standalone transport delivers its events to the loop it's created on self._attach() - def _candidate_of(self, native: 'wrtc.IceCandidateInit') -> 'webrtc.RTCIceCandidate': + def _candidate_of(self, native: wrtc.IceCandidateInit) -> webrtc.RTCIceCandidate: """The candidate object of a native candidate, the same one each time.""" kwargs = native.kwargs() known = _candidates.setdefault(self._native_obj, {}) @@ -69,11 +75,11 @@ def _candidate_of(self, native: 'wrtc.IceCandidateInit') -> 'webrtc.RTCIceCandid known[kwargs['candidate']] = RTCIceCandidate(**kwargs) return known[kwargs['candidate']] - def _remember(self, candidate: 'webrtc.RTCIceCandidate') -> None: + def _remember(self, candidate: webrtc.RTCIceCandidate) -> None: """Makes a candidate the object of its native candidate, unless there's one already.""" _candidates.setdefault(self._native_obj, {}).setdefault(candidate.candidate, candidate) - def _on_event(self, name: str, *args): + def _on_event(self, name: str, *args: object) -> None: # the states change along with their events if name == 'statechange': (state,) = args @@ -84,7 +90,7 @@ def _on_event(self, name: str, *args): elif name == 'icecandidate' and args and args[0] is not None: self._native_obj._surfaceCandidate() - def _create_event(self, name: str, *args): + def _create_event(self, name: str, *args: object) -> webrtc.Event | None: if name == 'icecandidate': candidate = self._candidate_of(args[0]) if args and args[0] is not None else None return RTCPeerConnectionIceEvent(name, candidate, None, target=self) @@ -92,17 +98,19 @@ def _create_event(self, name: str, *args): def _check_standalone(self, operation: str) -> None: if not self._native_obj._standalone: - raise InvalidStateError(f'Can not {operation}: the transport belongs to an RTCPeerConnection') + msg = f'Can not {operation}: the transport belongs to an RTCPeerConnection' + raise InvalidStateError(msg) def _check_open(self, operation: str) -> None: self._check_standalone(operation) if self.state == RTCIceTransportState.closed: - raise InvalidStateError(f'Can not {operation}: the transport is stopped') + msg = f'Can not {operation}: the transport is stopped' + raise InvalidStateError(msg) def gather( self, - gather_policy: Union['webrtc.RTCIceTransportPolicy', str] = 'all', - ice_servers: Optional[Sequence[Union['webrtc.RTCIceServer', Dict[str, Any]]]] = None, + gather_policy: webrtc.RTCIceTransportPolicy | str = 'all', + ice_servers: Sequence[webrtc.RTCIceServer | dict[str, Any]] | None = None, ) -> None: """Gathers the candidates of a standalone transport, sent in ``icecandidate`` events. @@ -112,22 +120,24 @@ def gather( A :obj:`dict` of the arguments of :obj:`webrtc.RTCIceServer` is accepted too. Raises: - :obj:`webrtc.InvalidStateError`: If it's stopped, gathering already, or belongs to a connection. - :obj:`webrtc.InvalidSyntaxError`: If an ICE server URL is invalid. - :obj:`webrtc.InvalidAccessError`: If a TURN server has no credentials. - :obj:`TypeError`: If the policy isn't a value of :obj:`webrtc.RTCIceTransportPolicy`. + webrtc.InvalidStateError: If it's stopped, gathering already, or belongs to a connection. + webrtc.InvalidSyntaxError: If an ICE server URL is invalid. + webrtc.InvalidAccessError: If a TURN server has no credentials. + TypeError: If the policy isn't a value of :obj:`webrtc.RTCIceTransportPolicy`. """ self._check_open('gather') if self.gathering_state != CricketIceGatheringState.new: - raise InvalidStateError('The transport gathers its candidates already') + msg = 'The transport gathers its candidates already' + raise InvalidStateError(msg) self._native_obj.gather(gather_policy, RTCIceServer._to_native_list(ice_servers or ())) def start( self, - remote_parameters: Union['webrtc.RTCIceParameters', Dict[str, str]], - role: Union['webrtc.RTCIceRole', str] = 'controlled', + remote_parameters: webrtc.RTCIceParameters | dict[str, str], + role: webrtc.RTCIceRole | str = 'controlled', ) -> None: """Starts connecting a standalone transport to the remote agent, with the candidates added, or later. + Remote parameters that differ from the ones given before remove the remote candidates. Args: @@ -137,22 +147,25 @@ def start( the same role, one of them switches. Raises: - :obj:`webrtc.InvalidStateError`: If it's stopped, started with another role, or belongs to a connection. - :obj:`webrtc.InvalidSyntaxError`: If the username fragment or the password is invalid. - :obj:`ValueError`: If the role is neither controlling nor controlled. + webrtc.InvalidStateError: If it's stopped, started with another role, or belongs to a connection. + webrtc.InvalidSyntaxError: If the username fragment or the password is invalid. + ValueError: If the role is neither controlling nor controlled. """ self._check_open('start') if isinstance(remote_parameters, dict): remote_parameters = RTCIceParameters(**remote_parameters) if not _UFRAG.fullmatch(remote_parameters.username_fragment): - raise InvalidSyntaxError(f'{remote_parameters.username_fragment!r} is not a valid ICE username fragment') + msg = f'{remote_parameters.username_fragment!r} is not a valid ICE username fragment' + raise InvalidSyntaxError(msg) if not _PASSWORD.fullmatch(remote_parameters.password): - raise InvalidSyntaxError('the ICE password is not valid') - if role not in (RTCIceRole.controlling, RTCIceRole.controlled): - raise ValueError('role must be controlling or controlled') + msg = 'the ICE password is not valid' + raise InvalidSyntaxError(msg) + if role not in {RTCIceRole.controlling, RTCIceRole.controlled}: + msg = 'role must be controlling or controlled' + raise ValueError(msg) self._native_obj.start(remote_parameters.username_fragment, remote_parameters.password, role) - def add_remote_candidate(self, candidate: Union['webrtc.RTCIceCandidate', Dict[str, Any]]) -> None: + def add_remote_candidate(self, candidate: webrtc.RTCIceCandidate | dict[str, Any]) -> None: """Adds a candidate of the remote agent to a standalone transport. Args: @@ -160,9 +173,9 @@ def add_remote_candidate(self, candidate: Union['webrtc.RTCIceCandidate', Dict[s (see :meth:`webrtc.RTCIceCandidate.to_json`). Raises: - :obj:`TypeError`: If the candidate has neither ``sdp_mid`` nor ``sdp_m_line_index``. - :obj:`webrtc.InvalidStateError`: If it's stopped, or belongs to a connection. - :obj:`webrtc.OperationError`: If the candidate can't be parsed. + TypeError: If the candidate has neither ``sdp_mid`` nor ``sdp_m_line_index``. + webrtc.InvalidStateError: If it's stopped, or belongs to a connection. + webrtc.OperationError: If the candidate can't be parsed. """ self._check_open('add a remote candidate') if not isinstance(candidate, RTCIceCandidate): @@ -176,12 +189,12 @@ def stop(self) -> None: """Stops a standalone transport: it's closed, without a ``statechange`` event. Raises: - :obj:`webrtc.InvalidStateError`: If it belongs to a connection. + webrtc.InvalidStateError: If it belongs to a connection. """ self._check_standalone('stop') self._native_obj.stop() - def get_selected_candidate_pair(self) -> Optional['webrtc.RTCIceCandidatePair']: + def get_selected_candidate_pair(self) -> webrtc.RTCIceCandidatePair | None: """Returns the local and the remote candidate the transport sends and receives with. Returns: @@ -197,7 +210,7 @@ def get_selected_candidate_pair(self) -> Optional['webrtc.RTCIceCandidatePair']: remote_candidate = RTCIceCandidate(**remote.kwargs()) return RTCIceCandidatePair(RTCIceCandidate(**local.kwargs()), remote_candidate) - def get_local_candidates(self) -> List['webrtc.RTCIceCandidate']: + def get_local_candidates(self) -> list[webrtc.RTCIceCandidate]: """Returns the candidates gathered for the transport, sent in ``icecandidate`` events of the connection. Returns: @@ -205,16 +218,17 @@ def get_local_candidates(self) -> List['webrtc.RTCIceCandidate']: """ return [self._candidate_of(c) for c in self._native_obj.getLocalCandidates()] - def get_remote_candidates(self) -> List['webrtc.RTCIceCandidate']: - """Returns the candidates the remote peer signaled for the transport, in its description or with - :meth:`webrtc.RTCPeerConnection.add_ice_candidate`. Peer-reflexive ones aren't. + def get_remote_candidates(self) -> list[webrtc.RTCIceCandidate]: + """Returns the candidates the remote peer signaled for the transport, but not peer-reflexive ones. + + They're signaled in its description or with :meth:`webrtc.RTCPeerConnection.add_ice_candidate`. Returns: :obj:`list` of :obj:`webrtc.RTCIceCandidate`: The candidates. """ return [self._candidate_of(c) for c in self._native_obj.getRemoteCandidates()] - def get_local_parameters(self) -> Optional['webrtc.RTCIceParameters']: + def get_local_parameters(self) -> webrtc.RTCIceParameters | None: """Returns the ICE parameters of the transport in the local description. Returns: @@ -223,7 +237,7 @@ def get_local_parameters(self) -> Optional['webrtc.RTCIceParameters']: parameters = self._native_obj.getLocalParameters() return RTCIceParameters(*parameters) if parameters is not None else None - def get_remote_parameters(self) -> Optional['webrtc.RTCIceParameters']: + def get_remote_parameters(self) -> webrtc.RTCIceParameters | None: """Returns the ICE parameters of the transport in the remote description. Returns: @@ -233,26 +247,28 @@ def get_remote_parameters(self) -> Optional['webrtc.RTCIceParameters']: return RTCIceParameters(*parameters) if parameters is not None else None @property - def component(self) -> 'webrtc.RTCIceComponent': + def component(self) -> webrtc.RTCIceComponent: """:obj:`webrtc.RTCIceComponent`: The ICE component being used by the transport, ``rtp`` or ``rtcp``.""" return self._native_obj.component @property - def gathering_state(self) -> 'webrtc.CricketIceGatheringState': + def gathering_state(self) -> webrtc.CricketIceGatheringState: """:obj:`webrtc.CricketIceGatheringState`: The gathering state of the ICE agent.""" return self._native_obj.gatheringState @property - def role(self) -> Optional['webrtc.RTCIceRole']: - """:obj:`webrtc.RTCIceRole`, optional: Whether the ICE agent is the one that makes the final decision as to - the candidate pair to use or not. :obj:`None` for a standalone transport that isn't started.""" + def role(self) -> webrtc.RTCIceRole | None: + """:obj:`webrtc.RTCIceRole`, optional: Whether the ICE agent decides the candidate pair to use. + + :obj:`None` for a standalone transport that isn't started. + """ role = self._native_obj.role if role == RTCIceRole.unknown and self._native_obj._standalone: return None return role @property - def state(self) -> 'webrtc.RTCIceTransportState': + def state(self) -> webrtc.RTCIceTransportState: """:obj:`webrtc.RTCIceTransportState`: The current state of the ICE agent. Note: diff --git a/python-webrtc/python/webrtc/interfaces/rtc_peer_connection.py b/python-webrtc/python/webrtc/interfaces/rtc_peer_connection.py index 5cf695c..74a590a 100644 --- a/python-webrtc/python/webrtc/interfaces/rtc_peer_connection.py +++ b/python-webrtc/python/webrtc/interfaces/rtc_peer_connection.py @@ -5,11 +5,16 @@ # that can be found in the LICENSE.md file in the root of the project. # +"""RTCPeerConnection of WebRTC.""" + +from __future__ import annotations + import asyncio import dataclasses import re -from typing import TYPE_CHECKING, Any, Dict, Iterable, List, Optional, Union +from typing import TYPE_CHECKING, Any, ClassVar, Union +import webrtc from webrtc import ( Event, InvalidAccessError, @@ -35,17 +40,20 @@ WebRTCObject, wrtc, ) +from webrtc.interfaces.rtc_data_channel import RTCDataChannelInit, check_utf8_length from webrtc.utils.events import EventTarget -from webrtc.utils.names import snake_case +from webrtc.utils.names import members from webrtc.utils.native_calls import call_native from webrtc.utils.operations import OperationsChain, later from webrtc.utils.task_queue import TaskQueue if TYPE_CHECKING: - import webrtc + from contextlib import AbstractAsyncContextManager + + from typing_extensions import Self #: A description, as the methods that set one take it -_Description = Union[RTCSessionDescription, RTCSessionDescriptionInit, Dict[str, Any]] +_Description = Union[RTCSessionDescription, RTCSessionDescriptionInit, dict[str, Any]] # the signaling states a local description of a type can be set in _LOCAL_DESCRIPTION_STATES = { @@ -64,10 +72,10 @@ _RID = re.compile(r'[A-Za-z0-9]{1,16}') -class RTCPeerConnection(WebRTCObject, EventTarget): - """The RTCPeerConnection interface represents a WebRTC connection between the local computer and a remote peer. - It provides methods to connect to a remote peer, maintain and monitor the connection, and close the connection - once it's no longer needed. +class RTCPeerConnection(WebRTCObject[wrtc.RTCPeerConnection], EventTarget): + """A WebRTC connection between the local computer and a remote peer. + + It connects to the remote peer, maintains and monitors the connection, and closes it once it's no longer needed. Events (see :meth:`on`): ``negotiationneeded`` (:obj:`webrtc.Event`): Negotiation (an offer/answer exchange) is needed. @@ -86,10 +94,10 @@ class RTCPeerConnection(WebRTCObject, EventTarget): configuration (:obj:`webrtc.RTCConfiguration`, optional): The configuration of the connection. Raises: - :obj:`webrtc.InvalidSyntaxError`: If an ICE server URL is invalid. - :obj:`webrtc.InvalidAccessError`: If a TURN server has no credentials. - :obj:`ValueError`: If a member of the configuration is out of range. - :obj:`TypeError`: If a member of the configuration has a wrong type, or a value its enum doesn't have. + webrtc.InvalidSyntaxError: If an ICE server URL is invalid. + webrtc.InvalidAccessError: If a TURN server has no credentials. + ValueError: If a member of the configuration is out of range. + TypeError: If a member of the configuration has a wrong type, or a value its enum doesn't have. """ _class = wrtc.RTCPeerConnection @@ -106,36 +114,43 @@ class RTCPeerConnection(WebRTCObject, EventTarget): ) # the native method that surfaces the state of each state event - _STATE_EVENTS = { + _STATE_EVENTS: ClassVar[dict[str, str]] = { 'signalingstatechange': '_surfaceSignalingState', 'iceconnectionstatechange': '_surfaceIceConnectionState', 'icegatheringstatechange': '_surfaceIceGatheringState', 'connectionstatechange': '_surfaceConnectionState', } + # the method that creates the event object of each event, from its native arguments + _EVENT_CREATORS: ClassVar[dict[str, str]] = { + 'negotiationneeded': '_negotiation_needed_event', + 'icecandidate': '_ice_candidate_event', + 'icecandidateerror': '_ice_candidate_error_event', + 'datachannel': '_data_channel_event', + 'track': '_track_event', + } #: The operations chain, created on first use (a wrapper of a native connection doesn't run __init__) - _chain: Optional[OperationsChain] = None + _chain: OperationsChain | None = None #: The id of a negotiationneeded event that waits for the operations chain to empty - _deferred_negotiation_id: Optional[int] = None + _deferred_negotiation_id: int | None = None - def __init__(self, configuration: Optional['webrtc.RTCConfiguration'] = None): + def __init__(self, configuration: webrtc.RTCConfiguration | None = None) -> None: super().__init__(self._class(configuration._to_native() if configuration is not None else None)) self._attach() @classmethod - def _wrap(cls, item) -> 'RTCPeerConnection': + def _wrap(cls, item: wrtc.RTCPeerConnection) -> Self: # the wrapper that owns the listeners (and so the operations chain), when there is one listeners = item._listeners if listeners is not None and isinstance(listeners.target, cls): return listeners.target # not attached, so the connection object of the application gets the listeners once it registers a handler connection = cls.__new__(cls) - WebRTCObject.__init__(connection, item) + connection._init_native(item) return connection - def _operation(self): - """Chains an operation (like setting a description) after the ones that are running - (see :obj:`webrtc.utils.operations.OperationsChain`).""" + def _operation(self) -> AbstractAsyncContextManager[None]: + """Chains an operation, like setting a description, after the running ones (see :obj:`OperationsChain`).""" if self._chain is None: self._chain = OperationsChain(self._chain_emptied) return self._chain.operation() @@ -147,17 +162,20 @@ def _chain_emptied(self) -> None: asyncio.get_running_loop().call_soon(self._dispatch, 'negotiationneeded', event_id) def _check_state(self, operation: str, *allowed: RTCSignalingState) -> None: - """Raises :obj:`webrtc.InvalidStateError` if the connection is closed, or not in one of the allowed states - when they're given.""" + """Checks the connection isn't closed, and is in one of the allowed states if they're given. + + Raises: + webrtc.InvalidStateError: If it isn't. + """ state = self.signaling_state if state == RTCSignalingState.closed: - raise InvalidStateError(f"Can not {operation}: the RTCPeerConnection's signalingState is 'closed'") + msg = f"Can not {operation}: the RTCPeerConnection's signalingState is 'closed'" + raise InvalidStateError(msg) if allowed and state not in allowed: - raise InvalidStateError(f'Can not {operation} in the {state} signaling state') - - def _on_event(self, name: str, *args): - from webrtc import RTCDataChannel + msg = f'Can not {operation} in the {state} signaling state' + raise InvalidStateError(msg) + def _on_event(self, name: str, *args: object) -> None: if name == '_gatheringcomplete': transports, state = args self._complete_gathering(transports, state) @@ -171,23 +189,21 @@ def _on_event(self, name: str, *args): _, descriptions = args # the descriptions as the change left them self._native_obj._applyDescriptions(descriptions) - elif name in ('icecandidate', 'icegatheringstatechange'): + elif name in {'icecandidate', 'icegatheringstatechange'}: # the local description gains candidates (and loses pending ones) along with these events self._native_obj._refreshDescriptions() elif name == 'datachannel': (channel,) = args - RTCDataChannel._wrap(channel) + webrtc.RTCDataChannel._wrap(channel) # the events of the channel follow the handlers of this one TaskQueue.post_to_running(channel._release) - def _complete_gathering( - self, transports: List['wrtc.RTCIceTransport'], state: 'webrtc.RTCIceGatheringState' - ) -> None: - """The ICE transports and the connection complete gathering, and the candidates end, in a single task: - every handler sees all of them complete.""" - from webrtc import RTCIceTransport + def _complete_gathering(self, transports: list[wrtc.RTCIceTransport], state: webrtc.RTCIceGatheringState) -> None: + """The ICE transports and the connection complete gathering, and the candidates end, in a single task. - ice_transports = RTCIceTransport._wrap_many(transports) + So every handler sees all of them complete. + """ + ice_transports = webrtc.RTCIceTransport._wrap_many(transports) for ice_transport in ice_transports: ice_transport._native_obj._surfaceGatheringState(state) self._native_obj._surfaceIceGatheringState(state) @@ -198,32 +214,16 @@ def _complete_gathering( # the end of candidates is an icecandidate event without a candidate self._dispatch('icecandidate') - def _create_event(self, name: str, *args): - from webrtc import RTCDataChannel - + def _create_event(self, name: str, *args: object) -> webrtc.Event | None: # events queued before close() aren't delivered after it if self._native_obj.signalingState == RTCSignalingState.closed: return None + creator = self._EVENT_CREATORS.get(name) + if creator is None: + return super()._create_event(name, *args) + return getattr(self, creator)(*args) - if name == 'negotiationneeded': - (event_id,) = args - return self._negotiation_needed_event(event_id) - if name == 'icecandidate': - return self._ice_candidate_event(*args) - if name == 'icecandidateerror': - address, port, url, error_code, error_text = args - return RTCPeerConnectionIceErrorEvent( - name, address or None, port or None, url, error_code, error_text, target=self - ) - if name == 'datachannel': - (channel,) = args - return RTCDataChannelEvent(name, RTCDataChannel._wrap(channel), target=self) - if name == 'track': - transceiver, receiver, streams = args - return self._track_event(transceiver, receiver, streams) - return super()._create_event(name, *args) - - def _negotiation_needed_event(self, event_id: int) -> Optional['webrtc.Event']: + def _negotiation_needed_event(self, event_id: int) -> webrtc.Event | None: # not while operations are chained, but once they're done, if it's still needed if self._chain is not None and self._chain.busy: self._deferred_negotiation_id = event_id @@ -232,30 +232,39 @@ def _negotiation_needed_event(self, event_id: int) -> Optional['webrtc.Event']: return None return Event('negotiationneeded', self) - def _ice_candidate_event(self, candidate: Optional['wrtc.IceCandidateInit'] = None) -> 'webrtc.Event': + def _ice_candidate_event(self, candidate: wrtc.IceCandidateInit | None = None) -> webrtc.Event: if candidate is None: return RTCPeerConnectionIceEvent('icecandidate', None, None, target=self) kwargs = candidate.kwargs() return RTCPeerConnectionIceEvent('icecandidate', RTCIceCandidate(**kwargs), kwargs['url'], target=self) - def _track_event( - self, transceiver: 'wrtc.RTCRtpTransceiver', receiver: 'wrtc.RTCRtpReceiver', streams: List['wrtc.MediaStream'] - ) -> 'webrtc.Event': - from webrtc import MediaStream, RTCRtpReceiver, RTCRtpTransceiver + def _ice_candidate_error_event(self, *native: object) -> webrtc.Event: + address, port, url, error_code, error_text = native + return RTCPeerConnectionIceErrorEvent( + 'icecandidateerror', address or None, port or None, url, error_code, error_text, target=self + ) + + def _data_channel_event(self, channel: wrtc.RTCDataChannel) -> webrtc.Event: + return RTCDataChannelEvent('datachannel', webrtc.RTCDataChannel._wrap(channel), target=self) - receiver = RTCRtpReceiver._wrap(receiver) + def _track_event( + self, transceiver: wrtc.RTCRtpTransceiver, receiver: wrtc.RTCRtpReceiver, streams: list[wrtc.MediaStream] + ) -> webrtc.Event: + wrapped_receiver = webrtc.RTCRtpReceiver._wrap(receiver) return RTCTrackEvent( 'track', - receiver, - receiver.track, - MediaStream._wrap_many(streams), - RTCRtpTransceiver._wrap(transceiver), + wrapped_receiver, + wrapped_receiver.track, + webrtc.MediaStream._wrap_many(streams), + webrtc.RTCRtpTransceiver._wrap(transceiver), target=self, ) - def _apply_legacy_offer_option(self, kind: 'webrtc.MediaType', receive: Optional[bool]) -> None: - """``offer_to_receive_audio`` or ``offer_to_receive_video`` of :meth:`create_offer`, as the specification - defines them in terms of transceivers.""" + def _apply_legacy_offer_option(self, kind: webrtc.MediaType, *, receive: bool | None) -> None: + """Applies ``offer_to_receive_audio`` or ``offer_to_receive_video`` of :meth:`create_offer`. + + As the specification defines them, in terms of transceivers. + """ if receive is None: return directions = TransceiverDirection @@ -266,7 +275,7 @@ def _apply_legacy_offer_option(self, kind: 'webrtc.MediaType', receive: Optional transceiver.direction = directions.sendonly elif transceiver.direction == directions.recvonly: transceiver.direction = directions.inactive - elif not any(t.direction in (directions.sendrecv, directions.recvonly) for t in transceivers): + elif not any(t.direction in {directions.sendrecv, directions.recvonly} for t in transceivers): self.add_transceiver(kind, RtpTransceiverInit(direction=directions.recvonly)) def _completed_description(self) -> None: @@ -281,11 +290,12 @@ async def create_offer( self, *, ice_restart: bool = False, - offer_to_receive_audio: Optional[bool] = None, - offer_to_receive_video: Optional[bool] = None, + offer_to_receive_audio: bool | None = None, + offer_to_receive_video: bool | None = None, voice_activity_detection: bool = True, - ) -> 'webrtc.RTCSessionDescriptionInit': + ) -> webrtc.RTCSessionDescriptionInit: """Initiates the creation of an SDP offer for the purpose of starting a new WebRTC connection to a remote peer. + The SDP offer includes information about any MediaStreamTrack objects already attached to the WebRTC session, codec, and options supported by the machine, as well as any candidates already gathered by the ICE agent, for the purpose of being sent over the signaling channel to a potential peer to request a connection or to update @@ -304,19 +314,20 @@ async def create_offer( :obj:`webrtc.RTCSessionDescriptionInit`: The offer, to set with :meth:`set_local_description`. Raises: - :obj:`webrtc.InvalidStateError`: If the signaling state is neither stable nor have-local-offer. + webrtc.InvalidStateError: If the signaling state is neither stable nor have-local-offer. """ async with self._operation(): self._check_state('create an offer', RTCSignalingState.stable, RTCSignalingState.have_local_offer) - self._apply_legacy_offer_option(MediaType.audio, offer_to_receive_audio) - self._apply_legacy_offer_option(MediaType.video, offer_to_receive_video) + self._apply_legacy_offer_option(MediaType.audio, receive=offer_to_receive_audio) + self._apply_legacy_offer_option(MediaType.video, receive=offer_to_receive_video) await later() return _init_of(await call_native(self._native_obj.createOffer, ice_restart, voice_activity_detection)) - async def create_answer(self, *, voice_activity_detection: bool = True) -> 'webrtc.RTCSessionDescriptionInit': - """Initiates the creation an SDP answer to an offer received from a remote peer during the offer/answer - negotiation of a WebRTC connection. The answer contains information about any media already attached to the - session, codecs and options supported by the machine, and any ICE candidates already gathered. + async def create_answer(self, *, voice_activity_detection: bool = True) -> webrtc.RTCSessionDescriptionInit: + """Creates an SDP answer to an offer received from the remote peer. + + The answer contains information about any media already attached to the session, codecs and options supported + by the machine, and any ICE candidates already gathered. Args: voice_activity_detection (:obj:`bool`, optional): Whether audio codecs may use voice activity detection. @@ -325,7 +336,7 @@ async def create_answer(self, *, voice_activity_detection: bool = True) -> 'webr :obj:`webrtc.RTCSessionDescriptionInit`: The answer, to set with :meth:`set_local_description`. Raises: - :obj:`webrtc.InvalidStateError`: If the connection is closed or has no remote offer. + webrtc.InvalidStateError: If the connection is closed or has no remote offer. """ async with self._operation(): self._check_state( @@ -334,9 +345,10 @@ async def create_answer(self, *, voice_activity_detection: bool = True) -> 'webr await later() return _init_of(await call_native(self._native_obj.createAnswer, voice_activity_detection)) - async def set_local_description(self, description: Optional[_Description] = None) -> None: - """Changes the local description associated with the connection. This description specifies the properties - of the local end of the connection, including the media format. + async def set_local_description(self, description: _Description | None = None) -> None: + """Changes the local description associated with the connection. + + This description specifies the properties of the local end of the connection, including the media format. Args: description (:obj:`webrtc.RTCSessionDescription`, optional): The description, as returned by @@ -346,10 +358,10 @@ async def set_local_description(self, description: Optional[_Description] = None is created and set. Raises: - :obj:`webrtc.InvalidStateError`: If the type doesn't match the signaling state, or the connection is closed. - :obj:`webrtc.InvalidModificationError`: If the SDP isn't the one :meth:`create_offer` + webrtc.InvalidStateError: If the type doesn't match the signaling state, or the connection is closed. + webrtc.InvalidModificationError: If the SDP isn't the one :meth:`create_offer` or :meth:`create_answer` returned last. - :obj:`webrtc.RTCError`: If the SDP can't be parsed (``sdp_syntax_error``). + webrtc.RTCError: If the SDP can't be parsed (``sdp_syntax_error``). """ init = _description_init(description, allow_implicit=True) async with self._operation(): @@ -360,8 +372,9 @@ async def set_local_description(self, description: Optional[_Description] = None self._completed_description() async def set_remote_description(self, description: _Description) -> None: - """Sets the specified session description as the remote peer's current offer or answer. The description - specifies the properties of the remote end of the connection, including the media format. + """Sets the specified session description as the remote peer's current offer or answer. + + The description specifies the properties of the remote end of the connection, including the media format. An offer set while there's a local offer rolls the local one back first. @@ -371,9 +384,9 @@ async def set_remote_description(self, description: _Description) -> None: is accepted too. Raises: - :obj:`webrtc.InvalidStateError`: If the type doesn't match the signaling state, or the connection is closed. - :obj:`webrtc.RTCError`: If the SDP can't be parsed (``sdp_syntax_error``). - :obj:`webrtc.InvalidAccessError`: If the description can't be applied. + webrtc.InvalidStateError: If the type doesn't match the signaling state, or the connection is closed. + webrtc.RTCError: If the SDP can't be parsed (``sdp_syntax_error``). + webrtc.InvalidAccessError: If the description can't be applied. """ init = _description_init(description, allow_implicit=False) async with self._operation(): @@ -384,9 +397,9 @@ async def set_remote_description(self, description: _Description) -> None: def add_track( self, - track: 'webrtc.MediaStreamTrack', - stream: Optional[Union['webrtc.MediaStream', List['webrtc.MediaStream']]] = None, - ) -> 'webrtc.RTCRtpSender': + track: webrtc.MediaStreamTrack, + stream: webrtc.MediaStream | list[webrtc.MediaStream] | None = None, + ) -> webrtc.RTCRtpSender: """Adds a new :obj:`webrtc.MediaStreamTrack` to the set of tracks which will be transmitted to the other peer. Args: @@ -399,8 +412,6 @@ def add_track( :obj:`webrtc.RTCRtpSender`: The :obj:`webrtc.RTCRtpSender` object which will be used to transmit the media data. """ - from webrtc import RTCRtpSender - if not stream: sender = self._native_obj.addTrack(track._native_obj, None) elif isinstance(stream, list): @@ -409,15 +420,16 @@ def add_track( else: sender = self._native_obj.addTrack(track._native_obj, stream._native_obj) - return RTCRtpSender._wrap(sender) + return webrtc.RTCRtpSender._wrap(sender) def add_transceiver( self, - track_or_kind: Union['webrtc.MediaStreamTrack', 'webrtc.MediaType'], - init: Optional[Union['webrtc.RtpTransceiverInit', Dict[str, Any]]] = None, - ) -> 'webrtc.RTCRtpTransceiver': - """Creates a new :obj:`webrtc.RTCRtpTransceiver` and adds it to the set of transceivers associated with the - connection. Each transceiver represents a bidirectional stream, with both an :obj:`webrtc.RTCRtpSender` and + track_or_kind: webrtc.MediaStreamTrack | webrtc.MediaType, + init: webrtc.RtpTransceiverInit | dict[str, Any] | None = None, + ) -> webrtc.RTCRtpTransceiver: + """Creates a new :obj:`webrtc.RTCRtpTransceiver` and adds it to the transceivers of the connection. + + Each transceiver represents a bidirectional stream, with both an :obj:`webrtc.RTCRtpSender` and an :obj:`webrtc.RTCRtpReceiver` associated with it. Args: @@ -433,16 +445,15 @@ def add_transceiver( :obj:`webrtc.RTCRtpTransceiver`: The new transceiver. Raises: - :obj:`TypeError`: If the kind is neither audio nor video. - :obj:`ValueError`: If a ``rid`` of the send encodings is invalid, or missing or repeated with several + TypeError: If the kind is neither audio nor video. + ValueError: If a ``rid`` of the send encodings is invalid, or missing or repeated with several encodings. - :obj:`webrtc.OperationError`: If the codec of a send encoding can't be sent. + webrtc.OperationError: If the codec of a send encoding can't be sent. """ - from webrtc import MediaStreamTrack, RTCRtpTransceiver - - kind = track_or_kind.kind if isinstance(track_or_kind, MediaStreamTrack) else track_or_kind - if kind not in (MediaType.audio, MediaType.video): - raise TypeError(f'{kind!r} is not a kind of track') + kind = track_or_kind.kind if isinstance(track_or_kind, webrtc.MediaStreamTrack) else track_or_kind + if kind not in {MediaType.audio, MediaType.video}: + msg = f'{kind!r} is not a kind of track' + raise TypeError(msg) native_init = None if isinstance(init, dict): init = _transceiver_init(init) @@ -452,16 +463,15 @@ def add_transceiver( encodings = [encoding._for_kind(kind) for encoding in init.send_encodings] native_init = RtpTransceiverInit(init.direction, encodings, init.streams)._native_obj - if isinstance(track_or_kind, MediaStreamTrack): + if isinstance(track_or_kind, webrtc.MediaStreamTrack): transceiver = self._native_obj.addTransceiver(track_or_kind._native_obj, native_init) else: transceiver = self._native_obj.addTransceiver(track_or_kind, native_init) - return RTCRtpTransceiver._wrap(transceiver) + return webrtc.RTCRtpTransceiver._wrap(transceiver) - def get_transceivers(self) -> List['webrtc.RTCRtpTransceiver']: - """Returns a :obj:`list` of the :obj:`webrtc.RTCRtpTransceiver` objects being used to send and - receive data on the connection. + def get_transceivers(self) -> list[webrtc.RTCRtpTransceiver]: + """Returns the transceivers the connection sends and receives media with. Returns: :obj:`list` of :obj:`webrtc.RTCRtpTransceiver`: An array of the :obj:`webrtc.RTCRtpTransceiver` objects @@ -469,14 +479,12 @@ def get_transceivers(self) -> List['webrtc.RTCRtpTransceiver']: on the :obj:`webrtc.RTCPeerConnection`. The list is in the order in which the transceivers were added to the connection. """ - from webrtc import RTCRtpTransceiver + return webrtc.RTCRtpTransceiver._wrap_many(self._native_obj.getTransceivers()) - return RTCRtpTransceiver._wrap_many(self._native_obj.getTransceivers()) + def get_senders(self) -> list[webrtc.RTCRtpSender]: + """Returns the senders of the connection, each of which sends the media of one track. - def get_senders(self) -> List['webrtc.RTCRtpSender']: - """Returns an array of :obj:`webrtc.RTCRtpSender` objects, each of which represents the RTP sender responsible - for transmitting one track's data. A sender object provides methods and properties for examining - and controlling the encoding and transmission of the track's data. + A sender examines and controls the encoding and transmission of the media of its track. Note: The order of the returned :obj:`webrtc.RTCRtpSender` objects is not defined by the specification, @@ -486,14 +494,13 @@ def get_senders(self) -> List['webrtc.RTCRtpSender']: :obj:`list` of :obj:`webrtc.RTCRtpSender`: An array of :obj:`webrtc.RTCRtpSender` objects, one for each track on the connection. The array is empty if there are no RTP senders on the connection. """ - from webrtc import RTCRtpSender + return webrtc.RTCRtpSender._wrap_many(self._native_obj.getSenders()) - return RTCRtpSender._wrap_many(self._native_obj.getSenders()) + def get_receivers(self) -> list[webrtc.RTCRtpReceiver]: + """Returns an array of :obj:`webrtc.RTCRtpReceiver` objects, each of which represents one RTP receiver. - def get_receivers(self) -> List['webrtc.RTCRtpReceiver']: - """Returns an array of :obj:`webrtc.RTCRtpReceiver` objects, each of which represents one RTP receiver. Each RTP - receiver manages the reception and decoding of data for a :obj:`webrtc.MediaStreamTrack` - on an :obj:`webrtc.RTCPeerConnection`. + Each RTP receiver manages the reception and decoding of data for a :obj:`webrtc.MediaStreamTrack` on an + :obj:`webrtc.RTCPeerConnection`. Note: The order of the returned :obj:`webrtc.RTCRtpReceiver` objects is not defined by the specification, @@ -503,26 +510,23 @@ def get_receivers(self) -> List['webrtc.RTCRtpReceiver']: :obj:`list` of :obj:`webrtc.RTCRtpReceiver`: An array of :obj:`webrtc.RTCRtpReceiver` objects, one for each track on the connection. The array is empty if there are no RTP receivers on the connection. """ - from webrtc import RTCRtpReceiver + return webrtc.RTCRtpReceiver._wrap_many(self._native_obj.getReceivers()) - return RTCRtpReceiver._wrap_many(self._native_obj.getReceivers()) + def remove_track(self, sender: webrtc.RTCRtpSender) -> None: + """Stops sending the track of a sender, which stays in :meth:`get_senders`. - def remove_track(self, sender: 'webrtc.RTCRtpSender') -> None: - """Stops sending the track of a sender, which stays in :meth:`get_senders`. Does nothing if the sender - has no track. + Does nothing if the sender has no track. Args: sender (:obj:`webrtc.RTCRtpSender`): A sender of this connection. Raises: - :obj:`webrtc.InvalidStateError`: If the connection is closed. - :obj:`webrtc.InvalidAccessError`: If the sender belongs to another connection. + webrtc.InvalidStateError: If the connection is closed. + webrtc.InvalidAccessError: If the sender belongs to another connection. """ - return self._native_obj.removeTrack(sender._native_obj) + self._native_obj.removeTrack(sender._native_obj) - async def add_ice_candidate( - self, candidate: Optional[Union['webrtc.RTCIceCandidate', Dict[str, Any]]] = None - ) -> None: + async def add_ice_candidate(self, candidate: webrtc.RTCIceCandidate | dict[str, Any] | None = None) -> None: """Adds a candidate received from the remote peer to the remote description. Args: @@ -531,82 +535,71 @@ async def add_ice_candidate( :attr:`webrtc.RTCIceCandidate.candidate`, or :obj:`None`, means the end of candidates. Raises: - :obj:`TypeError`: If a non-empty candidate has neither ``sdp_mid`` nor ``sdp_m_line_index``. - :obj:`webrtc.InvalidStateError`: If there's no remote description, or the connection is closed. - :obj:`webrtc.OperationError`: If the candidate can't be parsed or doesn't match a media section. + TypeError: If a non-empty candidate has neither ``sdp_mid`` nor ``sdp_m_line_index``. + webrtc.InvalidStateError: If there's no remote description, or the connection is closed. + webrtc.OperationError: If the candidate can't be parsed or doesn't match a media section. """ if candidate is None: candidate_str, sdp_mid, sdp_m_line_index, ufrag = '', None, None, None else: candidate_str, sdp_mid, sdp_m_line_index, ufrag = RTCIceCandidate._members_of(candidate) if candidate_str and sdp_mid is None and sdp_m_line_index is None: - raise TypeError('sdp_mid and sdp_m_line_index are both None') + msg = 'sdp_mid and sdp_m_line_index are both None' + raise TypeError(msg) async with self._operation(): self._check_state('add an ICE candidate') if self.remote_description is None: - raise InvalidStateError('A candidate can only be added once there is a remote description') + msg = 'A candidate can only be added once there is a remote description' + raise InvalidStateError(msg) await later() await call_native(self._native_obj.addIceCandidate, candidate_str, sdp_mid, sdp_m_line_index, ufrag) def create_data_channel( - self, - label: str, - *, - ordered: bool = True, - max_packet_life_time: Optional[int] = None, - max_retransmits: Optional[int] = None, - protocol: str = '', - negotiated: bool = False, - id: Optional[int] = None, - priority: 'webrtc.RTCPriorityType' = 'low', - ) -> 'webrtc.RTCDataChannel': - """Creates a channel to send messages to the remote peer, negotiated with the next offer - (unless ``negotiated`` is set). + self, label: str, options: webrtc.RTCDataChannelInit | dict[str, Any] | None = None + ) -> webrtc.RTCDataChannel: + """Creates a channel to send messages to the remote peer, negotiated with the next offer. + + Unless ``negotiated`` is set in the options. Args: label (:obj:`str`): The name of the channel, up to 65535 bytes in UTF-8. - ordered (:obj:`bool`, optional): Whether messages are delivered in order. - max_packet_life_time (:obj:`int`, optional): Makes the channel unreliable: how long in milliseconds - a message is retransmitted. - max_retransmits (:obj:`int`, optional): Makes the channel unreliable: how many times a message is - retransmitted. Can't be set along with ``max_packet_life_time``. - protocol (:obj:`str`, optional): The subprotocol name, up to 65535 bytes in UTF-8. - negotiated (:obj:`bool`, optional): Whether the application creates the channel on both ends - with the same ``id``, instead of announcing it to the remote peer. - id (:obj:`int`, optional): The SCTP stream id (0 to 65534) of a ``negotiated`` channel, which requires it. - Ignored otherwise, as the connection picks the id. - priority (:obj:`webrtc.RTCPriorityType`, optional): The priority of the channel. + options (:obj:`webrtc.RTCDataChannelInit` or :obj:`dict`, optional): How to create the channel, or + a dictionary of its members. Returns: :obj:`webrtc.RTCDataChannel`: The channel. Raises: - :obj:`ValueError`: If an argument is out of range, or both ``max_packet_life_time`` and + ValueError: If the label or an option is out of range, or both ``max_packet_life_time`` and ``max_retransmits`` are set, or ``negotiated`` is set without ``id``. - :obj:`webrtc.InvalidStateError`: If the connection is closed. - :obj:`webrtc.OperationError`: If the ``id`` is in use, or no id is left. + webrtc.InvalidStateError: If the connection is closed. + webrtc.OperationError: If the ``id`` is in use, or no id is left. """ - from webrtc import RTCDataChannel - - _check_data_channel_init(label, protocol, max_packet_life_time, max_retransmits, negotiated, id) + if isinstance(options, dict): + options = RTCDataChannelInit( + **members(options, [field.name for field in dataclasses.fields(RTCDataChannelInit)]) + ) + init = options or RTCDataChannelInit() + check_utf8_length('label', label) + init._check() native = self._native_obj.createDataChannel( label, - ordered, - max_packet_life_time, - max_retransmits, - protocol, - negotiated, - id if negotiated else None, - priority, + init.ordered, + init.max_packet_life_time, + init.max_retransmits, + init.protocol, + init.negotiated, + init.id if init.negotiated else None, + init.priority, ) - channel = RTCDataChannel._wrap(native) + channel = webrtc.RTCDataChannel._wrap(native) # handlers registered in this iteration of the loop get the first events if not TaskQueue.post_to_running(native._release): native._release() return channel - async def get_stats(self, selector: Optional['webrtc.MediaStreamTrack'] = None) -> 'webrtc.RTCStatsReport': + async def get_stats(self, selector: webrtc.MediaStreamTrack | None = None) -> webrtc.RTCStatsReport: """Collects the stats of the connection, or of the sender or the receiver of a track. Args: @@ -617,23 +610,25 @@ async def get_stats(self, selector: Optional['webrtc.MediaStreamTrack'] = None) :obj:`webrtc.RTCStatsReport`: The stats. Raises: - :obj:`webrtc.InvalidAccessError`: If no sender or receiver, or more than one, has the track. - :obj:`webrtc.InvalidStateError`: If the connection is closed. + webrtc.InvalidAccessError: If no sender or receiver, or more than one, has the track. + webrtc.InvalidStateError: If the connection is closed. """ if selector is not None: matches = [s for s in self.get_senders() if s.track == selector] matches += [r for r in self.get_receivers() if r.track == selector] if len(matches) != 1: - raise InvalidAccessError(f'{len(matches)} senders and receivers have the track, not exactly one') + msg = f'{len(matches)} senders and receivers have the track, not exactly one' + raise InvalidAccessError(msg) return await matches[0].get_stats() return RTCStatsReport._from_native(await call_native(self._native_obj.getStats), self.get_receivers()) @staticmethod async def generate_certificate( - algorithm: 'webrtc.models.rtc_certificate.Algorithm' = 'ECDSA', expires: Optional[float] = None - ) -> 'webrtc.RTCCertificate': - """Generates a certificate for :attr:`webrtc.RTCConfiguration.certificates`, the same as - :meth:`webrtc.RTCCertificate.generate`. + algorithm: webrtc.models.rtc_certificate.Algorithm = 'ECDSA', expires: float | None = None + ) -> webrtc.RTCCertificate: + """Generates a certificate for :attr:`webrtc.RTCConfiguration.certificates`. + + The same as :meth:`webrtc.RTCCertificate.generate`. Args: algorithm (:obj:`str` or :obj:`dict`, optional): The WebCrypto algorithm of the key. @@ -643,12 +638,12 @@ async def generate_certificate( :obj:`webrtc.RTCCertificate`: The certificate. Raises: - :obj:`webrtc.NotSupportedError`: If the algorithm isn't supported. - :obj:`ValueError`: If ``expires`` is negative. + webrtc.NotSupportedError: If the algorithm isn't supported. + ValueError: If ``expires`` is negative. """ return await RTCCertificate.generate(algorithm, expires) - def get_configuration(self) -> 'webrtc.RTCConfiguration': + def get_configuration(self) -> webrtc.RTCConfiguration: """Returns the configuration of the connection, as it was last set. Returns: @@ -656,7 +651,7 @@ def get_configuration(self) -> 'webrtc.RTCConfiguration': """ return RTCConfiguration._from_native(self._native_obj.getConfiguration()) - def set_configuration(self, configuration: Optional['webrtc.RTCConfiguration'] = None) -> None: + def set_configuration(self, configuration: webrtc.RTCConfiguration | None = None) -> None: """Changes the configuration of the connection. Members that aren't set get their default values. Changed ICE servers or ICE transport policy are used for the candidates gathered next, @@ -666,94 +661,95 @@ def set_configuration(self, configuration: Optional['webrtc.RTCConfiguration'] = configuration (:obj:`webrtc.RTCConfiguration`, optional): The new configuration. Raises: - :obj:`webrtc.InvalidStateError`: If the connection is closed. - :obj:`webrtc.InvalidModificationError`: If a member that can't be changed differs, like + webrtc.InvalidStateError: If the connection is closed. + webrtc.InvalidModificationError: If a member that can't be changed differs, like :attr:`webrtc.RTCConfiguration.bundle_policy` or :attr:`webrtc.RTCConfiguration.always_negotiate_data_channels`. - :obj:`webrtc.InvalidSyntaxError`: If an ICE server URL is invalid. - :obj:`webrtc.InvalidAccessError`: If a TURN server has no credentials. - :obj:`ValueError`: If a member of the configuration is out of range. - :obj:`TypeError`: If a member of the configuration has a wrong type, or a value its enum doesn't have. + webrtc.InvalidSyntaxError: If an ICE server URL is invalid. + webrtc.InvalidAccessError: If a TURN server has no credentials. + ValueError: If a member of the configuration is out of range. + TypeError: If a member of the configuration has a wrong type, or a value its enum doesn't have. """ self._native_obj.setConfiguration((configuration or RTCConfiguration())._to_native()) def restart_ice(self) -> None: """Allows to easily request that ICE candidate gathering be redone on both ends of the connection. - This simplifies the process by allowing the same method to be used by either the caller or the receiver - to trigger an ICE restart.""" - return self._native_obj.restartIce() - def close(self): - """Closes the current peer connection.""" - return self._native_obj.close() + This simplifies the process by allowing the same method to be used by either the caller or the receiver to + trigger an ICE restart. + """ + self._native_obj.restartIce() - @property - def sctp(self) -> Optional['webrtc.RTCSctpTransport']: - """:obj:`webrtc.RTCSctpTransport`, optional: An object describing the SCTP transport layer over which SCTP - data is being sent and received. If SCTP hasn't been negotiated, this value is :obj:`None`.""" - from webrtc import RTCSctpTransport + def close(self) -> None: + """Closes the connection.""" + self._native_obj.close() - return RTCSctpTransport._wrap_optional(self._native_obj.sctp) + @property + def sctp(self) -> webrtc.RTCSctpTransport | None: + """:obj:`webrtc.RTCSctpTransport`, optional: The SCTP transport of the data, :obj:`None` until negotiated.""" + return webrtc.RTCSctpTransport._wrap_optional(self._native_obj.sctp) @property - def local_description(self) -> Optional['webrtc.RTCSessionDescription']: - """:obj:`webrtc.RTCSessionDescription`, optional: The local end of the connection, including ICE candidates - gathered so far. :obj:`None` if the local description hasn't been set yet.""" + def local_description(self) -> webrtc.RTCSessionDescription | None: + """:obj:`webrtc.RTCSessionDescription`, optional: The local end of the connection, :obj:`None` until set. + + It includes the ICE candidates gathered so far. + """ return RTCSessionDescription._wrap_optional(self._native_obj.localDescription) @property - def remote_description(self) -> Optional['webrtc.RTCSessionDescription']: + def remote_description(self) -> webrtc.RTCSessionDescription | None: """:obj:`webrtc.RTCSessionDescription`, optional: The remote end of the connection. - :obj:`None` if the remote description hasn't been set yet.""" + + :obj:`None` if the remote description hasn't been set yet. + """ return RTCSessionDescription._wrap_optional(self._native_obj.remoteDescription) @property - def current_local_description(self) -> Optional['webrtc.RTCSessionDescription']: - """:obj:`webrtc.RTCSessionDescription`, optional: The local description negotiated the last time - the connection was in the stable state.""" + def current_local_description(self) -> webrtc.RTCSessionDescription | None: + """:obj:`webrtc.RTCSessionDescription`, optional: The local description negotiated last, in stable state.""" return RTCSessionDescription._wrap_optional(self._native_obj.currentLocalDescription) @property - def current_remote_description(self) -> Optional['webrtc.RTCSessionDescription']: - """:obj:`webrtc.RTCSessionDescription`, optional: The remote description negotiated the last time - the connection was in the stable state.""" + def current_remote_description(self) -> webrtc.RTCSessionDescription | None: + """:obj:`webrtc.RTCSessionDescription`, optional: The remote description negotiated last, in stable state.""" return RTCSessionDescription._wrap_optional(self._native_obj.currentRemoteDescription) @property - def pending_local_description(self) -> Optional['webrtc.RTCSessionDescription']: - """:obj:`webrtc.RTCSessionDescription`, optional: The local description being negotiated, :obj:`None` - in the stable state.""" + def pending_local_description(self) -> webrtc.RTCSessionDescription | None: + """:obj:`webrtc.RTCSessionDescription`, optional: The local description being negotiated, if any.""" return RTCSessionDescription._wrap_optional(self._native_obj.pendingLocalDescription) @property - def pending_remote_description(self) -> Optional['webrtc.RTCSessionDescription']: - """:obj:`webrtc.RTCSessionDescription`, optional: The remote description being negotiated, :obj:`None` - in the stable state.""" + def pending_remote_description(self) -> webrtc.RTCSessionDescription | None: + """:obj:`webrtc.RTCSessionDescription`, optional: The remote description being negotiated, if any.""" return RTCSessionDescription._wrap_optional(self._native_obj.pendingRemoteDescription) @property - def can_trickle_ice_candidates(self) -> Optional[bool]: - """:obj:`bool`, optional: Whether the remote peer takes candidates one by one (trickle ICE), - :obj:`None` until there's a remote description.""" + def can_trickle_ice_candidates(self) -> bool | None: + """:obj:`bool`, optional: Whether the remote peer takes candidates one by one (trickle ICE). + + :obj:`None` until there's a remote description. + """ return self._native_obj.canTrickleIceCandidates @property - def connection_state(self) -> 'webrtc.RTCPeerConnectionState': + def connection_state(self) -> webrtc.RTCPeerConnectionState: """:obj:`webrtc.RTCPeerConnectionState`: The current state of the connection.""" return self._native_obj.connectionState @property - def signaling_state(self) -> 'webrtc.RTCSignalingState': + def signaling_state(self) -> webrtc.RTCSignalingState: """:obj:`webrtc.RTCSignalingState`: The state of the signaling process.""" return self._native_obj.signalingState @property - def ice_connection_state(self) -> 'webrtc.RTCIceConnectionState': + def ice_connection_state(self) -> webrtc.RTCIceConnectionState: """:obj:`webrtc.RTCIceConnectionState`: The state of the ICE agent.""" return self._native_obj.iceConnectionState @property - def ice_gathering_state(self) -> 'webrtc.RTCIceGatheringState': + def ice_gathering_state(self) -> webrtc.RTCIceGatheringState: """:obj:`webrtc.RTCIceGatheringState`: The ICE candidate gathering state.""" return self._native_obj.iceGatheringState @@ -816,8 +812,8 @@ def ice_gathering_state(self) -> 'webrtc.RTCIceGatheringState': def _description_init( - description: Optional[_Description], allow_implicit: bool -) -> Optional['wrtc.RTCSessionDescriptionInit']: + description: _Description | None, *, allow_implicit: bool +) -> wrtc.RTCSessionDescriptionInit | None: """The native RTCSessionDescriptionInit of a description, or :obj:`None` for an implicit one.""" if description is None and allow_implicit: return None @@ -825,77 +821,54 @@ def _description_init( if description.get('type') is None: if allow_implicit and not description.get('sdp'): return None - raise TypeError('the type of a description is required') + msg = 'the type of a description is required' + raise TypeError(msg) description = RTCSessionDescription(description) if isinstance(description, RTCSessionDescription): return description._native_obj.init if isinstance(description, RTCSessionDescriptionInit): return description._native_obj - raise TypeError(f'expected an RTCSessionDescription, not {type(description).__name__}') + msg = f'expected an RTCSessionDescription, not {type(description).__name__}' + raise TypeError(msg) -def _init_of(description: 'wrtc.RTCSessionDescription') -> 'webrtc.RTCSessionDescriptionInit': +def _init_of(description: wrtc.RTCSessionDescription) -> webrtc.RTCSessionDescriptionInit: return RTCSessionDescriptionInit(description.type, description.sdp) -def _members(value: Dict[str, Any], names: Iterable[str]) -> Dict[str, Any]: - """The members of a dictionary, with snake_case or camelCase names: unknown ones are ignored, as in WebIDL""" - members = {snake_case(name): member for name, member in value.items()} - return {name: member for name, member in members.items() if name in names} - - -def _transceiver_init(init: Dict[str, Any]) -> RtpTransceiverInit: - """An init from a dictionary, as in browsers, with its encodings dictionaries too""" - members = _members(init, ('direction', 'send_encodings', 'streams')) - encodings = members.get('send_encodings') +def _transceiver_init(init: dict[str, Any]) -> RtpTransceiverInit: + """An init from a dictionary, as in browsers, with its encodings dictionaries too.""" + values = members(init, ('direction', 'send_encodings', 'streams')) + encodings = values.get('send_encodings') if encodings is not None: names = [field.name for field in dataclasses.fields(RTCRtpEncodingParameters)] - members['send_encodings'] = [ - RTCRtpEncodingParameters(**_members(e, names)) if isinstance(e, dict) else e for e in encodings + values['send_encodings'] = [ + RTCRtpEncodingParameters(**members(e, names)) if isinstance(e, dict) else e for e in encodings ] - return RtpTransceiverInit(**members) + return RtpTransceiverInit(**values) -def _check_send_encodings(encodings: List['webrtc.RTCRtpEncodingParameters'], kind: 'webrtc.MediaType') -> None: - """Validates the send encodings of a new transceiver, as the specification requires.""" - from webrtc import RTCRtpSender +def _check_send_encodings(encodings: list[webrtc.RTCRtpEncodingParameters], kind: webrtc.MediaType) -> None: + """Validates the send encodings of a new transceiver, as the specification requires. + Raises: + ValueError: If a ``rid`` is invalid, or missing or repeated with several encodings. + webrtc.OperationError: If the codec of an encoding can't be sent. + """ rids = [e.rid for e in encodings] for rid in rids: if rid is not None and not _RID.fullmatch(rid): - raise ValueError(f'{rid!r} is not a valid rid: 1 to 16 letters and digits') + msg = f'{rid!r} is not a valid rid: 1 to 16 letters and digits' + raise ValueError(msg) if len(encodings) > 1 and (None in rids or len(set(rids)) != len(rids)): - raise ValueError('every encoding needs a distinct rid when there are several') + msg = 'every encoding needs a distinct rid when there are several' + raise ValueError(msg) codecs = [e.codec for e in encodings if e.codec is not None] if codecs: - capabilities = RTCRtpSender.get_capabilities(kind) + capabilities = webrtc.RTCRtpSender.get_capabilities(kind) supported = capabilities.codecs if capabilities is not None else [] for codec in codecs: if not any(RTCRtpCodec._matches(c, codec) for c in supported): - raise OperationError(f'{codec.mime_type} can not be sent') - - -def _check_data_channel_init( - label: str, - protocol: str, - max_packet_life_time: Optional[int], - max_retransmits: Optional[int], - negotiated: bool, - id: Optional[int], -) -> None: - """Validates the arguments of :meth:`RTCPeerConnection.create_data_channel`, as the specification requires.""" - for name, value in (('label', label), ('protocol', protocol)): - if len(value.encode()) > 65535: - raise ValueError(f'{name} is longer than 65535 bytes') - for name, value in (('max_packet_life_time', max_packet_life_time), ('max_retransmits', max_retransmits)): - if value is not None and not 0 <= value <= 65535: - raise ValueError(f'{name} must be from 0 to 65535, not {value}') - if max_packet_life_time is not None and max_retransmits is not None: - raise ValueError('max_packet_life_time and max_retransmits can not both be set') - if not negotiated: - return - if id is None: - raise ValueError('a negotiated channel needs an id') - if not 0 <= id <= 65534: - raise ValueError(f'id must be from 0 to 65534, not {id}') + msg = f'{codec.mime_type} can not be sent' + raise OperationError(msg) diff --git a/python-webrtc/python/webrtc/interfaces/rtc_rtp_receiver.py b/python-webrtc/python/webrtc/interfaces/rtc_rtp_receiver.py index a9c6c41..c31fbe2 100644 --- a/python-webrtc/python/webrtc/interfaces/rtc_rtp_receiver.py +++ b/python-webrtc/python/webrtc/interfaces/rtc_rtp_receiver.py @@ -5,8 +5,11 @@ # that can be found in the LICENSE.md file in the root of the project. # -from typing import TYPE_CHECKING, List, Optional +"""RTCRtpReceiver of WebRTC.""" +from __future__ import annotations + +import webrtc from webrtc import ( InvalidRangeError, RTCRtpCapabilities, @@ -19,51 +22,46 @@ ) from webrtc.utils.native_calls import call_native -if TYPE_CHECKING: - import webrtc +#: The maximum jitter_buffer_target, in milliseconds +_MAX_JITTER_BUFFER_TARGET = 4000 -class RTCRtpReceiver(WebRTCObject): - """The :obj:`webrtc.RTCRtpReceiver` interface of the WebRTC API manages the reception and decoding of data - for a :obj:`webrtc.MediaStreamTrack` on an :obj:`webrtc.RTCPeerConnection`.""" +class RTCRtpReceiver(WebRTCObject[wrtc.RTCRtpReceiver]): + """Receives and decodes the media of a :obj:`webrtc.MediaStreamTrack` of an :obj:`webrtc.RTCPeerConnection`.""" _class = wrtc.RTCRtpReceiver - def _sources(self, synchronization: bool) -> List['webrtc.RTCRtpContributingSource']: + def _sources(self, *, synchronization: bool) -> list[webrtc.RTCRtpContributingSource]: cls = RTCRtpSynchronizationSource if synchronization else RTCRtpContributingSource # each native source starts with whether it's an SSRC return [cls._from_native(source) for source in self._native_obj._getSources() if source[0] == synchronization] @property - def track(self) -> 'webrtc.MediaStreamTrack': - """:obj:`webrtc.MediaStreamTrack`: The :obj:`webrtc.MediaStreamTrack` associated with the current - :obj:`webrtc.RTCRtpReceiver` instance.""" - from webrtc import MediaStreamTrack - - return MediaStreamTrack._wrap(self._native_obj.track) + def track(self) -> webrtc.MediaStreamTrack: + """:obj:`webrtc.MediaStreamTrack`: The track of the received media.""" + return webrtc.MediaStreamTrack._wrap(self._native_obj.track) @property - def transport(self) -> Optional['webrtc.RTCDtlsTransport']: - """:obj:`webrtc.RTCDtlsTransport`, optional: An object representing the underlying transport being used by - the receiver to exchange packets with the remote peer, or :obj:`None` if the receiver isn't yet connected - to transport.""" - from webrtc import RTCDtlsTransport - - return RTCDtlsTransport._wrap_optional(self._native_obj.transport) + def transport(self) -> webrtc.RTCDtlsTransport | None: + """:obj:`webrtc.RTCDtlsTransport`, optional: The transport of the packets, :obj:`None` until there's one.""" + return webrtc.RTCDtlsTransport._wrap_optional(self._native_obj.transport) @property - def jitter_buffer_target(self) -> Optional[float]: - """:obj:`float`, optional: How many milliseconds of media the receiver should buffer (0 to 4000), - trading latency for smoothness. :obj:`None` for the default.""" + def jitter_buffer_target(self) -> float | None: + """:obj:`float`, optional: How many milliseconds of media the receiver should buffer (0 to 4000). + + It trades latency for smoothness. :obj:`None` for the default. + """ return self._native_obj.jitterBufferTarget @jitter_buffer_target.setter - def jitter_buffer_target(self, value: Optional[float]): - if value is not None and not 0 <= value <= 4000: - raise InvalidRangeError(f'jitter_buffer_target must be from 0 to 4000 milliseconds, not {value}') + def jitter_buffer_target(self, value: float | None) -> None: + if value is not None and not 0 <= value <= _MAX_JITTER_BUFFER_TARGET: + msg = f'jitter_buffer_target must be from 0 to 4000 milliseconds, not {value}' + raise InvalidRangeError(msg) self._native_obj.jitterBufferTarget = value - def get_parameters(self) -> 'webrtc.RTCRtpReceiveParameters': + def get_parameters(self) -> webrtc.RTCRtpReceiveParameters: """Returns the parameters the receiver receives with. Returns: @@ -72,7 +70,7 @@ def get_parameters(self) -> 'webrtc.RTCRtpReceiveParameters': return RTCRtpReceiveParameters._from_native(self._native_obj.getParameters()) @staticmethod - def get_capabilities(kind: 'webrtc.MediaType') -> Optional['webrtc.RTCRtpCapabilities']: + def get_capabilities(kind: webrtc.MediaType) -> webrtc.RTCRtpCapabilities | None: """Returns the codecs and header extensions receivers of a kind support. Args: @@ -83,18 +81,18 @@ def get_capabilities(kind: 'webrtc.MediaType') -> Optional['webrtc.RTCRtpCapabil """ return RTCRtpCapabilities._supported(wrtc.RTCRtpReceiver, kind) - async def get_stats(self) -> 'webrtc.RTCStatsReport': + async def get_stats(self) -> webrtc.RTCStatsReport: """Collects the stats of the receiver, and of the objects its stats refer to. Returns: :obj:`webrtc.RTCStatsReport`: The stats. Raises: - :obj:`webrtc.InvalidStateError`: If the connection is closed. + webrtc.InvalidStateError: If the connection is closed. """ return RTCStatsReport._from_native(await call_native(self._native_obj.getStats), [self]) - def get_synchronization_sources(self) -> List['webrtc.RTCRtpSynchronizationSource']: + def get_synchronization_sources(self) -> list[webrtc.RTCRtpSynchronizationSource]: """Returns the synchronization sources (SSRCs) of the media received in the last 10 seconds. Returns: @@ -102,9 +100,10 @@ def get_synchronization_sources(self) -> List['webrtc.RTCRtpSynchronizationSourc """ return self._sources(synchronization=True) - def get_contributing_sources(self) -> List['webrtc.RTCRtpContributingSource']: - """Returns the contributing sources (CSRCs, like the participants mixed by a conference server) of the - media received in the last 10 seconds. + def get_contributing_sources(self) -> list[webrtc.RTCRtpContributingSource]: + """Returns the contributing sources (CSRCs) of the media received in the last 10 seconds. + + Like the participants a conference server mixes. Returns: :obj:`list` of :obj:`webrtc.RTCRtpContributingSource`: The sources, the most recent first. diff --git a/python-webrtc/python/webrtc/interfaces/rtc_rtp_sender.py b/python-webrtc/python/webrtc/interfaces/rtc_rtp_sender.py index 3912368..17d58bd 100644 --- a/python-webrtc/python/webrtc/interfaces/rtc_rtp_sender.py +++ b/python-webrtc/python/webrtc/interfaces/rtc_rtp_sender.py @@ -5,8 +5,13 @@ # that can be found in the LICENSE.md file in the root of the project. # -from typing import TYPE_CHECKING, List, Optional, Sequence +"""RTCRtpSender of WebRTC.""" +from __future__ import annotations + +from typing import TYPE_CHECKING + +import webrtc from webrtc import ( InvalidModificationError, InvalidRangeError, @@ -23,74 +28,38 @@ from webrtc.utils.task_queue import TaskQueue if TYPE_CHECKING: - import webrtc + from collections.abc import Sequence + +class RTCRtpSender(WebRTCObject[wrtc.RTCRtpSender]): + """Sends the media of a track, encoded, to the remote peer, and controls how it's sent. -class RTCRtpSender(WebRTCObject): - """The :obj:`webrtc.RTCRtpSender` interface sends the media of a track, encoded, to the remote peer, and - controls how it's sent. It's the sender of an :obj:`webrtc.RTCRtpTransceiver` of an - :obj:`webrtc.RTCPeerConnection`. + It's the sender of an :obj:`webrtc.RTCRtpTransceiver` of an :obj:`webrtc.RTCPeerConnection`. """ _class = wrtc.RTCRtpSender - def _check_parameters( - self, - parameters: 'webrtc.RTCRtpSendParameters', - last: 'wrtc.RtpParameters', - key_frames: Optional[Sequence[bool]], - ) -> None: - """Checks the parameters set_parameters() takes against the ones get_parameters() returned last, as the - specification requires, before libwebrtc gets them.""" - returned = RTCRtpSendParameters._from_native(last) - if parameters.transaction_id != returned.transaction_id: - raise InvalidModificationError("The transaction_id doesn't match the one of the last get_parameters()") - for name in ('codecs', 'header_extensions', 'rtcp'): - if getattr(parameters, name) != getattr(returned, name): - raise InvalidModificationError(f'{name} of the parameters can not be changed') - if [e.rid for e in parameters.encodings] != [e.rid for e in returned.encodings]: - raise InvalidModificationError('The number of encodings and their rid can not be changed') - if key_frames is not None and len(key_frames) != len(parameters.encodings): - raise InvalidModificationError('key_frames must have one value per encoding') - if self.kind != MediaType.video: - return - for encoding in parameters.encodings: - if encoding.scale_resolution_down_by is not None and encoding.scale_resolution_down_by < 1: - raise InvalidRangeError('scale_resolution_down_by must be at least 1') - if encoding.max_framerate is not None and encoding.max_framerate < 0: - raise InvalidRangeError('max_framerate must not be negative') - @property - def track(self) -> Optional['webrtc.MediaStreamTrack']: - """:obj:`webrtc.MediaStreamTrack`, optional: The :obj:`webrtc.MediaStreamTrack` which is being handled by - the :obj:`webrtc.RTCRtpSender`. If track is :obj:`None`, the :obj:`webrtc.RTCRtpSender` - doesn't transmit anything.""" - from webrtc import MediaStreamTrack - - return MediaStreamTrack._wrap_optional(self._native_obj.track) + def track(self) -> webrtc.MediaStreamTrack | None: + """:obj:`webrtc.MediaStreamTrack`, optional: The track the sender sends, :obj:`None` to send nothing.""" + return webrtc.MediaStreamTrack._wrap_optional(self._native_obj.track) @property - def transport(self) -> Optional['webrtc.RTCDtlsTransport']: - """:obj:`webrtc.RTCDtlsTransport`, optional: An object representing the underlying transport being used by - the sender to exchange packets with the remote peer, or :obj:`None` if the sender isn't yet connected - to transport.""" - from webrtc import RTCDtlsTransport - - return RTCDtlsTransport._wrap_optional(self._native_obj.transport) + def transport(self) -> webrtc.RTCDtlsTransport | None: + """:obj:`webrtc.RTCDtlsTransport`, optional: The transport of the packets, :obj:`None` until there's one.""" + return webrtc.RTCDtlsTransport._wrap_optional(self._native_obj.transport) @property - def dtmf(self) -> Optional['webrtc.RTCDTMFSender']: + def dtmf(self) -> webrtc.RTCDTMFSender | None: """:obj:`webrtc.RTCDTMFSender`, optional: Sends DTMF tones, for an audio sender.""" - from webrtc import RTCDTMFSender - - return RTCDTMFSender._wrap_optional(self._native_obj.dtmf) + return webrtc.RTCDTMFSender._wrap_optional(self._native_obj.dtmf) @property - def kind(self) -> 'webrtc.MediaType': + def kind(self) -> webrtc.MediaType: """:obj:`webrtc.MediaType`: The kind of media the sender sends, audio or video.""" return self._native_obj.kind - def get_parameters(self) -> 'webrtc.RTCRtpSendParameters': + def get_parameters(self) -> webrtc.RTCRtpSendParameters: """Returns the parameters the sender sends with. To change them, modify the returned parameters and pass them to :meth:`set_parameters` before the current @@ -107,7 +76,7 @@ def get_parameters(self) -> 'webrtc.RTCRtpSendParameters': return parameters async def set_parameters( - self, parameters: 'webrtc.RTCRtpSendParameters', *, key_frames: Optional[Sequence[bool]] = None + self, parameters: webrtc.RTCRtpSendParameters, *, key_frames: Sequence[bool] | None = None ) -> None: """Changes how the sender sends: its encodings and degradation preference. @@ -118,18 +87,22 @@ async def set_parameters( right away. Raises: - :obj:`webrtc.InvalidStateError`: If :meth:`get_parameters` wasn't called in the current task. - :obj:`webrtc.InvalidModificationError`: If the ``transaction_id``, the codecs, the header extensions, + webrtc.InvalidStateError: If :meth:`get_parameters` wasn't called in the current task. + webrtc.InvalidModificationError: If the ``transaction_id``, the codecs, the header extensions, the RTCP parameters, the number of encodings or their ``rid`` changed, the codec of an encoding isn't negotiated, or ``key_frames`` isn't one per encoding. - :obj:`webrtc.InvalidRangeError`: If a value is out of range, like ``scale_resolution_down_by`` below 1. + webrtc.InvalidRangeError: If a value is out of range, like ``scale_resolution_down_by`` below 1. """ if self._native_obj._transceiverStopped(): - raise InvalidStateError('The transceiver of the sender is stopped') + msg = 'The transceiver of the sender is stopped' + raise InvalidStateError(msg) last = self._native_obj._lastParameters() if last is None: - raise InvalidStateError('get_parameters() must be called before set_parameters(), in the same task') - self._check_parameters(parameters, last, key_frames) + msg = 'get_parameters() must be called before set_parameters(), in the same task' + raise InvalidStateError(msg) + _check_unchanged(parameters, RTCRtpSendParameters._from_native(last), key_frames) + if self.kind == MediaType.video: + _check_video_ranges(parameters.encodings) kind = self.kind # a copy (pybind returns one): changed, then set back @@ -142,7 +115,7 @@ async def set_parameters( last.degradationPreference = parameters.degradation_preference await call_native(self._native_obj.setParameters, last) - async def replace_track(self, track: Optional['webrtc.MediaStreamTrack']) -> None: + async def replace_track(self, track: webrtc.MediaStreamTrack | None) -> None: """Replaces the track the sender sends, without negotiation. The track is replaced in the operations chain of the connection, after the operations started before @@ -153,30 +126,30 @@ async def replace_track(self, track: Optional['webrtc.MediaStreamTrack']) -> Non to stop sending. Raises: - :obj:`TypeError`: If the track is of another kind. - :obj:`webrtc.InvalidStateError`: If the transceiver of the sender is stopped, or the connection closed. + TypeError: If the track is of another kind. + webrtc.InvalidStateError: If the transceiver of the sender is stopped, or the connection closed. """ - from webrtc import RTCPeerConnection - if track is not None and track.kind != self.kind: - raise TypeError(f'a {track.kind} track can not replace the track of a {self.kind} sender') + msg = f'a {track.kind} track can not replace the track of a {self.kind} sender' + raise TypeError(msg) - def replace(): + def replace() -> None: native_track = track._native_obj if track is not None else None if self._native_obj._transceiverStopped() or not self._native_obj.replaceTrack(native_track): - raise InvalidStateError('The track of a stopped sender can not be replaced') + msg = 'The track of a stopped sender can not be replaced' + raise InvalidStateError(msg) connection = wrtc.RTCPeerConnection._connectionOf(self._native_obj) if connection is None: replace() return - pc = RTCPeerConnection._wrap(connection) + pc = webrtc.RTCPeerConnection._wrap(connection) async with pc._operation(): pc._check_state('replace the track') await later() replace() - def set_streams(self, *streams: 'webrtc.MediaStream') -> None: + def set_streams(self, *streams: webrtc.MediaStream) -> None: """Sets the streams the remote peer associates the track of the sender with, from the next negotiation. Args: @@ -189,7 +162,7 @@ def set_streams(self, *streams: 'webrtc.MediaStream') -> None: self._native_obj.setStreams(ids) @staticmethod - def get_capabilities(kind: 'webrtc.MediaType') -> Optional['webrtc.RTCRtpCapabilities']: + def get_capabilities(kind: webrtc.MediaType) -> webrtc.RTCRtpCapabilities | None: """Returns the codecs and header extensions senders of a kind support. Args: @@ -200,14 +173,14 @@ def get_capabilities(kind: 'webrtc.MediaType') -> Optional['webrtc.RTCRtpCapabil """ return RTCRtpCapabilities._supported(wrtc.RTCRtpSender, kind) - async def get_stats(self) -> 'webrtc.RTCStatsReport': + async def get_stats(self) -> webrtc.RTCStatsReport: """Collects the stats of the sender, and of the objects its stats refer to. Returns: :obj:`webrtc.RTCStatsReport`: The stats. Raises: - :obj:`webrtc.InvalidStateError`: If the connection is closed. + webrtc.InvalidStateError: If the connection is closed. """ return RTCStatsReport._from_native(await call_native(self._native_obj.getStats)) @@ -225,7 +198,47 @@ async def get_stats(self) -> 'webrtc.RTCStatsReport': getCapabilities = get_capabilities -def _default_scale_resolution_down_by(encodings: List['webrtc.RTCRtpEncodingParameters']) -> None: +def _check_unchanged( + parameters: webrtc.RTCRtpSendParameters, + returned: webrtc.RTCRtpSendParameters, + key_frames: Sequence[bool] | None, +) -> None: + """Checks what set_parameters() can't change against the parameters get_parameters() returned last. + + Raises: + webrtc.InvalidModificationError: If something changed, or ``key_frames`` isn't one per encoding. + """ + if parameters.transaction_id != returned.transaction_id: + msg = "The transaction_id doesn't match the one of the last get_parameters()" + raise InvalidModificationError(msg) + for name in ('codecs', 'header_extensions', 'rtcp'): + if getattr(parameters, name) != getattr(returned, name): + msg = f'{name} of the parameters can not be changed' + raise InvalidModificationError(msg) + if [e.rid for e in parameters.encodings] != [e.rid for e in returned.encodings]: + msg = 'The number of encodings and their rid can not be changed' + raise InvalidModificationError(msg) + if key_frames is not None and len(key_frames) != len(parameters.encodings): + msg = 'key_frames must have one value per encoding' + raise InvalidModificationError(msg) + + +def _check_video_ranges(encodings: list[webrtc.RTCRtpEncodingParameters]) -> None: + """Checks the values of video encodings, before libwebrtc gets them. + + Raises: + webrtc.InvalidRangeError: If a value is out of range. + """ + for encoding in encodings: + if encoding.scale_resolution_down_by is not None and encoding.scale_resolution_down_by < 1: + msg = 'scale_resolution_down_by must be at least 1' + raise InvalidRangeError(msg) + if encoding.max_framerate is not None and encoding.max_framerate < 0: + msg = 'max_framerate must not be negative' + raise InvalidRangeError(msg) + + +def _default_scale_resolution_down_by(encodings: list[webrtc.RTCRtpEncodingParameters]) -> None: """Video encodings without scale_resolution_down_by scale by 1, or by descending powers of 2 if none has one.""" if all(e.scale_resolution_down_by is None for e in encodings): for i, encoding in enumerate(encodings): diff --git a/python-webrtc/python/webrtc/interfaces/rtc_rtp_transceiver.py b/python-webrtc/python/webrtc/interfaces/rtc_rtp_transceiver.py index 67b9424..8ff4b06 100644 --- a/python-webrtc/python/webrtc/interfaces/rtc_rtp_transceiver.py +++ b/python-webrtc/python/webrtc/interfaces/rtc_rtp_transceiver.py @@ -5,49 +5,41 @@ # that can be found in the LICENSE.md file in the root of the project. # -from typing import TYPE_CHECKING, List, Optional +"""RTCRtpTransceiver of WebRTC.""" -from webrtc import InvalidModificationError, RTCRtpCodec, RTCRtpHeaderExtensionCapability, WebRTCObject, wrtc +from __future__ import annotations -if TYPE_CHECKING: - import webrtc +import webrtc +from webrtc import InvalidModificationError, RTCRtpCodec, RTCRtpHeaderExtensionCapability, WebRTCObject, wrtc -class RTCRtpTransceiver(WebRTCObject): - """The WebRTC interface :obj:`webrtc.RTCRtpTransceiver` describes a permanent pairing of - an :obj:`webrtc.RTCRtpSender` and an :obj:`webrtc.RTCRtpReceiver`, along with some shared state. - """ +class RTCRtpTransceiver(WebRTCObject[wrtc.RTCRtpTransceiver]): + """A permanent pair of an :obj:`webrtc.RTCRtpSender` and an :obj:`webrtc.RTCRtpReceiver`, with shared state.""" _class = wrtc.RTCRtpTransceiver @property - def mid(self) -> Optional[str]: + def mid(self) -> str | None: """A :obj:`str` which uniquely identifies the pairing of source and destination of the transceiver's stream. - Its value is taken from the media ID of the SDP m-line. This value is :obj:`None` if negotiation - has not completed.""" + + Its value is taken from the media ID of the SDP m-line. This value is :obj:`None` if negotiation has not + completed. + """ return self._native_obj.mid @property - def receiver(self) -> 'webrtc.RTCRtpReceiver': - """A :obj:`webrtc.RTCRtpReceiver` object which is responsible for receiving and decoding incoming media data - whose media ID is the same as the current value of :attr:`mid`.""" - from webrtc import RTCRtpReceiver - - return RTCRtpReceiver._wrap(self._native_obj.receiver) + def receiver(self) -> webrtc.RTCRtpReceiver: + """:obj:`webrtc.RTCRtpReceiver`: Receives and decodes the incoming media of the :attr:`mid`.""" + return webrtc.RTCRtpReceiver._wrap(self._native_obj.receiver) @property - def sender(self) -> 'webrtc.RTCRtpSender': - """A :obj:`webrtc.RTCRtpSender` object used to encode and send media whose media ID matches - the current value of :attr:`mid`.""" - from webrtc import RTCRtpSender - - return RTCRtpSender._wrap(self._native_obj.sender) + def sender(self) -> webrtc.RTCRtpSender: + """:obj:`webrtc.RTCRtpSender`: Encodes and sends the media of the :attr:`mid`.""" + return webrtc.RTCRtpSender._wrap(self._native_obj.sender) @property def stopped(self) -> bool: - """A :obj:`bool` value which is :obj:`True` if the transceiver's :attr:`sender` will no longer send data, - and its :attr:`receiver` will no longer receive data. If either or both are still at work, - the result is :obj:`False`. + """:obj:`bool`: Whether both the :attr:`sender` and the :attr:`receiver` stopped for good. Warning: Deprecated: This feature is no longer recommended. @@ -55,7 +47,7 @@ def stopped(self) -> bool: return self._native_obj.stopped @property - def direction(self) -> 'webrtc.TransceiverDirection': + def direction(self) -> webrtc.TransceiverDirection: """A member of :obj:`webrtc.TransceiverDirection` enum, indicating the transceiver's preferred direction. Note: @@ -64,18 +56,16 @@ def direction(self) -> 'webrtc.TransceiverDirection': return self._native_obj.direction @direction.setter - def direction(self, new_direction: 'webrtc.TransceiverDirection'): + def direction(self, new_direction: webrtc.TransceiverDirection) -> None: self._native_obj.direction = new_direction @property - def current_direction(self) -> Optional['webrtc.TransceiverDirection']: - """A member of :obj:`webrtc.TransceiverDirection` enum, indicating - the current directionality of the transceiver.""" + def current_direction(self) -> webrtc.TransceiverDirection | None: + """:obj:`webrtc.TransceiverDirection`, optional: The negotiated direction of the transceiver.""" return self._native_obj.currentDirection def stop(self) -> None: - """Permanently stops the transceiver by stopping both the associated :obj:`webrtc.RTCRtpSender` - and :obj:`webrtc.RTCRtpReceiver`. + """Stops the transceiver for good, its :obj:`webrtc.RTCRtpSender` and its :obj:`webrtc.RTCRtpReceiver`. Note: To check whether the transceiver is stopped, compare :attr:`currentDirection` with @@ -84,11 +74,11 @@ def stop(self) -> None: self._native_obj.stop() @property - def kind(self) -> 'webrtc.MediaType': + def kind(self) -> webrtc.MediaType: """:obj:`webrtc.MediaType`: The kind of media the transceiver sends and receives, audio or video.""" return self._native_obj.kind - def set_codec_preferences(self, codecs: List['webrtc.RTCRtpCodec']) -> None: + def set_codec_preferences(self, codecs: list[webrtc.RTCRtpCodec]) -> None: """Sets the codecs to negotiate, in order of preference, from the next negotiation. Args: @@ -97,7 +87,7 @@ def set_codec_preferences(self, codecs: List['webrtc.RTCRtpCodec']) -> None: for the kind of the transceiver. An empty list restores the default preferences. Raises: - :obj:`webrtc.InvalidModificationError`: If a codec isn't supported, or only resiliency codecs + webrtc.InvalidModificationError: If a codec isn't supported, or only resiliency codecs (like RTX or FEC) are given. """ kind = self.kind @@ -109,11 +99,12 @@ def set_codec_preferences(self, codecs: List['webrtc.RTCRtpCodec']) -> None: for codec in codecs: native = next((n for n in natives if RTCRtpCodec._from_native(n)._matches(codec)), None) if native is None: - raise InvalidModificationError(f'{codec.mime_type} is not a {kind} codec that can be negotiated') + msg = f'{codec.mime_type} is not a {kind} codec that can be negotiated' + raise InvalidModificationError(msg) preferences.append(native) self._native_obj.setCodecPreferences(preferences) - def get_header_extensions_to_negotiate(self) -> List['webrtc.RTCRtpHeaderExtensionCapability']: + def get_header_extensions_to_negotiate(self) -> list[webrtc.RTCRtpHeaderExtensionCapability]: """Returns the header extensions offered or accepted in the next negotiation. Returns: @@ -124,7 +115,7 @@ def get_header_extensions_to_negotiate(self) -> List['webrtc.RTCRtpHeaderExtensi RTCRtpHeaderExtensionCapability._from_native(e) for e in self._native_obj.getHeaderExtensionsToNegotiate() ] - def set_header_extensions_to_negotiate(self, extensions: List['webrtc.RTCRtpHeaderExtensionCapability']) -> None: + def set_header_extensions_to_negotiate(self, extensions: list[webrtc.RTCRtpHeaderExtensionCapability]) -> None: """Changes the directions the header extensions are negotiated in, from the next negotiation. Args: @@ -132,15 +123,16 @@ def set_header_extensions_to_negotiate(self, extensions: List['webrtc.RTCRtpHead :meth:`get_header_extensions_to_negotiate` returns, with directions changed. Raises: - :obj:`ValueError`: If an extension has an empty URI. - :obj:`webrtc.InvalidModificationError`: If the extensions or their order differ, or a mandatory + ValueError: If an extension has an empty URI. + webrtc.InvalidModificationError: If the extensions or their order differ, or a mandatory extension is stopped. """ current = {e.uri: e for e in self._native_obj.getHeaderExtensionsToNegotiate()} natives = [] for extension in extensions: if not extension.uri: - raise ValueError('the URI of a header extension must not be empty') + msg = 'the URI of a header extension must not be empty' + raise ValueError(msg) native = wrtc.RtpHeaderExtensionCapability() native.uri = extension.uri if extension.uri in current: @@ -149,7 +141,7 @@ def set_header_extensions_to_negotiate(self, extensions: List['webrtc.RTCRtpHead natives.append(native) self._native_obj.setHeaderExtensionsToNegotiate(natives) - def get_negotiated_header_extensions(self) -> List['webrtc.RTCRtpHeaderExtensionCapability']: + def get_negotiated_header_extensions(self) -> list[webrtc.RTCRtpHeaderExtensionCapability]: """Returns the header extensions negotiated last, and their directions. Returns: diff --git a/python-webrtc/python/webrtc/interfaces/rtc_sctp_transport.py b/python-webrtc/python/webrtc/interfaces/rtc_sctp_transport.py index 610c00f..31fd59d 100644 --- a/python-webrtc/python/webrtc/interfaces/rtc_sctp_transport.py +++ b/python-webrtc/python/webrtc/interfaces/rtc_sctp_transport.py @@ -5,20 +5,20 @@ # that can be found in the LICENSE.md file in the root of the project. # -from typing import TYPE_CHECKING, Optional +"""RTCSctpTransport of WebRTC.""" +from __future__ import annotations + +import webrtc from webrtc import WebRTCObject, wrtc from webrtc.utils.events import EventTarget -if TYPE_CHECKING: - import webrtc +class RTCSctpTransport(WebRTCObject[wrtc.RTCSctpTransport], EventTarget): + """The Stream Control Transmission Protocol (SCTP) transport of a :obj:`webrtc.RTCPeerConnection`. -class RTCSctpTransport(WebRTCObject, EventTarget): - """The :obj:`webrtc.RTCSctpTransport` interface provides information which describes a Stream Control - Transmission Protocol (SCTP) transport. This provides information about limitations of the transport, - but also provides a way to access the underlying Datagram Transport Layer Security (DTLS) transport over - which SCTP packets for all of an :obj:`webrtc.RTCPeerConnection`'s data channels are sent and received. + It tells the limitations of the transport, and gives the Datagram Transport Layer Security (DTLS) transport + over which the SCTP packets of all the data channels of the connection are sent and received. Events (see :meth:`on`): ``statechange`` (:obj:`webrtc.Event`): :attr:`state` changed. @@ -27,34 +27,29 @@ class RTCSctpTransport(WebRTCObject, EventTarget): _class = wrtc.RTCSctpTransport _events = ('statechange',) - def _on_event(self, name: str, *args): + def _on_event(self, _name: str, *args: object) -> None: (state,) = args # the state changes along with its event self._native_obj._surfaceState(state) @property - def transport(self) -> 'webrtc.RTCDtlsTransport': - """:obj:`webrtc.RTCDtlsTransport`: An object representing the DTLS transport used for the transmission - and receipt of data packets.""" - from webrtc import RTCDtlsTransport - - return RTCDtlsTransport._wrap(self._native_obj.transport) + def transport(self) -> webrtc.RTCDtlsTransport: + """:obj:`webrtc.RTCDtlsTransport`: The DTLS transport the data packets are sent and received over.""" + return webrtc.RTCDtlsTransport._wrap(self._native_obj.transport) @property - def state(self) -> 'webrtc.SctpTransportState': + def state(self) -> webrtc.SctpTransportState: """:obj:`webrtc.SctpTransportState`: An enumerated value indicating the state of the SCTP transport.""" return self._native_obj.state @property - def max_message_size(self) -> Optional[float]: - """:obj:`float`, optional: The maximum size, in bytes, of a message which can be sent using the - :meth:`webrtc.RTCDataChannel.send` method.""" + def max_message_size(self) -> float | None: + """:obj:`float`, optional: The maximum size in bytes of a message :meth:`webrtc.RTCDataChannel.send` sends.""" return self._native_obj.maxMessageSize @property - def max_channels(self) -> Optional[int]: - """:obj:`int`, optional: An integer value indicating the maximum number of :obj:`webrtc.RTCDataChannel` that - can be open simultaneously.""" + def max_channels(self) -> int | None: + """:obj:`int`, optional: The maximum number of :obj:`webrtc.RTCDataChannel` open at the same time.""" return self._native_obj.maxChannels #: Alias for :attr:`max_message_size` diff --git a/python-webrtc/python/webrtc/interfaces/track_generator.py b/python-webrtc/python/webrtc/interfaces/track_generator.py index f0292f5..c41fa5d 100644 --- a/python-webrtc/python/webrtc/interfaces/track_generator.py +++ b/python-webrtc/python/webrtc/interfaces/track_generator.py @@ -7,39 +7,48 @@ """Tracks of media the application writes: VideoTrackGenerator, and MediaStreamTrackGenerator of Chrome.""" +from __future__ import annotations + from dataclasses import dataclass -from typing import Any, Union +from typing import TYPE_CHECKING from webrtc import AudioData, AudioSampleFormat, MediaStreamTrack, MediaType, VideoFrame, wrtc from webrtc.exceptions import NotSupportedError from webrtc.streams import WritableStream +if TYPE_CHECKING: + from webrtc.streams import WritableStreamDefaultController + class _TrackSink: - """The underlying sink of a generator's writable stream: sends each chunk on the track, closing it""" + """The underlying sink of a generator's writable stream: sends each chunk on the track, closing it.""" - def __init__(self, native: 'wrtc.TrackGenerator'): + def __init__(self, native: wrtc.TrackGenerator) -> None: self._native = native - def write(self, chunk: Any, controller) -> None: + def write(self, chunk: object, _controller: WritableStreamDefaultController) -> None: if self._native.kind == 'video': self._write_video(chunk) else: self._write_audio(chunk) - def _write_video(self, frame: Any) -> None: + def _write_video(self, frame: object) -> None: if not isinstance(frame, VideoFrame): - raise TypeError(f'A video generator takes VideoFrame, not {type(frame).__name__}') + msg = f'A video generator takes VideoFrame, not {type(frame).__name__}' + raise TypeError(msg) if frame._resource is None: - raise TypeError('The frame is closed') + msg = 'The frame is closed' + raise TypeError(msg) timestamp, rotation = frame.timestamp, frame.rotation self._native.writeVideo(frame._take_resource(), timestamp, rotation) - def _write_audio(self, data: Any) -> None: + def _write_audio(self, data: object) -> None: if not isinstance(data, AudioData): - raise TypeError(f'An audio generator takes AudioData, not {type(data).__name__}') + msg = f'An audio generator takes AudioData, not {type(data).__name__}' + raise TypeError(msg) if data._data is None: - raise TypeError('The data is closed') + msg = 'The data is closed' + raise TypeError(msg) audio = data._take() if audio.format == AudioSampleFormat.s16: samples = audio._data @@ -60,13 +69,14 @@ def close(self) -> None: # ends the tracks of the generator self._native.close() - def abort(self, reason: Any) -> None: + def abort(self, _reason: object) -> None: self._native.close() class VideoTrackGenerator: - """A video track of the frames written to a stream - (https://developer.mozilla.org/en-US/docs/Web/API/VideoTrackGenerator). + """A video track of the frames written to a stream. + + See https://developer.mozilla.org/en-US/docs/Web/API/VideoTrackGenerator. Each frame written is sent on :attr:`track` and closed. Closing or aborting :attr:`writable` ends the track. @@ -78,7 +88,7 @@ class VideoTrackGenerator: await writer.write(webrtc.VideoFrame(i420, format='I420', coded_width=640, coded_height=480, timestamp=0)) """ - def __init__(self): + def __init__(self) -> None: self._native = wrtc.TrackGenerator('video') # the native generator doesn't keep the track, Python does self._track = MediaStreamTrack._wrap(self._native.track) @@ -100,7 +110,7 @@ def muted(self) -> bool: return self._native.muted @muted.setter - def muted(self, value: bool): + def muted(self, value: bool) -> None: self._native.muted = bool(value) @@ -116,8 +126,9 @@ class MediaStreamTrackGeneratorInit: class MediaStreamTrackGenerator(MediaStreamTrack): - """A track of the media written to a stream, :obj:`webrtc.VideoFrame` or :obj:`webrtc.AudioData` objects. It's - Chrome's API (https://developer.mozilla.org/en-US/docs/Web/API/MediaStreamTrackGenerator), the only one for + """A track of the media written to a stream, :obj:`webrtc.VideoFrame` or :obj:`webrtc.AudioData` objects. + + It's Chrome's API (https://developer.mozilla.org/en-US/docs/Web/API/MediaStreamTrackGenerator), the only one for audio: for video, :obj:`VideoTrackGenerator` is the standard one. Audio is sent in 10 ms frames: samples short of one wait for the next ones written. @@ -127,16 +138,17 @@ class MediaStreamTrackGenerator(MediaStreamTrack): or the init with it. A dictionary of the init's members is taken too. Raises: - :obj:`TypeError`: If the kind isn't audio or video. + TypeError: If the kind isn't audio or video. """ - def __init__(self, kind: Union[str, MediaType, MediaStreamTrackGeneratorInit, dict]): + def __init__(self, kind: str | MediaType | MediaStreamTrackGeneratorInit | dict[str, str]) -> None: if isinstance(kind, MediaStreamTrackGeneratorInit): kind = kind.kind elif isinstance(kind, dict): kind = kind.get('kind') - if kind not in ('audio', 'video'): - raise TypeError(f"The kind must be 'audio' or 'video', not {kind!r}") + if kind not in {'audio', 'video'}: + msg = f"The kind must be 'audio' or 'video', not {kind!r}" + raise TypeError(msg) self._generator = wrtc.TrackGenerator(MediaType(kind).value) super().__init__(self._generator.track) self._attach() diff --git a/python-webrtc/python/webrtc/models/audio_data.py b/python-webrtc/python/webrtc/models/audio_data.py index 1860364..426df19 100644 --- a/python-webrtc/python/webrtc/models/audio_data.py +++ b/python-webrtc/python/webrtc/models/audio_data.py @@ -7,22 +7,26 @@ """AudioData of WebCodecs (https://developer.mozilla.org/en-US/docs/Web/API/AudioData) and its dictionaries.""" +from __future__ import annotations + import math import warnings from dataclasses import dataclass -from typing import Any, NamedTuple, Optional, Union +from typing import Any, ClassVar, NamedTuple from webrtc import AudioSampleFormat, InvalidRangeError, InvalidStateError, NotSupportedError, wrtc -from webrtc.utils.names import alias, snake_case +from webrtc.models.closable import Closable +from webrtc.utils.names import Alias, alias, snake_case _SAMPLE_BYTES = {'u8': 1, 's16': 2, 's32': 4, 'f32': 4} -def _sample_format(value: Any) -> AudioSampleFormat: +def _sample_format(value: object) -> AudioSampleFormat: try: return AudioSampleFormat(value) except ValueError: - raise TypeError(f'{value!r} is not an AudioSampleFormat') from None + msg = f'{value!r} is not an AudioSampleFormat' + raise TypeError(msg) from None def _sample_bytes(format: AudioSampleFormat) -> int: @@ -54,11 +58,11 @@ class AudioDataInit: data: Any #: Alias for :attr:`sample_rate` - sampleRate = alias('sample_rate') + sampleRate: ClassVar[Alias[float]] = alias('sample_rate') #: Alias for :attr:`number_of_frames` - numberOfFrames = alias('number_of_frames') + numberOfFrames: ClassVar[Alias[int]] = alias('number_of_frames') #: Alias for :attr:`number_of_channels` - numberOfChannels = alias('number_of_channels') + numberOfChannels: ClassVar[Alias[int]] = alias('number_of_channels') @dataclass @@ -74,37 +78,87 @@ class AudioDataCopyToOptions: plane_index: int frame_offset: int = 0 - frame_count: Optional[int] = None - format: Optional[AudioSampleFormat] = None + frame_count: int | None = None + format: AudioSampleFormat | None = None #: Alias for :attr:`plane_index` - planeIndex = alias('plane_index') + planeIndex: ClassVar[Alias[int]] = alias('plane_index') #: Alias for :attr:`frame_offset` - frameOffset = alias('frame_offset') + frameOffset: ClassVar[Alias[int]] = alias('frame_offset') #: Alias for :attr:`frame_count` - frameCount = alias('frame_count') + frameCount: ClassVar[Alias[int | None]] = alias('frame_count') -def _copy_options(value: Any) -> AudioDataCopyToOptions: +def _copy_options(value: object) -> AudioDataCopyToOptions: if isinstance(value, AudioDataCopyToOptions): return value if isinstance(value, dict): kwargs = {snake_case(k): v for k, v in value.items() if v is not None} if 'plane_index' not in kwargs: - raise TypeError('plane_index is required') + msg = 'plane_index is required' + raise TypeError(msg) try: return AudioDataCopyToOptions(**kwargs) except TypeError as e: - raise TypeError(f'Invalid AudioDataCopyToOptions: {e}') from None - raise TypeError(f'{value!r} is not an AudioDataCopyToOptions') + msg = f'Invalid AudioDataCopyToOptions: {e}' + raise TypeError(msg) from None + msg = f'{value!r} is not an AudioDataCopyToOptions' + raise TypeError(msg) -def _unsigned(value: Any, name: str) -> int: +def _unsigned(value: object, name: str) -> int: if isinstance(value, bool) or not isinstance(value, int) or value < 0: - raise TypeError(f'{name} must be a non-negative integer, not {value!r}') + msg = f'{name} must be a non-negative integer, not {value!r}' + raise TypeError(msg) return value +def _audio_data_init(init: object, options: dict[str, object]) -> AudioDataInit: + if init is None: + init = options + elif options: + msg = 'Pass either an init or keyword arguments' + raise TypeError(msg) + if isinstance(init, dict): + kwargs = {snake_case(k): v for k, v in init.items() if k != 'transfer'} + try: + init = AudioDataInit(**kwargs) + except TypeError as e: + msg = f'Invalid AudioDataInit: {e}' + raise TypeError(msg) from None + if not isinstance(init, AudioDataInit): + msg = f'{init!r} is not an AudioDataInit' + raise TypeError(msg) + return init + + +def _sample_rate(value: object) -> float: + if isinstance(value, bool) or not isinstance(value, (int, float)) or not 0 < value < math.inf: + msg = 'sample_rate must be positive and finite' + raise TypeError(msg) + return float(value) + + +def _buffer(data: object, size: int) -> bytes: + """The first bytes of a bytes-like buffer.""" + try: + view = memoryview(data).cast('B') + except TypeError: + msg = 'data must be a bytes-like buffer' + raise TypeError(msg) from None + if view.nbytes < size: + msg = f'data must be at least {size} bytes for this format and size' + raise TypeError(msg) + return bytes(view[:size]) + + +class _Layout(NamedTuple): + format: AudioSampleFormat + sample_rate: float + frames: int + channels: int + + class _CopyPlan(NamedTuple): format: AudioSampleFormat plane_index: int @@ -113,7 +167,7 @@ class _CopyPlan(NamedTuple): size: int -class AudioData: +class AudioData(Closable): """Audio samples and their metadata (https://developer.mozilla.org/en-US/docs/Web/API/AudioData). Samples read from a track hold memory until :meth:`close`, like a :obj:`webrtc.VideoFrame`. @@ -123,82 +177,60 @@ class AudioData: arguments, can be passed instead. Raises: - :obj:`TypeError`: If the init isn't valid, or the data is too small for it. + TypeError: If the init isn't valid, or the data is too small for it. Example:: - data = webrtc.AudioData(format='s16', sample_rate=48000, number_of_frames=480, number_of_channels=1, - timestamp=0, data=bytes(960)) + data = webrtc.AudioData( + format='s16', sample_rate=48000, number_of_frames=480, number_of_channels=1, timestamp=0, data=bytes(960) + ) """ - def __init__(self, init: Any = None, **options): - if init is None: - init = options - elif options: - raise TypeError('Pass either an init or keyword arguments') - if isinstance(init, dict): - kwargs = {snake_case(k): v for k, v in init.items() if k != 'transfer'} - try: - init = AudioDataInit(**kwargs) - except TypeError as e: - raise TypeError(f'Invalid AudioDataInit: {e}') from None - if not isinstance(init, AudioDataInit): - raise TypeError(f'{init!r} is not an AudioDataInit') - + def __init__(self, init: AudioDataInit | dict[str, object] | None = None, **options: object) -> None: + init = _audio_data_init(init, options) format = _sample_format(init.format) - sample_rate = init.sample_rate - if isinstance(sample_rate, bool) or not isinstance(sample_rate, (int, float)) or not 0 < sample_rate < math.inf: - raise TypeError('sample_rate must be positive and finite') + sample_rate = _sample_rate(init.sample_rate) frames = _unsigned(init.number_of_frames, 'number_of_frames') channels = _unsigned(init.number_of_channels, 'number_of_channels') if frames == 0 or channels == 0: - raise TypeError('number_of_frames and number_of_channels must be positive') + msg = 'number_of_frames and number_of_channels must be positive' + raise TypeError(msg) if isinstance(init.timestamp, bool) or not isinstance(init.timestamp, int): - raise TypeError('The timestamp is an integer of microseconds') - try: - view = memoryview(init.data).cast('B') - except TypeError: - raise TypeError('data must be a bytes-like buffer') from None - size = frames * channels * _sample_bytes(format) - if view.nbytes < size: - raise TypeError(f'data must be at least {size} bytes for this format and size') - self._set(bytes(view[:size]), format, float(sample_rate), frames, channels, init.timestamp) - - def _set( - self, data: bytes, format: AudioSampleFormat, sample_rate: float, frames: int, channels: int, timestamp: int - ) -> None: - self._data: Optional[bytes] = data - self._format = format - self._sample_rate = sample_rate - self._frames = frames - self._channels = channels + msg = 'The timestamp is an integer of microseconds' + raise TypeError(msg) + data = _buffer(init.data, frames * channels * _sample_bytes(format)) + self._set(data, _Layout(format, sample_rate, frames, channels), init.timestamp) + + def _set(self, data: bytes, layout: _Layout, timestamp: int) -> None: + self._data: bytes | None = data + self._format, self._sample_rate, self._frames, self._channels = layout self._timestamp = timestamp @classmethod - def _from_native( - cls, data: bytes, bits_per_sample: int, sample_rate: int, channels: int, frames: int, timestamp: int - ) -> 'AudioData': - """Samples of a track, interleaved""" + def _from_native(cls, native: tuple[bytes, int, int, int, int, int]) -> AudioData: + """Samples of a track, interleaved: the data, bits per sample, sample rate, channels, frames and timestamp.""" + data, bits_per_sample, sample_rate, channels, frames, timestamp = native audio = cls.__new__(cls) format = {8: AudioSampleFormat.u8, 16: AudioSampleFormat.s16, 32: AudioSampleFormat.s32}[bits_per_sample] - audio._set(data, format, float(sample_rate), frames, channels, timestamp) + audio._set(data, _Layout(format, float(sample_rate), frames, channels), timestamp) audio._warn_unclosed = True return audio - def _take(self) -> 'AudioData': - """The samples, for a generator, which closes the data""" + def _take(self) -> AudioData: + """The samples, for a generator, which closes the data.""" if self._data is None: - raise InvalidStateError('The data is closed') + msg = 'The data is closed' + raise InvalidStateError(msg) copy = self.clone() self.close() return copy - def __del__(self): + def __del__(self) -> None: if getattr(self, '_data', None) is not None and getattr(self, '_warn_unclosed', False): warnings.warn('An AudioData was garbage collected without being closed', ResourceWarning, stacklevel=2) @property - def format(self) -> Optional[AudioSampleFormat]: + def format(self) -> AudioSampleFormat | None: """:obj:`webrtc.AudioSampleFormat`, optional: The type and layout of the samples, :obj:`None` once closed.""" return self._format if self._data is not None else None @@ -229,57 +261,62 @@ def timestamp(self) -> int: """:obj:`int`: The presentation time in microseconds.""" return self._timestamp - def _plan_copy(self, options: Any) -> _CopyPlan: + def _plan_copy(self, options: AudioDataCopyToOptions | dict[str, object]) -> _CopyPlan: if self._data is None: - raise InvalidStateError('The data is closed') + msg = 'The data is closed' + raise InvalidStateError(msg) options = _copy_options(options) plane_index = _unsigned(options.plane_index, 'plane_index') frame_offset = _unsigned(options.frame_offset, 'frame_offset') destination = self._format if options.format is None else _sample_format(options.format) if _is_planar(destination): if plane_index >= self._channels: - raise InvalidRangeError(f'plane_index must be less than the {self._channels} channels') + msg = f'plane_index must be less than the {self._channels} channels' + raise InvalidRangeError(msg) elif plane_index > 0: - raise InvalidRangeError('plane_index must be 0 for an interleaved format') + msg = 'plane_index must be 0 for an interleaved format' + raise InvalidRangeError(msg) if frame_offset >= self._frames: - raise InvalidRangeError(f'frame_offset must be less than the {self._frames} frames') + msg = f'frame_offset must be less than the {self._frames} frames' + raise InvalidRangeError(msg) remaining = self._frames - frame_offset frame_count = remaining if options.frame_count is not None: frame_count = _unsigned(options.frame_count, 'frame_count') if frame_count > remaining: - raise InvalidRangeError(f'frame_count must be at most the {remaining} frames from frame_offset') + msg = f'frame_count must be at most the {remaining} frames from frame_offset' + raise InvalidRangeError(msg) elements = frame_count if _is_planar(destination) else frame_count * self._channels return _CopyPlan(destination, plane_index, frame_offset, frame_count, elements * _sample_bytes(destination)) - def allocation_size(self, options: Any) -> int: + def allocation_size(self, options: AudioDataCopyToOptions | dict[str, object]) -> int: """Returns how many bytes :meth:`copy_to` needs. + Raises :obj:`webrtc.InvalidStateError` if the data is closed, :obj:`TypeError` if the options aren't valid + and :obj:`webrtc.InvalidRangeError` if the plane or the frames don't exist. + Args: options (:obj:`AudioDataCopyToOptions`): What is copied. - - Raises: - :obj:`webrtc.InvalidStateError`: If the data is closed. - :obj:`TypeError`: If the options aren't valid. - :obj:`webrtc.InvalidRangeError`: If the plane or the frames don't exist. """ return self._plan_copy(options).size - def copy_to(self, destination: Union[bytearray, memoryview], options: Any) -> None: + def copy_to(self, destination: bytearray | memoryview, options: AudioDataCopyToOptions | dict[str, object]) -> None: """Copies samples into a buffer, converting them to another format if asked. + Raises the errors of :meth:`allocation_size` too. + Args: destination (:obj:`bytearray` or writable :obj:`memoryview`): The buffer. options (:obj:`AudioDataCopyToOptions`): What is copied. Raises: - :obj:`webrtc.InvalidStateError`: If the data is closed. - :obj:`webrtc.InvalidRangeError`: If the plane or the frames don't exist, or the buffer is too small. - :obj:`webrtc.NotSupportedError`: If the samples can't be converted to the format. + webrtc.InvalidRangeError: If the plane or the frames don't exist, or the buffer is too small. + webrtc.NotSupportedError: If the samples can't be converted to the format. """ plan = self._plan_copy(options) if memoryview(destination).nbytes < plan.size: - raise InvalidRangeError(f'The destination must be at least {plan.size} bytes') + msg = f'The destination must be at least {plan.size} bytes' + raise InvalidRangeError(msg) try: wrtc.copyAudioSamples( self._data, @@ -295,14 +332,15 @@ def copy_to(self, destination: Union[bytearray, memoryview], options: Any) -> No except ValueError as e: raise NotSupportedError(str(e)) from None - def clone(self) -> 'AudioData': + def clone(self) -> AudioData: """Returns another data of the same samples, which is closed separately. Raises: - :obj:`webrtc.InvalidStateError`: If the data is closed. + webrtc.InvalidStateError: If the data is closed. """ if self._data is None: - raise InvalidStateError('The data is closed') + msg = 'The data is closed' + raise InvalidStateError(msg) audio = AudioData.__new__(AudioData) audio.__dict__.update(self.__dict__) return audio @@ -311,13 +349,7 @@ def close(self) -> None: """Releases the samples. Closing a closed data does nothing.""" self._data = None - def __enter__(self) -> 'AudioData': - return self - - def __exit__(self, *exc_info) -> None: - self.close() - - def __repr__(self): + def __repr__(self) -> str: if self._data is None: return '' return ( diff --git a/python-webrtc/python/webrtc/models/blob.py b/python-webrtc/python/webrtc/models/blob.py index 8c334df..7fe49e9 100644 --- a/python-webrtc/python/webrtc/models/blob.py +++ b/python-webrtc/python/webrtc/models/blob.py @@ -5,16 +5,28 @@ # that can be found in the LICENSE.md file in the root of the project. # -"""Blob (https://developer.mozilla.org/en-US/docs/Web/API/Blob): immutable bytes, like the binary messages of a data -channel with a ``binaryType`` of ``'blob'``.""" +"""Blob (https://developer.mozilla.org/en-US/docs/Web/API/Blob): immutable bytes with a MIME type. + +Like the binary messages of a data channel with a ``binaryType`` of ``'blob'``. +""" + +from __future__ import annotations import asyncio -from typing import Any, Iterable, Optional, Union +from typing import TYPE_CHECKING, TypeVar, Union + +if TYPE_CHECKING: + from collections.abc import Iterable BlobPart = Union[str, bytes, bytearray, memoryview, 'Blob'] +_T = TypeVar('_T') + +# the printable ASCII range of a MIME type +_MIN_TYPE_CHAR = 0x20 +_MAX_TYPE_CHAR = 0x7E -def _done(value: Any) -> asyncio.Future: +def _done(value: _T) -> asyncio.Future[_T]: future = asyncio.get_running_loop().create_future() future.set_result(value) return future @@ -25,10 +37,10 @@ class Blob: Args: parts (iterable, optional): Strings (encoded as UTF-8), bytes-like objects and blobs, concatenated. - type (:obj:`str`, optional): The MIME type, lowercased; empty if it has characters outside of U+0020–U+007E. + type (:obj:`str`, optional): The MIME type, lowercased; empty if it has characters outside of U+0020-U+007E. """ - def __init__(self, parts: Optional[Iterable[BlobPart]] = None, type: str = ''): + def __init__(self, parts: Iterable[BlobPart] | None = None, type: str = '') -> None: chunks = [] for part in parts or (): if isinstance(part, Blob): @@ -40,7 +52,7 @@ def __init__(self, parts: Optional[Iterable[BlobPart]] = None, type: str = ''): chunks.append(bytes(memoryview(part))) self._bytes = b''.join(chunks) type = str(type) - self._type = type.lower() if all(0x20 <= ord(c) <= 0x7E for c in type) else '' + self._type = type.lower() if all(_MIN_TYPE_CHAR <= ord(c) <= _MAX_TYPE_CHAR for c in type) else '' @property def size(self) -> int: @@ -52,7 +64,7 @@ def type(self) -> str: """:obj:`str`: The MIME type, empty if unknown.""" return self._type - def slice(self, start: int = 0, end: Optional[int] = None, content_type: str = '') -> 'Blob': + def slice(self, start: int = 0, end: int | None = None, content_type: str = '') -> Blob: """Returns a blob of a range of the bytes. Args: @@ -65,15 +77,15 @@ def slice(self, start: int = 0, end: Optional[int] = None, content_type: str = ' end = size if end is None else (max(size + end, 0) if end < 0 else min(end, size)) return Blob([self._bytes[start : max(start, end)]], content_type) - def array_buffer(self) -> asyncio.Future: + def array_buffer(self) -> asyncio.Future[bytes]: """Returns a future of the bytes, as :obj:`bytes`.""" return _done(self._bytes) - def bytes(self) -> asyncio.Future: + def bytes(self) -> asyncio.Future[bytes]: """Returns a future of the bytes, as :obj:`bytes`.""" return _done(self._bytes) - def text(self) -> asyncio.Future: + def text(self) -> asyncio.Future[str]: """Returns a future of the bytes decoded as UTF-8.""" return _done(self._bytes.decode('utf-8', 'replace')) @@ -83,15 +95,15 @@ def __bytes__(self) -> bytes: def __len__(self) -> int: return len(self._bytes) - def __eq__(self, other): + def __eq__(self, other: object) -> bool: if isinstance(other, Blob): return self._bytes == other._bytes and self._type == other._type return NotImplemented - def __hash__(self): + def __hash__(self) -> int: return hash((self._bytes, self._type)) - def __repr__(self): + def __repr__(self) -> str: return f'' #: Alias for :meth:`array_buffer` diff --git a/python-webrtc/python/webrtc/models/closable.py b/python-webrtc/python/webrtc/models/closable.py new file mode 100644 index 0000000..aaf0106 --- /dev/null +++ b/python-webrtc/python/webrtc/models/closable.py @@ -0,0 +1,30 @@ +# +# Copyright 2026 Ilya (Marshal) . All rights reserved. +# +# Use of this source code is governed by a BSD-style license +# that can be found in the LICENSE.md file in the root of the project. +# + +"""Objects holding memory until closed, like media.""" + +from __future__ import annotations + +from abc import ABC, abstractmethod +from typing import TYPE_CHECKING + +if TYPE_CHECKING: + from typing_extensions import Self + + +class Closable(ABC): + """An object holding memory until :meth:`close`, which is called on leaving a ``with`` block too.""" + + @abstractmethod + def close(self) -> None: + """Releases the memory. Closing a closed object does nothing.""" + + def __enter__(self) -> Self: + return self + + def __exit__(self, *exc_info: object) -> None: + self.close() diff --git a/python-webrtc/python/webrtc/models/events.py b/python-webrtc/python/webrtc/models/events.py index 65401b0..3ec3daf 100644 --- a/python-webrtc/python/webrtc/models/events.py +++ b/python-webrtc/python/webrtc/models/events.py @@ -7,9 +7,12 @@ """Event objects passed to the handlers registered with ``on()`` (see :obj:`webrtc.utils.events.EventTarget`).""" -from typing import TYPE_CHECKING, Any, List, Optional, Union +from __future__ import annotations -from webrtc.utils.names import alias +from dataclasses import dataclass +from typing import TYPE_CHECKING, ClassVar + +from webrtc.utils.names import Alias, alias if TYPE_CHECKING: import webrtc @@ -23,15 +26,16 @@ class Event: target (:obj:`object`, optional): The object that emitted the event. """ - def __init__(self, type: str, target: Any = None): + def __init__(self, type: str, target: webrtc.EventTarget | None = None) -> None: self.type = type self.target = target - def __repr__(self): + def __repr__(self) -> str: fields = ', '.join(f'{k}={v!r}' for k, v in vars(self).items() if k != 'target') return f'{type(self).__name__}({fields})' +@dataclass(eq=False, repr=False) class RTCPeerConnectionIceEvent(Event): """An ``icecandidate`` event of :obj:`webrtc.RTCPeerConnection`. @@ -42,18 +46,13 @@ class RTCPeerConnectionIceEvent(Event): target (:obj:`object`, optional): The object that emitted the event. """ - def __init__( - self, - type: str, - candidate: Optional['webrtc.RTCIceCandidate'] = None, - url: Optional[str] = None, - target: Any = None, - ): - super().__init__(type, target) - self.candidate = candidate - self.url = url + type: str + candidate: webrtc.RTCIceCandidate | None = None + url: str | None = None + target: webrtc.EventTarget | None = None +@dataclass(eq=False, repr=False) class RTCPeerConnectionIceErrorEvent(Event): """An ``icecandidateerror`` event of :obj:`webrtc.RTCPeerConnection`: a STUN or TURN server failed. @@ -67,29 +66,21 @@ class RTCPeerConnectionIceErrorEvent(Event): target (:obj:`object`, optional): The object that emitted the event. """ - def __init__( - self, - type: str, - address: Optional[str], - port: Optional[int], - url: str, - error_code: int, - error_text: str, - target: Any = None, - ): - super().__init__(type, target) - self.address = address - self.port = port - self.url = url - self.error_code = error_code - self.error_text = error_text + type: str + address: str | None + port: int | None + url: str + error_code: int + error_text: str + target: webrtc.EventTarget | None = None #: Alias for :attr:`error_code` - errorCode = alias('error_code') + errorCode: ClassVar[Alias[int]] = alias('error_code') #: Alias for :attr:`error_text` - errorText = alias('error_text') + errorText: ClassVar[Alias[str]] = alias('error_text') +@dataclass(eq=False, repr=False) class MessageEvent(Event): """A ``message`` event of :obj:`webrtc.RTCDataChannel`. @@ -99,11 +90,12 @@ class MessageEvent(Event): target (:obj:`object`, optional): The object that emitted the event. """ - def __init__(self, type: str, data: Union[str, bytes], target: Any = None): - super().__init__(type, target) - self.data = data + type: str + data: str | bytes + target: webrtc.EventTarget | None = None +@dataclass(eq=False, repr=False) class RTCDataChannelEvent(Event): """A ``datachannel`` event of :obj:`webrtc.RTCPeerConnection`: the remote peer created a channel. @@ -113,11 +105,12 @@ class RTCDataChannelEvent(Event): target (:obj:`object`, optional): The object that emitted the event. """ - def __init__(self, type: str, channel: 'webrtc.RTCDataChannel', target: Any = None): - super().__init__(type, target) - self.channel = channel + type: str + channel: webrtc.RTCDataChannel + target: webrtc.EventTarget | None = None +@dataclass(eq=False, repr=False) class MediaStreamTrackEvent(Event): """An ``addtrack`` or ``removetrack`` event of :obj:`webrtc.MediaStream`. @@ -127,11 +120,12 @@ class MediaStreamTrackEvent(Event): target (:obj:`object`, optional): The object that emitted the event. """ - def __init__(self, type: str, track: 'webrtc.MediaStreamTrack', target: Any = None): - super().__init__(type, target) - self.track = track + type: str + track: webrtc.MediaStreamTrack + target: webrtc.EventTarget | None = None +@dataclass(eq=False, repr=False) class RTCDTMFToneChangeEvent(Event): """A ``tonechange`` event of :obj:`webrtc.RTCDTMFSender`. @@ -141,11 +135,12 @@ class RTCDTMFToneChangeEvent(Event): target (:obj:`object`, optional): The object that emitted the event. """ - def __init__(self, type: str, tone: str = '', target: Any = None): - super().__init__(type, target) - self.tone = tone + type: str + tone: str = '' + target: webrtc.EventTarget | None = None +@dataclass(eq=False, repr=False) class RTCErrorEvent(Event): """An ``error`` event, carrying the :obj:`webrtc.RTCError` that occurred. @@ -155,11 +150,12 @@ class RTCErrorEvent(Event): target (:obj:`object`, optional): The object that emitted the event. """ - def __init__(self, type: str, error: 'webrtc.RTCError', target: Any = None): - super().__init__(type, target) - self.error = error + type: str + error: webrtc.RTCError + target: webrtc.EventTarget | None = None +@dataclass(eq=False, repr=False) class RTCTrackEvent(Event): """A ``track`` event of :obj:`webrtc.RTCPeerConnection`: a remote track was negotiated. @@ -172,17 +168,9 @@ class RTCTrackEvent(Event): target (:obj:`object`, optional): The object that emitted the event. """ - def __init__( - self, - type: str, - receiver: 'webrtc.RTCRtpReceiver', - track: 'webrtc.MediaStreamTrack', - streams: List['webrtc.MediaStream'], - transceiver: 'webrtc.RTCRtpTransceiver', - target: Any = None, - ): - super().__init__(type, target) - self.receiver = receiver - self.track = track - self.streams = streams - self.transceiver = transceiver + type: str + receiver: webrtc.RTCRtpReceiver + track: webrtc.MediaStreamTrack + streams: list[webrtc.MediaStream] + transceiver: webrtc.RTCRtpTransceiver + target: webrtc.EventTarget | None = None diff --git a/python-webrtc/python/webrtc/models/media_track_constraints.py b/python-webrtc/python/webrtc/models/media_track_constraints.py index 5e20b1f..adfb83e 100644 --- a/python-webrtc/python/webrtc/models/media_track_constraints.py +++ b/python-webrtc/python/webrtc/models/media_track_constraints.py @@ -5,16 +5,20 @@ # that can be found in the LICENSE.md file in the root of the project. # -"""Settings, capabilities and constraints of tracks -(https://developer.mozilla.org/en-US/docs/Web/API/Media_Capture_and_Streams_API/Constraints).""" +"""Settings, capabilities and constraints of tracks. + +See https://developer.mozilla.org/en-US/docs/Web/API/Media_Capture_and_Streams_API/Constraints. +""" + +from __future__ import annotations from dataclasses import dataclass, fields -from typing import Any, Dict, List, Optional, Union +from typing import Any, ClassVar, Union -from webrtc.utils.names import alias, snake_case +from webrtc.utils.names import Alias, alias, snake_case #: A value, or a constraint on it: a :obj:`dict` with any of ``exact``, ``ideal``, ``min`` and ``max`` -ConstrainValue = Union[float, int, str, bool, Dict[str, Any]] +ConstrainValue = Union[float, int, str, bool, dict[str, Any]] @dataclass @@ -26,8 +30,8 @@ class ULongRange: max (:obj:`int`, optional): The highest value. """ - min: Optional[int] = None - max: Optional[int] = None + min: int | None = None + max: int | None = None @dataclass @@ -39,14 +43,15 @@ class DoubleRange: max (:obj:`float`, optional): The highest value. """ - min: Optional[float] = None - max: Optional[float] = None + min: float | None = None + max: float | None = None @dataclass class MediaTrackSettings: - """What a track carries, as far as it's known (:meth:`webrtc.MediaStreamTrack.get_settings`). Members are - :obj:`None` when they don't apply to the track. + """What a track carries, as far as it's known (:meth:`webrtc.MediaStreamTrack.get_settings`). + + Members are :obj:`None` when they don't apply to the track. Args: width (:obj:`int`, optional): The width of the video. @@ -64,48 +69,50 @@ class MediaTrackSettings: noise_suppression (:obj:`bool`, optional): Whether noise is suppressed. """ - width: Optional[int] = None - height: Optional[int] = None - aspect_ratio: Optional[float] = None - frame_rate: Optional[float] = None - resize_mode: Optional[str] = None - device_id: Optional[str] = None - group_id: Optional[str] = None - sample_rate: Optional[int] = None - sample_size: Optional[int] = None - channel_count: Optional[int] = None - echo_cancellation: Optional[bool] = None - auto_gain_control: Optional[bool] = None - noise_suppression: Optional[bool] = None + width: int | None = None + height: int | None = None + aspect_ratio: float | None = None + frame_rate: float | None = None + resize_mode: str | None = None + device_id: str | None = None + group_id: str | None = None + sample_rate: int | None = None + sample_size: int | None = None + channel_count: int | None = None + echo_cancellation: bool | None = None + auto_gain_control: bool | None = None + noise_suppression: bool | None = None #: Alias for :attr:`aspect_ratio` - aspectRatio = alias('aspect_ratio') + aspectRatio: ClassVar[Alias[float | None]] = alias('aspect_ratio') #: Alias for :attr:`frame_rate` - frameRate = alias('frame_rate') + frameRate: ClassVar[Alias[float | None]] = alias('frame_rate') #: Alias for :attr:`resize_mode` - resizeMode = alias('resize_mode') + resizeMode: ClassVar[Alias[str | None]] = alias('resize_mode') #: Alias for :attr:`device_id` - deviceId = alias('device_id') + deviceId: ClassVar[Alias[str | None]] = alias('device_id') #: Alias for :attr:`group_id` - groupId = alias('group_id') + groupId: ClassVar[Alias[str | None]] = alias('group_id') #: Alias for :attr:`sample_rate` - sampleRate = alias('sample_rate') + sampleRate: ClassVar[Alias[int | None]] = alias('sample_rate') #: Alias for :attr:`sample_size` - sampleSize = alias('sample_size') + sampleSize: ClassVar[Alias[int | None]] = alias('sample_size') #: Alias for :attr:`channel_count` - channelCount = alias('channel_count') + channelCount: ClassVar[Alias[int | None]] = alias('channel_count') #: Alias for :attr:`echo_cancellation` - echoCancellation = alias('echo_cancellation') + echoCancellation: ClassVar[Alias[bool | None]] = alias('echo_cancellation') #: Alias for :attr:`auto_gain_control` - autoGainControl = alias('auto_gain_control') + autoGainControl: ClassVar[Alias[bool | None]] = alias('auto_gain_control') #: Alias for :attr:`noise_suppression` - noiseSuppression = alias('noise_suppression') + noiseSuppression: ClassVar[Alias[bool | None]] = alias('noise_suppression') @dataclass class MediaTrackCapabilities: - """What the source of a track can do (:meth:`webrtc.MediaStreamTrack.get_capabilities`): the synthetic camera and - microphone of :func:`webrtc.get_user_media` have capabilities, other tracks don't control their source. + """What the source of a track can do (:meth:`webrtc.MediaStreamTrack.get_capabilities`). + + The synthetic camera and microphone of :func:`webrtc.get_user_media` have capabilities, other tracks don't control + their source. Args: width (:obj:`ULongRange`, optional): The widths of the video. @@ -123,49 +130,50 @@ class MediaTrackCapabilities: noise_suppression (:obj:`list` of :obj:`bool`, optional): Whether noise can be suppressed. """ - width: Optional[ULongRange] = None - height: Optional[ULongRange] = None - aspect_ratio: Optional[DoubleRange] = None - frame_rate: Optional[DoubleRange] = None - resize_mode: Optional[List[str]] = None - device_id: Optional[str] = None - group_id: Optional[str] = None - sample_rate: Optional[ULongRange] = None - sample_size: Optional[ULongRange] = None - channel_count: Optional[ULongRange] = None - echo_cancellation: Optional[List[bool]] = None - auto_gain_control: Optional[List[bool]] = None - noise_suppression: Optional[List[bool]] = None + width: ULongRange | None = None + height: ULongRange | None = None + aspect_ratio: DoubleRange | None = None + frame_rate: DoubleRange | None = None + resize_mode: list[str] | None = None + device_id: str | None = None + group_id: str | None = None + sample_rate: ULongRange | None = None + sample_size: ULongRange | None = None + channel_count: ULongRange | None = None + echo_cancellation: list[bool] | None = None + auto_gain_control: list[bool] | None = None + noise_suppression: list[bool] | None = None #: Alias for :attr:`aspect_ratio` - aspectRatio = alias('aspect_ratio') + aspectRatio: ClassVar[Alias[DoubleRange | None]] = alias('aspect_ratio') #: Alias for :attr:`frame_rate` - frameRate = alias('frame_rate') + frameRate: ClassVar[Alias[DoubleRange | None]] = alias('frame_rate') #: Alias for :attr:`resize_mode` - resizeMode = alias('resize_mode') + resizeMode: ClassVar[Alias[list[str] | None]] = alias('resize_mode') #: Alias for :attr:`device_id` - deviceId = alias('device_id') + deviceId: ClassVar[Alias[str | None]] = alias('device_id') #: Alias for :attr:`group_id` - groupId = alias('group_id') + groupId: ClassVar[Alias[str | None]] = alias('group_id') #: Alias for :attr:`sample_rate` - sampleRate = alias('sample_rate') + sampleRate: ClassVar[Alias[ULongRange | None]] = alias('sample_rate') #: Alias for :attr:`sample_size` - sampleSize = alias('sample_size') + sampleSize: ClassVar[Alias[ULongRange | None]] = alias('sample_size') #: Alias for :attr:`channel_count` - channelCount = alias('channel_count') + channelCount: ClassVar[Alias[ULongRange | None]] = alias('channel_count') #: Alias for :attr:`echo_cancellation` - echoCancellation = alias('echo_cancellation') + echoCancellation: ClassVar[Alias[list[bool] | None]] = alias('echo_cancellation') #: Alias for :attr:`auto_gain_control` - autoGainControl = alias('auto_gain_control') + autoGainControl: ClassVar[Alias[list[bool] | None]] = alias('auto_gain_control') #: Alias for :attr:`noise_suppression` - noiseSuppression = alias('noise_suppression') + noiseSuppression: ClassVar[Alias[list[bool] | None]] = alias('noise_suppression') @dataclass class MediaTrackConstraints: - """What a track is asked to be (:meth:`webrtc.MediaStreamTrack.apply_constraints`). Each member is a value (an - ideal one) or a :obj:`dict` of ``exact``, ``ideal``, ``min`` and ``max``: the required ones make the constraints - fail if the source can't satisfy them. + """What a track is asked to be (:meth:`webrtc.MediaStreamTrack.apply_constraints`). + + Each member is a value (an ideal one) or a :obj:`dict` of ``exact``, ``ideal``, ``min`` and ``max``: the required + ones make the constraints fail if the source can't satisfy them. Args: width (optional): The width of the video. @@ -185,54 +193,55 @@ class MediaTrackConstraints: be satisfied. """ - width: Optional[ConstrainValue] = None - height: Optional[ConstrainValue] = None - aspect_ratio: Optional[ConstrainValue] = None - frame_rate: Optional[ConstrainValue] = None - resize_mode: Optional[ConstrainValue] = None - device_id: Optional[ConstrainValue] = None - group_id: Optional[ConstrainValue] = None - sample_rate: Optional[ConstrainValue] = None - sample_size: Optional[ConstrainValue] = None - channel_count: Optional[ConstrainValue] = None - echo_cancellation: Optional[ConstrainValue] = None - auto_gain_control: Optional[ConstrainValue] = None - noise_suppression: Optional[ConstrainValue] = None - advanced: Optional[List[Dict[str, Any]]] = None + width: ConstrainValue | None = None + height: ConstrainValue | None = None + aspect_ratio: ConstrainValue | None = None + frame_rate: ConstrainValue | None = None + resize_mode: ConstrainValue | None = None + device_id: ConstrainValue | None = None + group_id: ConstrainValue | None = None + sample_rate: ConstrainValue | None = None + sample_size: ConstrainValue | None = None + channel_count: ConstrainValue | None = None + echo_cancellation: ConstrainValue | None = None + auto_gain_control: ConstrainValue | None = None + noise_suppression: ConstrainValue | None = None + advanced: list[dict[str, Any]] | None = None @classmethod - def _parse(cls, value: Any) -> 'MediaTrackConstraints': - """Constraints from an instance or a dictionary, with snake_case or camelCase names""" + def _parse(cls, value: object) -> MediaTrackConstraints: + """Constraints from an instance or a dictionary, with snake_case or camelCase names.""" if value is None: return cls() if isinstance(value, cls): return value if not isinstance(value, dict): - raise TypeError(f'{value!r} is not a MediaTrackConstraints') + msg = f'{value!r} is not a MediaTrackConstraints' + raise TypeError(msg) names = {f.name for f in fields(cls)} members = {snake_case(k): v for k, v in value.items()} # unknown members are ignored, as the specification says return cls(**{k: v for k, v in members.items() if k in names}) #: Alias for :attr:`aspect_ratio` - aspectRatio = alias('aspect_ratio') + aspectRatio: ClassVar[Alias[ConstrainValue | None]] = alias('aspect_ratio') #: Alias for :attr:`frame_rate` - frameRate = alias('frame_rate') + frameRate: ClassVar[Alias[ConstrainValue | None]] = alias('frame_rate') #: Alias for :attr:`resize_mode` - resizeMode = alias('resize_mode') + resizeMode: ClassVar[Alias[ConstrainValue | None]] = alias('resize_mode') #: Alias for :attr:`device_id` - deviceId = alias('device_id') + deviceId: ClassVar[Alias[ConstrainValue | None]] = alias('device_id') #: Alias for :attr:`group_id` - groupId = alias('group_id') + groupId: ClassVar[Alias[ConstrainValue | None]] = alias('group_id') #: Alias for :attr:`sample_rate` - sampleRate = alias('sample_rate') + sampleRate: ClassVar[Alias[ConstrainValue | None]] = alias('sample_rate') #: Alias for :attr:`sample_size` - sampleSize = alias('sample_size') + sampleSize: ClassVar[Alias[ConstrainValue | None]] = alias('sample_size') #: Alias for :attr:`channel_count` - channelCount = alias('channel_count') + channelCount: ClassVar[Alias[ConstrainValue | None]] = alias('channel_count') #: Alias for :attr:`echo_cancellation` - echoCancellation = alias('echo_cancellation') + echoCancellation: ClassVar[Alias[ConstrainValue | None]] = alias('echo_cancellation') #: Alias for :attr:`auto_gain_control` - autoGainControl = alias('auto_gain_control') + autoGainControl: ClassVar[Alias[ConstrainValue | None]] = alias('auto_gain_control') #: Alias for :attr:`noise_suppression` - noiseSuppression = alias('noise_suppression') + noiseSuppression: ClassVar[Alias[ConstrainValue | None]] = alias('noise_suppression') diff --git a/python-webrtc/python/webrtc/models/rtc_certificate.py b/python-webrtc/python/webrtc/models/rtc_certificate.py index b9e55f4..d8ee752 100644 --- a/python-webrtc/python/webrtc/models/rtc_certificate.py +++ b/python-webrtc/python/webrtc/models/rtc_certificate.py @@ -5,16 +5,22 @@ # that can be found in the LICENSE.md file in the root of the project. # +"""Certificates of DTLS.""" + +from __future__ import annotations + import asyncio import time +from collections.abc import Mapping from dataclasses import dataclass -from typing import Any, List, Mapping, Optional, Union +from itertools import starmap +from typing import Union from webrtc import NotSupportedError, WebRTCObject, wrtc from webrtc.utils.names import snake_case #: A WebCrypto algorithm: its name (like ``'ECDSA'``), or a dictionary with its name and parameters. -Algorithm = Union[str, Mapping[str, Any]] +Algorithm = Union[str, Mapping[str, object]] @dataclass(frozen=True) @@ -30,47 +36,65 @@ class RTCDtlsFingerprint: value: str -def _member(algorithm: Mapping[str, Any], name: str, default=None): +#: The key type, modulus length and public exponent, as the native generate() takes them. +_KeyParams = tuple[str, int, int] + + +def _member(algorithm: Mapping[str, object], name: str, default: object = None) -> object: """A member of a WebCrypto algorithm dictionary, by its camelCase or snake_case name.""" return algorithm.get(name, algorithm.get(snake_case(name), default)) -def _key_params(algorithm: Algorithm): - """The key type, modulus length and public exponent for an algorithm, as the native generate() takes them.""" +def _ecdsa_params(algorithm: Mapping[str, object]) -> _KeyParams: + curve = _member(algorithm, 'namedCurve', 'P-256') + if curve != 'P-256': + msg = f'the {curve} curve is not supported, only P-256 is' + raise NotSupportedError(msg) + return 'ecdsa', 0, 0 + + +def _rsa_params(algorithm: Mapping[str, object]) -> _KeyParams: + hash_name = _member(algorithm, 'hash') + if isinstance(hash_name, Mapping): + hash_name = hash_name.get('name') + modulus_length = _member(algorithm, 'modulusLength') + exponent = _member(algorithm, 'publicExponent') + if hash_name is None or modulus_length is None or exponent is None: + msg = 'RSASSA-PKCS1-v1_5 needs a hash, a modulus length and a public exponent' + raise NotSupportedError(msg) + if str(hash_name).upper() != 'SHA-256': + msg = f'the {hash_name} hash is not supported, only SHA-256 is' + raise NotSupportedError(msg) + if isinstance(exponent, (bytes, bytearray, memoryview)): + exponent = int.from_bytes(bytes(exponent), 'big') + return 'rsa', int(modulus_length), int(exponent) + + +_KEY_PARAMS = {'ECDSA': _ecdsa_params, 'RSASSA-PKCS1-V1_5': _rsa_params} + + +def _key_params(algorithm: Algorithm) -> _KeyParams: + """The key parameters for an algorithm.""" if isinstance(algorithm, str): algorithm = {'name': algorithm} - name = str(_member(algorithm, 'name', '')).upper() - if name == 'ECDSA': - curve = _member(algorithm, 'namedCurve', 'P-256') - if curve != 'P-256': - raise NotSupportedError(f'the {curve} curve is not supported, only P-256 is') - return 'ecdsa', 0, 0 - if name == 'RSASSA-PKCS1-V1_5': - hash_name = _member(algorithm, 'hash') - if isinstance(hash_name, Mapping): - hash_name = hash_name.get('name') - modulus_length = _member(algorithm, 'modulusLength') - exponent = _member(algorithm, 'publicExponent') - if hash_name is None or modulus_length is None or exponent is None: - raise NotSupportedError('RSASSA-PKCS1-v1_5 needs a hash, a modulus length and a public exponent') - if str(hash_name).upper() != 'SHA-256': - raise NotSupportedError(f'the {hash_name} hash is not supported, only SHA-256 is') - if isinstance(exponent, (bytes, bytearray, memoryview)): - exponent = int.from_bytes(bytes(exponent), 'big') - return 'rsa', int(modulus_length), int(exponent) - raise NotSupportedError( - f'the {algorithm.get("name")!r} algorithm is not supported, ECDSA and RSASSA-PKCS1-v1_5 are' - ) + key_params = _KEY_PARAMS.get(str(_member(algorithm, 'name', '')).upper()) + if key_params is None: + msg = f'the {algorithm.get("name")!r} algorithm is not supported, ECDSA and RSASSA-PKCS1-v1_5 are' + raise NotSupportedError(msg) + return key_params(algorithm) class RTCCertificate(WebRTCObject): - """A certificate a connection uses to authenticate with DTLS, from :meth:`generate`, and set with - :attr:`webrtc.RTCConfiguration.certificates`. Without one, a connection generates its own.""" + """A certificate a connection uses to authenticate with DTLS. + + Generated with :meth:`generate` and set with :attr:`webrtc.RTCConfiguration.certificates`. Without one, + a connection generates its own. + """ _class = wrtc.RTCCertificate @classmethod - async def generate(cls, algorithm: Algorithm = 'ECDSA', expires: Optional[float] = None) -> 'RTCCertificate': + async def generate(cls, algorithm: Algorithm = 'ECDSA', expires: float | None = None) -> RTCCertificate: """Generates a key and a self-signed certificate, on a worker thread. Args: @@ -84,12 +108,13 @@ async def generate(cls, algorithm: Algorithm = 'ECDSA', expires: Optional[float] :obj:`webrtc.RTCCertificate`: The certificate. Raises: - :obj:`webrtc.NotSupportedError`: If the algorithm isn't supported. - :obj:`ValueError`: If ``expires`` is negative. + webrtc.NotSupportedError: If the algorithm isn't supported. + ValueError: If ``expires`` is negative. """ key_type, modulus_length, exponent = _key_params(algorithm) if expires is not None and expires < 0: - raise ValueError(f'expires must not be negative, not {expires}') + msg = f'expires must not be negative, not {expires}' + raise ValueError(msg) native = await asyncio.get_running_loop().run_in_executor( None, cls._class.generate, @@ -99,7 +124,8 @@ async def generate(cls, algorithm: Algorithm = 'ECDSA', expires: Optional[float] int(expires) if expires is not None else None, ) if native is None: - raise NotSupportedError('the key could not be generated with these parameters') + msg = 'the key could not be generated with these parameters' + raise NotSupportedError(msg) return cls._wrap(native) @property @@ -112,13 +138,13 @@ def expired(self) -> bool: """:obj:`bool`: Whether the certificate has expired.""" return self.expires <= time.time() * 1000 - def get_fingerprints(self) -> List[RTCDtlsFingerprint]: + def get_fingerprints(self) -> list[RTCDtlsFingerprint]: """Returns the fingerprints of the certificate. Returns: :obj:`list` of :obj:`webrtc.RTCDtlsFingerprint`: The fingerprints. """ - return [RTCDtlsFingerprint(algorithm, value) for algorithm, value in self._native_obj.fingerprints()] + return list(starmap(RTCDtlsFingerprint, self._native_obj.fingerprints())) #: Alias for :attr:`get_fingerprints` getFingerprints = get_fingerprints diff --git a/python-webrtc/python/webrtc/models/rtc_configuration.py b/python-webrtc/python/webrtc/models/rtc_configuration.py index b5343f7..d45ed53 100644 --- a/python-webrtc/python/webrtc/models/rtc_configuration.py +++ b/python-webrtc/python/webrtc/models/rtc_configuration.py @@ -5,10 +5,14 @@ # that can be found in the LICENSE.md file in the root of the project. # +"""The configuration of a connection.""" + +from __future__ import annotations + import ipaddress import re from dataclasses import dataclass, field -from typing import Any, Dict, Iterable, List, Optional, Tuple, Union +from typing import TYPE_CHECKING, Any, ClassVar from webrtc import ( InvalidAccessError, @@ -21,10 +25,15 @@ RTCRtpHeaderEncryptionPolicy, wrtc, ) -from webrtc.utils.names import alias +from webrtc.utils.names import Alias, alias + +if TYPE_CHECKING: + from collections.abc import Iterable # the longest TURN username, as browsers limit it _MAX_USERNAME_LENGTH = 509 +_MAX_PORT = 65535 +_MAX_CANDIDATE_POOL_SIZE = 255 # RFC 7064 and RFC 7065: scheme ":" host [ ":" port ] [ "?transport=" transport ], with udp and tcp transports only _URL = re.compile( @@ -38,25 +47,30 @@ def _check_url(url: str) -> str: """Returns the scheme of a STUN or TURN server URL. Raises: - :obj:`webrtc.InvalidSyntaxError`: If the URL doesn't match RFC 7064 or RFC 7065. + webrtc.InvalidSyntaxError: If the URL doesn't match RFC 7064 or RFC 7065. """ match = _URL.fullmatch(url) - if not match or match['scheme'] not in ('stun', 'stuns', 'turn', 'turns'): - raise InvalidSyntaxError(f'{url!r} is not a valid STUN or TURN URL') + if not match or match['scheme'] not in {'stun', 'stuns', 'turn', 'turns'}: + msg = f'{url!r} is not a valid STUN or TURN URL' + raise InvalidSyntaxError(msg) host, port, transport = match['host'], match['port'], match['transport'] if host.startswith('['): try: ipaddress.IPv6Address(host[1:-1]) except ValueError: - raise InvalidSyntaxError(f'{url!r} has an invalid IPv6 address') from None + msg = f'{url!r} has an invalid IPv6 address' + raise InvalidSyntaxError(msg) from None elif not _REG_NAME.fullmatch(host): - raise InvalidSyntaxError(f'{url!r} has an invalid host') - - if port is not None and (not port or int(port) > 65535): - raise InvalidSyntaxError(f'{url!r} has an invalid port') - if transport is not None and (match['scheme'].startswith('stun') or transport not in ('udp', 'tcp')): - raise InvalidSyntaxError(f'{url!r} has an invalid transport') + msg = f'{url!r} has an invalid host' + raise InvalidSyntaxError(msg) + + if port is not None and (not port or int(port) > _MAX_PORT): + msg = f'{url!r} has an invalid port' + raise InvalidSyntaxError(msg) + if transport is not None and (match['scheme'].startswith('stun') or transport not in {'udp', 'tcp'}): + msg = f'{url!r} has an invalid transport' + raise InvalidSyntaxError(msg) return match['scheme'] @@ -73,9 +87,9 @@ class RTCOAuthCredential: access_token: str #: Alias for :attr:`mac_key` - macKey = alias('mac_key') + macKey: ClassVar[Alias[str]] = alias('mac_key') #: Alias for :attr:`access_token` - accessToken = alias('access_token') + accessToken: ClassVar[Alias[str]] = alias('access_token') @dataclass @@ -92,34 +106,40 @@ class RTCIceServer: doesn't support. """ - urls: Union[str, List[str]] - username: Optional[str] = None - credential: Optional[Union[str, RTCOAuthCredential]] = None + urls: str | list[str] + username: str | None = None + credential: str | RTCOAuthCredential | None = None credential_type: str = 'password' @classmethod - def _to_native_list(cls, servers: Iterable[Union['RTCIceServer', Dict[str, Any]]]) -> List['wrtc.IceServerInit']: + def _to_native_list(cls, servers: Iterable[RTCIceServer | dict[str, Any]]) -> list[wrtc.IceServerInit]: """The native servers of a list of servers, or of their keyword arguments.""" return [(cls(**server) if isinstance(server, dict) else server)._to_native() for server in servers] - def _to_native(self) -> 'wrtc.IceServerInit': + def _to_native(self) -> wrtc.IceServerInit: urls = [self.urls] if isinstance(self.urls, str) else list(self.urls) if not urls: - raise InvalidSyntaxError('urls of an ICE server must not be empty') + msg = 'urls of an ICE server must not be empty' + raise InvalidSyntaxError(msg) # every URL is parsed before the credentials are checked schemes = [_check_url(url) for url in urls] - if self.credential_type not in ('password', 'oauth'): - raise ValueError(f"credential_type must be 'password' or 'oauth', not {self.credential_type!r}") - if any(scheme in ('turn', 'turns') for scheme in schemes): + if self.credential_type not in {'password', 'oauth'}: + msg = f"credential_type must be 'password' or 'oauth', not {self.credential_type!r}" + raise ValueError(msg) + if any(scheme in {'turn', 'turns'} for scheme in schemes): if self.credential_type == 'oauth': if not isinstance(self.credential, RTCOAuthCredential): - raise InvalidAccessError('an OAuth TURN server needs an RTCOAuthCredential') - raise NotSupportedError('libwebrtc does not support OAuth credentials of TURN servers') + msg = 'an OAuth TURN server needs an RTCOAuthCredential' + raise InvalidAccessError(msg) + msg = 'libwebrtc does not support OAuth credentials of TURN servers' + raise NotSupportedError(msg) if self.username is None or not self.credential: - raise InvalidAccessError('a TURN server needs a username and a credential') + msg = 'a TURN server needs a username and a credential' + raise InvalidAccessError(msg) if len(self.username) > _MAX_USERNAME_LENGTH: - raise InvalidAccessError(f'the username of a TURN server is longer than {_MAX_USERNAME_LENGTH}') + msg = f'the username of a TURN server is longer than {_MAX_USERNAME_LENGTH}' + raise InvalidAccessError(msg) native = wrtc.IceServerInit() native.urls = urls @@ -128,7 +148,7 @@ def _to_native(self) -> 'wrtc.IceServerInit': return native #: Alias for :attr:`credential_type` - credentialType = alias('credential_type') + credentialType: ClassVar[Alias[str]] = alias('credential_type') @dataclass @@ -157,24 +177,25 @@ class RTCConfiguration: a remote description without it fails (``require``). Can't be changed. """ - ice_servers: List[Union[RTCIceServer, Dict[str, Any]]] = field(default_factory=list) + ice_servers: list[RTCIceServer | dict[str, Any]] = field(default_factory=list) ice_transport_policy: RTCIceTransportPolicy = RTCIceTransportPolicy.all bundle_policy: RTCBundlePolicy = RTCBundlePolicy.balanced rtcp_mux_policy: RTCRtcpMuxPolicy = RTCRtcpMuxPolicy.require ice_candidate_pool_size: int = 0 - port_range: Optional[Tuple[int, int]] = None - certificates: Optional[List[RTCCertificate]] = None + port_range: tuple[int, int] | None = None + certificates: list[RTCCertificate] | None = None always_negotiate_data_channels: bool = False rtp_header_encryption_policy: RTCRtpHeaderEncryptionPolicy = RTCRtpHeaderEncryptionPolicy.negotiate - def _to_native(self) -> 'wrtc.ConfigurationInit': + def _to_native(self) -> wrtc.ConfigurationInit: """Validates the configuration and creates the native one. + Returns: + :obj:`wrtc.ConfigurationInit`: The native configuration. + Raises: - :obj:`TypeError`: If a member has a wrong type. - :obj:`ValueError`: If ``ice_candidate_pool_size`` or ``port_range`` is out of range. - :obj:`webrtc.InvalidSyntaxError`: If an ICE server URL is invalid. - :obj:`webrtc.InvalidAccessError`: If a TURN server has no credentials. + ValueError: If ``ice_candidate_pool_size`` or ``port_range`` is out of range. + webrtc.InvalidAccessError: If a certificate has expired. """ native = wrtc.ConfigurationInit() native.iceServers = RTCIceServer._to_native_list(self.ice_servers) @@ -182,27 +203,31 @@ def _to_native(self) -> 'wrtc.ConfigurationInit': native.bundlePolicy = self.bundle_policy native.rtcpMuxPolicy = self.rtcp_mux_policy - if not isinstance(self.ice_candidate_pool_size, int) or not 0 <= self.ice_candidate_pool_size <= 255: - raise ValueError(f'ice_candidate_pool_size must be from 0 to 255, not {self.ice_candidate_pool_size}') - native.iceCandidatePoolSize = self.ice_candidate_pool_size + pool_size = self.ice_candidate_pool_size + if not isinstance(pool_size, int) or not 0 <= pool_size <= _MAX_CANDIDATE_POOL_SIZE: + msg = f'ice_candidate_pool_size must be from 0 to {_MAX_CANDIDATE_POOL_SIZE}, not {pool_size}' + raise ValueError(msg) + native.iceCandidatePoolSize = pool_size native.alwaysNegotiateDataChannels = bool(self.always_negotiate_data_channels) native.rtpHeaderEncryptionPolicy = self.rtp_header_encryption_policy if self.certificates is not None: for certificate in self.certificates: if certificate.expired: - raise InvalidAccessError('the certificate has expired') + msg = 'the certificate has expired' + raise InvalidAccessError(msg) native.certificates = [certificate._native_obj for certificate in self.certificates] if self.port_range is not None: low, high = self.port_range - if not 0 <= low <= high <= 65535: - raise ValueError(f'port_range must be two ports from low to high, not {self.port_range}') + if not 0 <= low <= high <= _MAX_PORT: + msg = f'port_range must be two ports from low to high, not {self.port_range}' + raise ValueError(msg) native.portRange = (low, high) return native @classmethod - def _from_native(cls, native: 'wrtc.ConfigurationInit') -> 'RTCConfiguration': + def _from_native(cls, native: wrtc.ConfigurationInit) -> RTCConfiguration: return cls( ice_servers=[ RTCIceServer( @@ -223,18 +248,18 @@ def _from_native(cls, native: 'wrtc.ConfigurationInit') -> 'RTCConfiguration': ) #: Alias for :attr:`ice_servers` - iceServers = alias('ice_servers') + iceServers: ClassVar[Alias[list[RTCIceServer | dict[str, Any]]]] = alias('ice_servers') #: Alias for :attr:`ice_transport_policy` - iceTransportPolicy = alias('ice_transport_policy') + iceTransportPolicy: ClassVar[Alias[RTCIceTransportPolicy]] = alias('ice_transport_policy') #: Alias for :attr:`bundle_policy` - bundlePolicy = alias('bundle_policy') + bundlePolicy: ClassVar[Alias[RTCBundlePolicy]] = alias('bundle_policy') #: Alias for :attr:`rtcp_mux_policy` - rtcpMuxPolicy = alias('rtcp_mux_policy') + rtcpMuxPolicy: ClassVar[Alias[RTCRtcpMuxPolicy]] = alias('rtcp_mux_policy') #: Alias for :attr:`ice_candidate_pool_size` - iceCandidatePoolSize = alias('ice_candidate_pool_size') + iceCandidatePoolSize: ClassVar[Alias[int]] = alias('ice_candidate_pool_size') #: Alias for :attr:`port_range` - portRange = alias('port_range') + portRange: ClassVar[Alias[tuple[int, int] | None]] = alias('port_range') #: Alias for :attr:`always_negotiate_data_channels` - alwaysNegotiateDataChannels = alias('always_negotiate_data_channels') + alwaysNegotiateDataChannels: ClassVar[Alias[bool]] = alias('always_negotiate_data_channels') #: Alias for :attr:`rtp_header_encryption_policy` - rtpHeaderEncryptionPolicy = alias('rtp_header_encryption_policy') + rtpHeaderEncryptionPolicy: ClassVar[Alias[RTCRtpHeaderEncryptionPolicy]] = alias('rtp_header_encryption_policy') diff --git a/python-webrtc/python/webrtc/models/rtc_ice_candidate.py b/python-webrtc/python/webrtc/models/rtc_ice_candidate.py index 5c1c51d..3332efd 100644 --- a/python-webrtc/python/webrtc/models/rtc_ice_candidate.py +++ b/python-webrtc/python/webrtc/models/rtc_ice_candidate.py @@ -5,9 +5,14 @@ # that can be found in the LICENSE.md file in the root of the project. # +"""ICE candidates and parameters.""" + +from __future__ import annotations + import re from dataclasses import dataclass -from typing import Any, Dict, Optional, Tuple, Union +from enum import Enum +from typing import Any, ClassVar, TypeVar, Union from webrtc import ( RTCIceCandidateType, @@ -16,78 +21,124 @@ RTCIceServerTransportProtocol, RTCIceTcpCandidateType, ) -from webrtc.utils.names import alias +from webrtc.utils.names import Alias, alias _FOUNDATION = re.compile(r'[A-Za-z0-9+/]{1,32}') _DIGITS = re.compile(r'[0-9]+') _TOKEN = re.compile(r"[!#$%&'*+\-.^_`{|}~A-Za-z0-9]+") +_MIN_TOKENS = 8 # foundation, component, transport, priority, address, port, "typ" and the type +_RELATED_TOKENS = 4 # "raddr", the address, "rport" and the port +_TCP_TYPE_TOKENS = 2 # "tcptype" and the type +_PORTS = range(65536) +_COMPONENTS = {1: 'rtp', 2: 'rtcp'} +_PROTOCOLS = frozenset({'udp', 'tcp'}) +_TYPES = frozenset({'host', 'srflx', 'prflx', 'relay'}) +_TCP_TYPES = frozenset({'active', 'passive', 'so'}) + +_T = TypeVar('_T') +_EnumT = TypeVar('_EnumT', bound=Enum) +#: The fields parsed from a candidate-attribute, by the names of the properties of RTCIceCandidate +_CandidateFields = dict[str, Union[str, int, None]] + + +class _InvalidCandidateError(ValueError): + """A candidate-attribute that doesn't parse.""" + -def _number(token: str, max_digits: int, low: int, high: int) -> Optional[int]: +def _number(token: str, max_digits: int, valid: range) -> int | None: if len(token) > max_digits or not _DIGITS.fullmatch(token): return None value = int(token) - return value if low <= value <= high else None + return value if value in valid else None -def _parse_candidate(value: str, strict: bool = True) -> Optional[Dict[str, Any]]: - """Parses a candidate-attribute (RFC 8839, with the tcptype of RFC 6544), or returns :obj:`None`. - Not strict, a candidate other than a host one may have no related address, as libwebrtc describes - peer-reflexive candidates.""" - if not value.startswith('candidate:'): - return None - tokens = value[len('candidate:') :].split(' ') - if len(tokens) < 8 or tokens[6] != 'typ' or any(not t for t in tokens): - return None +def _required(value: _T | None) -> _T: + if value is None: + raise _InvalidCandidateError + return value - foundation, component, transport, priority, address, port, _, cand_type, *rest = tokens - component_id = _number(component, 3, 1, 256) - fields = { - 'foundation': foundation if _FOUNDATION.fullmatch(foundation) else None, - 'component': {1: 'rtp', 2: 'rtcp'}.get(component_id), - 'priority': _number(priority, 10, 1, 2**31 - 1), + +def _one_of(value: str, allowed: frozenset[str]) -> str: + if value not in allowed: + raise _InvalidCandidateError + return value + + +def _parse_related(fields: _CandidateFields, rest: list[str], *, strict: bool) -> list[str]: + """Parses the related address and port into the fields, returns the tokens after them.""" + if rest[:1] == ['raddr']: + if len(rest) < _RELATED_TOKENS or rest[2] != 'rport': + raise _InvalidCandidateError + fields['related_address'] = rest[1] + fields['related_port'] = _required(_number(rest[3], 5, _PORTS)) + return rest[_RELATED_TOKENS:] + if fields['type'] != 'host' and strict: + raise _InvalidCandidateError + return rest + + +def _parse_tcp_type(fields: _CandidateFields, rest: list[str]) -> list[str]: + """Parses the tcptype into the fields, returns the tokens after it.""" + if rest[:1] == ['tcptype']: + if len(rest) < _TCP_TYPE_TOKENS: + raise _InvalidCandidateError + fields['tcp_type'] = _one_of(rest[1].lower(), _TCP_TYPES) + return rest[_TCP_TYPE_TOKENS:] + if fields['protocol'] == 'tcp' and fields['type'] != 'relay': + raise _InvalidCandidateError + return rest + + +def _base_fields(tokens: list[str]) -> _CandidateFields: + """The fields of the tokens up to the type.""" + foundation, component, transport, priority, address, port, _, cand_type = tokens + component_id = _required(_number(component, 3, range(1, 257))) + return { + 'foundation': _required(foundation if _FOUNDATION.fullmatch(foundation) else None), + 'component': _COMPONENTS.get(component_id), + 'priority': _required(_number(priority, 10, range(1, 2**31))), 'address': address, - 'protocol': transport.lower(), - 'port': _number(port, 5, 0, 65535), - 'type': cand_type.lower(), + 'protocol': _one_of(transport.lower(), _PROTOCOLS), + 'port': _required(_number(port, 5, _PORTS)), + 'type': _one_of(cand_type.lower(), _TYPES), 'tcp_type': None, 'related_address': None, 'related_port': None, } - if ( - None in (fields['foundation'], fields['priority'], fields['port']) - or component_id is None - or fields['protocol'] not in ('udp', 'tcp') - or fields['type'] not in ('host', 'srflx', 'prflx', 'relay') - ): - return None - if rest[:1] == ['raddr']: - if len(rest) < 4 or rest[2] != 'rport': - return None - fields['related_address'] = rest[1] - fields['related_port'] = _number(rest[3], 5, 0, 65535) - if fields['related_port'] is None: - return None - rest = rest[4:] - elif fields['type'] != 'host' and strict: - return None - if rest[:1] == ['tcptype']: - if len(rest) < 2 or rest[1].lower() not in ('active', 'passive', 'so'): - return None - fields['tcp_type'] = rest[1].lower() - rest = rest[2:] - elif fields['protocol'] == 'tcp' and fields['type'] != 'relay': - return None +def _parse_fields(value: str, *, strict: bool) -> _CandidateFields: + if not value.startswith('candidate:'): + raise _InvalidCandidateError + tokens = value[len('candidate:') :].split(' ') + if len(tokens) < _MIN_TOKENS or tokens[6] != 'typ' or not all(tokens): + raise _InvalidCandidateError + fields = _base_fields(tokens[:_MIN_TOKENS]) + rest = _parse_tcp_type(fields, _parse_related(fields, tokens[_MIN_TOKENS:], strict=strict)) # extensions are pairs of a token and a value without spaces (like an ufrag, with "/" and "+") if len(rest) % 2 or not all(_TOKEN.fullmatch(t) for t in rest[::2]): - return None + raise _InvalidCandidateError return fields -def _member_or_none(cls, value): +def _parse_candidate(value: str, *, strict: bool = True) -> _CandidateFields | None: + """Parses a candidate-attribute (RFC 8839, with the tcptype of RFC 6544). + + Not strict, a candidate other than a host one may have no related address, as libwebrtc describes + peer-reflexive candidates. + + Returns: + :obj:`dict`: The fields, or :obj:`None` if the candidate doesn't parse. + """ + try: + return _parse_fields(value, strict=strict) + except _InvalidCandidateError: + return None + + +def _member_or_none(cls: type[_EnumT], value: object) -> _EnumT | None: # candidates with values the enum doesn't have are valid try: return cls(value) if value is not None else None @@ -108,7 +159,7 @@ class RTCIceParameters: password: str #: Alias for :attr:`username_fragment` - usernameFragment = alias('username_fragment') + usernameFragment: ClassVar[Alias[str]] = alias('username_fragment') @dataclass(frozen=True) @@ -120,8 +171,8 @@ class RTCIceCandidatePair: remote (:obj:`webrtc.RTCIceCandidate`): The remote candidate. """ - local: 'RTCIceCandidate' - remote: 'RTCIceCandidate' + local: RTCIceCandidate + remote: RTCIceCandidate @dataclass(frozen=True, repr=False) @@ -143,19 +194,20 @@ class RTCIceCandidate: url (:obj:`str`, optional): For a local candidate, the STUN or TURN server that gathered it. Raises: - :obj:`TypeError`: If both ``sdp_mid`` and ``sdp_m_line_index`` are :obj:`None`. + TypeError: If both ``sdp_mid`` and ``sdp_m_line_index`` are :obj:`None`. """ candidate: str = '' - sdp_mid: Optional[str] = None - sdp_m_line_index: Optional[int] = None - username_fragment: Optional[str] = None - relay_protocol: Optional[RTCIceServerTransportProtocol] = None - url: Optional[str] = None + sdp_mid: str | None = None + sdp_m_line_index: int | None = None + username_fragment: str | None = None + relay_protocol: RTCIceServerTransportProtocol | None = None + url: str | None = None - def __post_init__(self): + def __post_init__(self) -> None: if self.sdp_mid is None and self.sdp_m_line_index is None: - raise TypeError('sdp_mid and sdp_m_line_index are both None') + msg = 'sdp_mid and sdp_m_line_index are both None' + raise TypeError(msg) object.__setattr__(self, 'candidate', str(self.candidate)) object.__setattr__(self, 'relay_protocol', _member_or_none(RTCIceServerTransportProtocol, self.relay_protocol)) # the fields parsed from the candidate, not a field of the dataclass @@ -163,8 +215,8 @@ def __post_init__(self): @staticmethod def _members_of( - candidate: Union['RTCIceCandidate', Dict[str, Any]], - ) -> Tuple[str, Optional[str], Optional[int], Optional[str]]: + candidate: RTCIceCandidate | dict[str, Any], + ) -> tuple[str, str | None, int | None, str | None]: """The candidate, sdp_mid, sdp_m_line_index and username_fragment of a candidate or of its JSON form.""" if isinstance(candidate, RTCIceCandidate): return candidate.candidate, candidate.sdp_mid, candidate.sdp_m_line_index, candidate.username_fragment @@ -175,24 +227,31 @@ def _members_of( candidate.get('sdpMLineIndex'), candidate.get('usernameFragment'), ) - raise TypeError(f'candidate must be an RTCIceCandidate or a dict, not {type(candidate).__name__}') + msg = f'candidate must be an RTCIceCandidate or a dict, not {type(candidate).__name__}' + raise TypeError(msg) @classmethod - def _peer_reflexive(cls, kwargs: Dict[str, Any]) -> 'RTCIceCandidate': - """A remote peer-reflexive candidate, only known from connectivity checks: its candidate string and - address aren't exposed, as the remote peer didn't signal them. + def _peer_reflexive(cls, kwargs: dict[str, Any]) -> RTCIceCandidate: + """A remote peer-reflexive candidate, only known from connectivity checks. + + Its candidate string and address aren't exposed, as the remote peer didn't signal them. Args: - kwargs (:obj:`dict`): The arguments of the constructor, as the native candidate gives them.""" + kwargs (:obj:`dict`): The arguments of the constructor, as the native candidate gives them. + + Returns: + :obj:`webrtc.RTCIceCandidate`: The candidate. + """ fields = _parse_candidate(kwargs.get('candidate', ''), strict=False) or {} candidate = cls(**{**kwargs, 'candidate': ''}) # libwebrtc knows no related address of it, which is port 0 parsed = {**fields, 'address': None, 'related_address': None, 'related_port': 0} - object.__setattr__(candidate, '_parsed', parsed) + # past the frozen __setattr__, as __post_init__ parsed the empty candidate string + vars(candidate)['_parsed'] = parsed return candidate @classmethod - def from_json(cls, init: Dict[str, Any]) -> 'RTCIceCandidate': + def from_json(cls, init: dict[str, Any]) -> RTCIceCandidate: """Creates a candidate from its JSON form, as :meth:`to_json` returns it. Args: @@ -201,63 +260,60 @@ def from_json(cls, init: Dict[str, Any]) -> 'RTCIceCandidate': Returns: :obj:`webrtc.RTCIceCandidate`: The candidate. - - Raises: - :obj:`TypeError`: If both ``sdpMid`` and ``sdpMLineIndex`` are missing or :obj:`None`. """ return cls(*cls._members_of(init)) @property - def foundation(self) -> Optional[str]: + def foundation(self) -> str | None: """:obj:`str`, optional: An identifier of candidates of the same type, base and server.""" return self._parsed.get('foundation') @property - def component(self) -> Optional[RTCIceComponent]: + def component(self) -> RTCIceComponent | None: """:obj:`webrtc.RTCIceComponent`, optional: Whether the candidate is for RTP or RTCP.""" return _member_or_none(RTCIceComponent, self._parsed.get('component')) @property - def priority(self) -> Optional[int]: + def priority(self) -> int | None: """:obj:`int`, optional: The priority of the candidate.""" return self._parsed.get('priority') @property - def address(self) -> Optional[str]: + def address(self) -> str | None: """:obj:`str`, optional: The IP address or the host name of the candidate.""" return self._parsed.get('address') @property - def protocol(self) -> Optional[RTCIceProtocol]: + def protocol(self) -> RTCIceProtocol | None: """:obj:`webrtc.RTCIceProtocol`, optional: The transport protocol of the candidate.""" return _member_or_none(RTCIceProtocol, self._parsed.get('protocol')) @property - def port(self) -> Optional[int]: + def port(self) -> int | None: """:obj:`int`, optional: The port of the candidate.""" return self._parsed.get('port') @property - def type(self) -> Optional[RTCIceCandidateType]: + def type(self) -> RTCIceCandidateType | None: """:obj:`webrtc.RTCIceCandidateType`, optional: The type of the candidate.""" return _member_or_none(RTCIceCandidateType, self._parsed.get('type')) @property - def tcp_type(self) -> Optional[RTCIceTcpCandidateType]: + def tcp_type(self) -> RTCIceTcpCandidateType | None: """:obj:`webrtc.RTCIceTcpCandidateType`, optional: The type of a TCP candidate.""" return _member_or_none(RTCIceTcpCandidateType, self._parsed.get('tcp_type')) @property - def related_address(self) -> Optional[str]: + def related_address(self) -> str | None: """:obj:`str`, optional: For a candidate that isn't a host one, the address it's derived from.""" return self._parsed.get('related_address') @property - def related_port(self) -> Optional[int]: + def related_port(self) -> int | None: """:obj:`int`, optional: For a candidate that isn't a host one, the port it's derived from.""" return self._parsed.get('related_port') - def to_json(self) -> Dict[str, Any]: + def to_json(self) -> dict[str, Any]: """The candidate as a JSON-serializable dictionary, to send to the remote peer. Returns: @@ -270,26 +326,26 @@ def to_json(self) -> Dict[str, Any]: 'usernameFragment': self.username_fragment, } - def __repr__(self): + def __repr__(self) -> str: return ( f'RTCIceCandidate({self.candidate!r}, sdp_mid={self.sdp_mid!r}, sdp_m_line_index={self.sdp_m_line_index!r})' ) #: Alias for :attr:`sdp_mid` - sdpMid = alias('sdp_mid') + sdpMid: ClassVar[Alias[str | None]] = alias('sdp_mid') #: Alias for :attr:`sdp_m_line_index` - sdpMLineIndex = alias('sdp_m_line_index') + sdpMLineIndex: ClassVar[Alias[int | None]] = alias('sdp_m_line_index') #: Alias for :attr:`username_fragment` - usernameFragment = alias('username_fragment') + usernameFragment: ClassVar[Alias[str | None]] = alias('username_fragment') #: Alias for :attr:`relay_protocol` - relayProtocol = alias('relay_protocol') + relayProtocol: ClassVar[Alias[RTCIceServerTransportProtocol | None]] = alias('relay_protocol') #: Alias for :attr:`tcp_type` - tcpType = tcp_type + tcpType: ClassVar = tcp_type #: Alias for :attr:`related_address` - relatedAddress = related_address + relatedAddress: ClassVar = related_address #: Alias for :attr:`related_port` - relatedPort = related_port + relatedPort: ClassVar = related_port #: Alias for :attr:`to_json` - toJSON = to_json + toJSON: ClassVar = to_json #: Alias for :attr:`from_json` - fromJSON = from_json + fromJSON: ClassVar = from_json diff --git a/python-webrtc/python/webrtc/models/rtc_session_description.py b/python-webrtc/python/webrtc/models/rtc_session_description.py index f8b8b15..33d2a01 100644 --- a/python-webrtc/python/webrtc/models/rtc_session_description.py +++ b/python-webrtc/python/webrtc/models/rtc_session_description.py @@ -5,7 +5,11 @@ # that can be found in the LICENSE.md file in the root of the project. # -from typing import TYPE_CHECKING, Any, Dict, Union +"""The description of one end of a connection.""" + +from __future__ import annotations + +from typing import TYPE_CHECKING, Any from webrtc import RTCSessionDescriptionInit, WebRTCObject, wrtc @@ -14,8 +18,9 @@ class RTCSessionDescription(WebRTCObject): - """The :obj:`webrtc.RTCSessionDescription` interface describes one end of a connection or potential - connection and how it's configured. Each :obj:`webrtc.RTCSessionDescription` consists of + """One end of a connection or potential connection and how it's configured. + + Each :obj:`webrtc.RTCSessionDescription` consists of a description type indicating which part of the offer/answer negotiation process it describes and of the SDP descriptor of the session. @@ -38,26 +43,28 @@ class RTCSessionDescription(WebRTCObject): is set. Raises: - :obj:`TypeError`: If the type is missing, or the SDP is :obj:`None`. + TypeError: If the type is missing, or the SDP is :obj:`None`. """ _class = wrtc.RTCSessionDescription def __init__( self, - type: Union['webrtc.RTCSdpType', 'webrtc.RTCSessionDescriptionInit', Dict[str, Any]], + type: webrtc.RTCSdpType | webrtc.RTCSessionDescriptionInit | dict[str, Any], sdp: str = '', - ): + ) -> None: if isinstance(type, dict): if type.get('type') is None: - raise TypeError('RTCSessionDescriptionInit requires a type') + msg = 'RTCSessionDescriptionInit requires a type' + raise TypeError(msg) type, sdp = type['type'], type.get('sdp', '') if sdp is None: - raise TypeError('The SDP of a description may not be None') + msg = 'The SDP of a description may not be None' + raise TypeError(msg) init = type if isinstance(type, RTCSessionDescriptionInit) else RTCSessionDescriptionInit(type, sdp) super().__init__(self._class(init._native_obj)) - def to_json(self) -> Dict[str, str]: + def to_json(self) -> dict[str, str]: """The description as a JSON-serializable dictionary, to send to the remote peer. Returns: @@ -66,12 +73,12 @@ def to_json(self) -> Dict[str, str]: return {'type': self.type.value, 'sdp': self.sdp} @property - def type(self) -> 'webrtc.RTCSdpType': + def type(self) -> webrtc.RTCSdpType: """:obj:`webrtc.RTCSdpType`: A member of the :obj:`webrtc.RTCSdpType` enum.""" return self._native_obj.type @property - def sdp(self): + def sdp(self) -> str: """:obj:`str`: A string containing a SDP message describing the session.""" return self._native_obj.sdp diff --git a/python-webrtc/python/webrtc/models/rtc_session_description_init.py b/python-webrtc/python/webrtc/models/rtc_session_description_init.py index 4c0dce1..700f05d 100644 --- a/python-webrtc/python/webrtc/models/rtc_session_description_init.py +++ b/python-webrtc/python/webrtc/models/rtc_session_description_init.py @@ -5,7 +5,11 @@ # that can be found in the LICENSE.md file in the root of the project. # -from typing import TYPE_CHECKING, Dict +"""The type and the SDP of a description.""" + +from __future__ import annotations + +from typing import TYPE_CHECKING from webrtc import WebRTCObject, wrtc @@ -14,8 +18,9 @@ class RTCSessionDescriptionInit(WebRTCObject): - """The type and the SDP of a description, as :meth:`webrtc.RTCPeerConnection.create_offer` and - :meth:`webrtc.RTCPeerConnection.create_answer` return them. + """The type and the SDP of a description. + + As :meth:`webrtc.RTCPeerConnection.create_offer` and :meth:`webrtc.RTCPeerConnection.create_answer` return them. Args: type (:obj:`webrtc.RTCSdpType`): The type of the description, or its value (like ``'offer'``). @@ -24,29 +29,31 @@ class RTCSessionDescriptionInit(WebRTCObject): _class = wrtc.RTCSessionDescriptionInit - def __init__(self, type: 'webrtc.RTCSdpType', sdp: str = ''): + def __init__(self, type: webrtc.RTCSdpType, sdp: str = '') -> None: super().__init__(self._class(type, sdp)) @property - def type(self) -> 'webrtc.RTCSdpType': + def type(self) -> webrtc.RTCSdpType: """:obj:`webrtc.RTCSdpType`: A member of the :obj:`webrtc.RTCSdpType` enum.""" return self._native_obj.type @type.setter - def type(self, value: 'webrtc.RTCSdpType'): + def type(self, value: webrtc.RTCSdpType) -> None: self._native_obj.type = value @property def sdp(self) -> str: """:obj:`str`: A string containing a SDP message describing the session. - This value is an empty string by default and may not be :obj:`None`.""" + + This value is an empty string by default and may not be :obj:`None`. + """ return self._native_obj.sdp @sdp.setter - def sdp(self, value: str): + def sdp(self, value: str) -> None: self._native_obj.sdp = value - def to_json(self) -> Dict[str, str]: + def to_json(self) -> dict[str, str]: """The description as a JSON-serializable dictionary, to send to the remote peer. Returns: @@ -54,7 +61,7 @@ def to_json(self) -> Dict[str, str]: """ return {'type': self.type.value, 'sdp': self.sdp} - def __repr__(self): + def __repr__(self) -> str: return f'RTCSessionDescriptionInit(type={self.type.value!r}, sdp={len(self.sdp)} characters)' #: Alias for :attr:`to_json` diff --git a/python-webrtc/python/webrtc/models/rtc_stats.py b/python-webrtc/python/webrtc/models/rtc_stats.py index b77079a..967a0d6 100644 --- a/python-webrtc/python/webrtc/models/rtc_stats.py +++ b/python-webrtc/python/webrtc/models/rtc_stats.py @@ -5,18 +5,29 @@ # that can be found in the LICENSE.md file in the root of the project. # +"""The stats of the WebRTC Statistics specification.""" + +from __future__ import annotations + import json -from typing import TYPE_CHECKING, Any, Dict, Iterable, Iterator, List, Mapping +from collections.abc import Iterable, Iterator, Mapping +from typing import TYPE_CHECKING, Union from webrtc.utils.names import camel_case if TYPE_CHECKING: import webrtc +_CANDIDATE_TYPES = frozenset({'local-candidate', 'remote-candidate'}) +# a value of the JSON libwebrtc serializes the stats to +StatsValue = Union[str, int, float, bool, None, list['StatsValue'], dict[str, 'StatsValue']] + -class RTCStats(Dict[str, Any]): - """Stats of one object, like an outbound RTP stream: a :obj:`dict` of the members of the stats dictionary - of the WebRTC Statistics specification, by their names there (like ``'bytesSent'``). +class RTCStats(dict[str, StatsValue]): + """Stats of one object, like an outbound RTP stream. + + A :obj:`dict` of the members of the stats dictionary of the WebRTC Statistics specification, by their names + there (like ``'bytesSent'``). Members can also be read as attributes with snake_case names:: @@ -25,12 +36,13 @@ class RTCStats(Dict[str, Any]): ``id``, ``type`` and ``timestamp`` (milliseconds since the epoch) are always present. """ - def __getattr__(self, name: str) -> Any: + def __getattr__(self, name: str) -> StatsValue: key = camel_case(name) try: return self[key] except KeyError: - raise AttributeError(f'{type(self).__name__} of type {self.get("type")!r} has no {name!r}') from None + msg = f'{type(self).__name__} of type {self.get("type")!r} has no {name!r}' + raise AttributeError(msg) from None @property def id(self) -> str: @@ -49,25 +61,31 @@ def timestamp(self) -> float: class RTCStatsReport(Mapping[str, RTCStats]): - """The stats of a connection, or of a sender or a receiver: a read-only mapping of their ids to - :obj:`webrtc.RTCStats`.""" + """The stats of a connection, or of a sender or a receiver. + + A read-only mapping of their ids to :obj:`webrtc.RTCStats`. + """ - def __init__(self, stats: Mapping[str, RTCStats]): + def __init__(self, stats: Mapping[str, RTCStats]) -> None: self._stats = dict(stats) @classmethod - def _from_native(cls, report: str, receivers: Iterable['webrtc.RTCRtpReceiver'] = ()) -> 'RTCStatsReport': + def _from_native(cls, report: str, receivers: Iterable[webrtc.RTCRtpReceiver] = ()) -> RTCStatsReport: """The report from the JSON libwebrtc serializes it to, with the receivers whose tracks it refers to.""" # remote tracks have their own ids, rather than the libwebrtc ones in the stats track_ids = {receiver.track._native_obj._nativeId: receiver.track.id for receiver in receivers} stats = [RTCStats(entry) for entry in json.loads(report or '[]')] for entry in stats: # libwebrtc serializes microseconds - entry['timestamp'] = entry['timestamp'] / 1000 + entry['timestamp'] /= 1000 if entry.get('type') == 'inbound-rtp' and entry.get('trackIdentifier') in track_ids: entry['trackIdentifier'] = track_ids[entry['trackIdentifier']] # libwebrtc leaves the addresses of candidates it doesn't expose (like peer-reflexive ones) empty - if entry.get('type') in ('local-candidate', 'remote-candidate') and entry.get('address') == '': + if ( + entry.get('type') in {'local-candidate', 'remote-candidate'} + and 'address' in entry + and not entry['address'] + ): entry['address'] = None return cls({entry['id']: entry for entry in stats}) @@ -80,7 +98,7 @@ def __iter__(self) -> Iterator[str]: def __len__(self) -> int: return len(self._stats) - def of_type(self, stats_type: str) -> List[RTCStats]: + def of_type(self, stats_type: str) -> list[RTCStats]: """Returns the stats of a type. Args: @@ -91,7 +109,7 @@ def of_type(self, stats_type: str) -> List[RTCStats]: """ return [stats for stats in self._stats.values() if stats['type'] == stats_type] - def __repr__(self): + def __repr__(self) -> str: return f'RTCStatsReport({len(self)} stats)' #: Alias for :attr:`of_type` diff --git a/python-webrtc/python/webrtc/models/rtp_parameters.py b/python-webrtc/python/webrtc/models/rtp_parameters.py index 943631a..07c4522 100644 --- a/python-webrtc/python/webrtc/models/rtp_parameters.py +++ b/python-webrtc/python/webrtc/models/rtp_parameters.py @@ -7,13 +7,17 @@ """RTP parameters and capabilities of senders, receivers and transceivers.""" +from __future__ import annotations + import dataclasses import math from dataclasses import dataclass, field -from typing import Any, Dict, List, Optional +from typing import Any, ClassVar, TypeVar from webrtc import MediaType, RTCDegradationPreference, RTCPriorityType, TransceiverDirection, wrtc -from webrtc.utils.names import alias +from webrtc.utils.names import Alias, alias + +_NativeCodecT = TypeVar('_NativeCodecT', bound='wrtc.RtpCodec') # the bitrate priorities of libwebrtc for RTCRtpEncodingParameters.priority, as Chromium maps them _BITRATE_PRIORITY = { @@ -24,10 +28,15 @@ } -def _parse_fmtp(line: Optional[str]) -> Dict[str, str]: +def _is_unsigned_long(value: object) -> bool: + """Whether a value is an [EnforceRange] unsigned long of WebIDL.""" + return not isinstance(value, bool) and isinstance(value, int) and 0 <= value < 2**32 + + +def _parse_fmtp(line: str | None) -> dict[str, str]: parameters = {} - for item in (line or '').split(';'): - item = item.strip() + for raw_item in (line or '').split(';'): + item = raw_item.strip() if not item: continue if '=' in item: @@ -39,13 +48,13 @@ def _parse_fmtp(line: Optional[str]) -> Dict[str, str]: return parameters -def _format_fmtp(parameters: Dict[str, str]) -> Optional[str]: +def _format_fmtp(parameters: dict[str, str]) -> str | None: if not parameters: return None return ';'.join(f'{key}={value}' if key else value for key, value in parameters.items()) -def _codec_members(native: 'wrtc.RtpCodec') -> Dict[str, Any]: +def _codec_members(native: wrtc.RtpCodec) -> dict[str, Any]: """The members RTCRtpCodec and RTCRtpCodecParameters share.""" return { 'mime_type': native.mimeType, @@ -69,18 +78,19 @@ class RTCRtpCodec: mime_type: str clock_rate: int - channels: Optional[int] = None - sdp_fmtp_line: Optional[str] = None + channels: int | None = None + sdp_fmtp_line: str | None = None @classmethod - def _from_native(cls, native: 'wrtc.RtpCodec') -> 'RTCRtpCodec': + def _from_native(cls, native: wrtc.RtpCodec) -> RTCRtpCodec: return cls(**_codec_members(native)) - def _to_native(self, native_class: type): + def _to_native(self, native_class: type[_NativeCodecT]) -> _NativeCodecT: """Creates a native codec of a class: ``wrtc.RtpCodec`` or ``wrtc.RtpCodecCapability``.""" kind, _, name = self.mime_type.partition('/') if not name: - raise ValueError(f'{self.mime_type!r} is not a valid MIME type of a codec') + msg = f'{self.mime_type!r} is not a valid MIME type of a codec' + raise ValueError(msg) native = native_class() native.kind = MediaType.video if kind.lower() == 'video' else MediaType.audio native.name = name @@ -89,7 +99,7 @@ def _to_native(self, native_class: type): native.parameters = _parse_fmtp(self.sdp_fmtp_line) return native - def _matches(self, other: 'RTCRtpCodec') -> bool: + def _matches(self, other: RTCRtpCodec) -> bool: # codecs match with a case-insensitive MIME type and the same parameters, in any order return ( self.mime_type.lower() == other.mime_type.lower() @@ -99,11 +109,11 @@ def _matches(self, other: 'RTCRtpCodec') -> bool: ) #: Alias for :attr:`mime_type` - mimeType = alias('mime_type') + mimeType: ClassVar[Alias[str]] = alias('mime_type') #: Alias for :attr:`clock_rate` - clockRate = alias('clock_rate') + clockRate: ClassVar[Alias[int]] = alias('clock_rate') #: Alias for :attr:`sdp_fmtp_line` - sdpFmtpLine = alias('sdp_fmtp_line') + sdpFmtpLine: ClassVar[Alias[str | None]] = alias('sdp_fmtp_line') @dataclass @@ -121,21 +131,21 @@ class RTCRtpCodecParameters: payload_type: int mime_type: str clock_rate: int - channels: Optional[int] = None - sdp_fmtp_line: Optional[str] = None + channels: int | None = None + sdp_fmtp_line: str | None = None @classmethod - def _from_native(cls, native: 'wrtc.RtpCodecParameters') -> 'RTCRtpCodecParameters': + def _from_native(cls, native: wrtc.RtpCodecParameters) -> RTCRtpCodecParameters: return cls(payload_type=native.payloadType, **_codec_members(native)) #: Alias for :attr:`payload_type` - payloadType = alias('payload_type') + payloadType: ClassVar[Alias[int]] = alias('payload_type') #: Alias for :attr:`mime_type` - mimeType = alias('mime_type') + mimeType: ClassVar[Alias[str]] = alias('mime_type') #: Alias for :attr:`clock_rate` - clockRate = alias('clock_rate') + clockRate: ClassVar[Alias[int]] = alias('clock_rate') #: Alias for :attr:`sdp_fmtp_line` - sdpFmtpLine = alias('sdp_fmtp_line') + sdpFmtpLine: ClassVar[Alias[str | None]] = alias('sdp_fmtp_line') @dataclass @@ -153,7 +163,7 @@ class RTCRtpHeaderExtensionParameters: encrypted: bool = False @classmethod - def _from_native(cls, native: 'wrtc.RtpExtension') -> 'RTCRtpHeaderExtensionParameters': + def _from_native(cls, native: wrtc.RtpExtension) -> RTCRtpHeaderExtensionParameters: return cls(uri=native.uri, id=native.id, encrypted=native.encrypt) @@ -166,11 +176,11 @@ class RTCRtcpParameters: reduced_size (:obj:`bool`, optional): Whether reduced-size RTCP is negotiated. """ - cname: Optional[str] = None - reduced_size: Optional[bool] = None + cname: str | None = None + reduced_size: bool | None = None #: Alias for :attr:`reduced_size` - reducedSize = alias('reduced_size') + reducedSize: ClassVar[Alias[bool | None]] = alias('reduced_size') @dataclass @@ -191,18 +201,18 @@ class RTCRtpEncodingParameters: """ active: bool = True - max_bitrate: Optional[int] = None - max_framerate: Optional[float] = None - rid: Optional[str] = None - scale_resolution_down_by: Optional[float] = None + max_bitrate: int | None = None + max_framerate: float | None = None + rid: str | None = None + scale_resolution_down_by: float | None = None priority: RTCPriorityType = RTCPriorityType.low network_priority: RTCPriorityType = RTCPriorityType.low - scalability_mode: Optional[str] = None + scalability_mode: str | None = None adaptive_ptime: bool = False - codec: Optional[RTCRtpCodec] = None + codec: RTCRtpCodec | None = None @classmethod - def _from_native(cls, native: 'wrtc.RtpEncodingParameters') -> 'RTCRtpEncodingParameters': + def _from_native(cls, native: wrtc.RtpEncodingParameters) -> RTCRtpEncodingParameters: priority = min(_BITRATE_PRIORITY, key=lambda p: abs(_BITRATE_PRIORITY[p] - native.bitratePriority)) return cls( active=native.active, @@ -217,25 +227,28 @@ def _from_native(cls, native: 'wrtc.RtpEncodingParameters') -> 'RTCRtpEncodingPa codec=RTCRtpCodec._from_native(native.codec) if native.codec is not None else None, ) - def _for_kind(self, kind: MediaType) -> 'RTCRtpEncodingParameters': - """The encoding for a sender of a kind: members of video encodings are ignored for audio, whatever - their value.""" + def _for_kind(self, kind: MediaType) -> RTCRtpEncodingParameters: + """The encoding for a sender of a kind. + + Returns: + :obj:`RTCRtpEncodingParameters`: The encoding, without the members of video ones for audio. + """ if kind == MediaType.video: return self return dataclasses.replace(self, max_framerate=None, scale_resolution_down_by=None) - def _apply(self, native: 'wrtc.RtpEncodingParameters') -> 'wrtc.RtpEncodingParameters': + def _apply(self, native: wrtc.RtpEncodingParameters) -> wrtc.RtpEncodingParameters: """:meth:`_to_native` into an existing native encoding: sets the members that can be changed.""" # the WebIDL types: an [EnforceRange] unsigned long and restricted doubles bitrate = self.max_bitrate - if bitrate is not None and ( - isinstance(bitrate, bool) or not isinstance(bitrate, int) or not 0 <= bitrate < 2**32 - ): - raise TypeError(f'max_bitrate must be an unsigned 32-bit integer, not {bitrate!r}') + if bitrate is not None and not _is_unsigned_long(bitrate): + msg = f'max_bitrate must be an unsigned 32-bit integer, not {bitrate!r}' + raise TypeError(msg) for name in ('max_framerate', 'scale_resolution_down_by'): value = getattr(self, name) if value is not None and (not isinstance(value, (int, float)) or not math.isfinite(value)): - raise TypeError(f'{name} must be a finite number, not {value!r}') + msg = f'{name} must be a finite number, not {value!r}' + raise TypeError(msg) native.active = bool(self.active) native.maxBitrate = min(bitrate, 2**31 - 1) if bitrate is not None else None native.maxFramerate = self.max_framerate @@ -247,23 +260,23 @@ def _apply(self, native: 'wrtc.RtpEncodingParameters') -> 'wrtc.RtpEncodingParam native.codec = self.codec._to_native(wrtc.RtpCodec) if self.codec is not None else None return native - def _to_native(self) -> 'wrtc.RtpEncodingParameters': + def _to_native(self) -> wrtc.RtpEncodingParameters: native = self._apply(wrtc.RtpEncodingParameters()) native.rid = self.rid or '' return native #: Alias for :attr:`max_bitrate` - maxBitrate = alias('max_bitrate') + maxBitrate: ClassVar[Alias[int | None]] = alias('max_bitrate') #: Alias for :attr:`max_framerate` - maxFramerate = alias('max_framerate') + maxFramerate: ClassVar[Alias[float | None]] = alias('max_framerate') #: Alias for :attr:`scale_resolution_down_by` - scaleResolutionDownBy = alias('scale_resolution_down_by') + scaleResolutionDownBy: ClassVar[Alias[float | None]] = alias('scale_resolution_down_by') #: Alias for :attr:`network_priority` - networkPriority = alias('network_priority') + networkPriority: ClassVar[Alias[RTCPriorityType]] = alias('network_priority') #: Alias for :attr:`scalability_mode` - scalabilityMode = alias('scalability_mode') + scalabilityMode: ClassVar[Alias[str | None]] = alias('scalability_mode') #: Alias for :attr:`adaptive_ptime` - adaptivePtime = alias('adaptive_ptime') + adaptivePtime: ClassVar[Alias[bool]] = alias('adaptive_ptime') @dataclass @@ -276,12 +289,12 @@ class RTCRtpReceiveParameters: rtcp (:obj:`webrtc.RTCRtcpParameters`): The RTCP parameters. """ - codecs: List[RTCRtpCodecParameters] = field(default_factory=list) - header_extensions: List[RTCRtpHeaderExtensionParameters] = field(default_factory=list) + codecs: list[RTCRtpCodecParameters] = field(default_factory=list) + header_extensions: list[RTCRtpHeaderExtensionParameters] = field(default_factory=list) rtcp: RTCRtcpParameters = field(default_factory=RTCRtcpParameters) @classmethod - def _from_native(cls, native: 'wrtc.RtpParameters') -> 'RTCRtpReceiveParameters': + def _from_native(cls, native: wrtc.RtpParameters) -> RTCRtpReceiveParameters: return cls( codecs=[RTCRtpCodecParameters._from_native(c) for c in native.codecs], header_extensions=[RTCRtpHeaderExtensionParameters._from_native(e) for e in native.headerExtensions], @@ -289,7 +302,7 @@ def _from_native(cls, native: 'wrtc.RtpParameters') -> 'RTCRtpReceiveParameters' ) #: Alias for :attr:`header_extensions` - headerExtensions = alias('header_extensions') + headerExtensions: ClassVar[Alias[list[RTCRtpHeaderExtensionParameters]]] = alias('header_extensions') @dataclass @@ -309,14 +322,14 @@ class RTCRtpSendParameters: """ transaction_id: str - encodings: List[RTCRtpEncodingParameters] = field(default_factory=list) - codecs: List[RTCRtpCodecParameters] = field(default_factory=list) - header_extensions: List[RTCRtpHeaderExtensionParameters] = field(default_factory=list) + encodings: list[RTCRtpEncodingParameters] = field(default_factory=list) + codecs: list[RTCRtpCodecParameters] = field(default_factory=list) + header_extensions: list[RTCRtpHeaderExtensionParameters] = field(default_factory=list) rtcp: RTCRtcpParameters = field(default_factory=RTCRtcpParameters) - degradation_preference: Optional[RTCDegradationPreference] = None + degradation_preference: RTCDegradationPreference | None = None @classmethod - def _from_native(cls, native: 'wrtc.RtpParameters') -> 'RTCRtpSendParameters': + def _from_native(cls, native: wrtc.RtpParameters) -> RTCRtpSendParameters: return cls( transaction_id=native.transactionId, encodings=[RTCRtpEncodingParameters._from_native(e) for e in native.encodings], @@ -327,11 +340,11 @@ def _from_native(cls, native: 'wrtc.RtpParameters') -> 'RTCRtpSendParameters': ) #: Alias for :attr:`transaction_id` - transactionId = alias('transaction_id') + transactionId: ClassVar[Alias[str]] = alias('transaction_id') #: Alias for :attr:`header_extensions` - headerExtensions = alias('header_extensions') + headerExtensions: ClassVar[Alias[list[RTCRtpHeaderExtensionParameters]]] = alias('header_extensions') #: Alias for :attr:`degradation_preference` - degradationPreference = alias('degradation_preference') + degradationPreference: ClassVar[Alias[RTCDegradationPreference | None]] = alias('degradation_preference') @dataclass @@ -348,7 +361,7 @@ class RTCRtpHeaderExtensionCapability: direction: TransceiverDirection = TransceiverDirection.sendrecv @classmethod - def _from_native(cls, native: 'wrtc.RtpHeaderExtensionCapability') -> 'RTCRtpHeaderExtensionCapability': + def _from_native(cls, native: wrtc.RtpHeaderExtensionCapability) -> RTCRtpHeaderExtensionCapability: return cls(uri=native.uri, direction=native.direction) @@ -361,11 +374,11 @@ class RTCRtpCapabilities: header_extensions (:obj:`list` of :obj:`webrtc.RTCRtpHeaderExtensionCapability`): The header extensions. """ - codecs: List[RTCRtpCodec] = field(default_factory=list) - header_extensions: List[RTCRtpHeaderExtensionCapability] = field(default_factory=list) + codecs: list[RTCRtpCodec] = field(default_factory=list) + header_extensions: list[RTCRtpHeaderExtensionCapability] = field(default_factory=list) @classmethod - def _from_native(cls, native: 'wrtc.RtpCapabilities') -> 'RTCRtpCapabilities': + def _from_native(cls, native: wrtc.RtpCapabilities) -> RTCRtpCapabilities: return cls( codecs=[RTCRtpCodec._from_native(c) for c in native.codecs], # only the URIs: the directions are the ones of a transceiver (see get_header_extensions_to_negotiate) @@ -373,11 +386,16 @@ def _from_native(cls, native: 'wrtc.RtpCapabilities') -> 'RTCRtpCapabilities': ) @classmethod - def _supported(cls, native_class: type, kind: MediaType) -> Optional['RTCRtpCapabilities']: - """The capabilities of ``wrtc.RTCRtpSender`` or ``wrtc.RTCRtpReceiver`` for a kind, :obj:`None` for - another kind.""" + def _supported( + cls, native_class: type[wrtc.RTCRtpSender | wrtc.RTCRtpReceiver], kind: MediaType + ) -> RTCRtpCapabilities | None: + """The capabilities of ``wrtc.RTCRtpSender`` or ``wrtc.RTCRtpReceiver`` for a kind. + + Returns: + :obj:`RTCRtpCapabilities`: The capabilities, :obj:`None` for another kind. + """ native = native_class.getCapabilities(str(kind)) return cls._from_native(native) if native is not None else None #: Alias for :attr:`header_extensions` - headerExtensions = alias('header_extensions') + headerExtensions: ClassVar[Alias[list[RTCRtpHeaderExtensionCapability]]] = alias('header_extensions') diff --git a/python-webrtc/python/webrtc/models/rtp_source.py b/python-webrtc/python/webrtc/models/rtp_source.py index 953d464..b4d17b2 100644 --- a/python-webrtc/python/webrtc/models/rtp_source.py +++ b/python-webrtc/python/webrtc/models/rtp_source.py @@ -5,10 +5,17 @@ # that can be found in the LICENSE.md file in the root of the project. # +"""The contributing and synchronization sources of the media a receiver received.""" + +from __future__ import annotations + from dataclasses import dataclass -from typing import Optional, Tuple +from typing import ClassVar + +from webrtc.utils.names import Alias, alias -from webrtc.utils.names import alias +# RFC 6464 and RFC 6465 levels are -dBov, 127 being silence +_SILENT_LEVEL = 127 @dataclass(frozen=True) @@ -27,24 +34,25 @@ class RTCRtpContributingSource: timestamp: float source: int rtp_timestamp: int - audio_level: Optional[float] = None + audio_level: float | None = None @classmethod - def _from_native(cls, native: Tuple[bool, int, float, int, Optional[int]]) -> 'RTCRtpContributingSource': + def _from_native(cls, native: tuple[bool, int, float, int, int | None]) -> RTCRtpContributingSource: """A source from the native one: whether it's an SSRC, the source, timestamp, RTP timestamp and level.""" _, source, timestamp, rtp_timestamp, level = native - # RFC 6464 and RFC 6465 levels are -dBov, 127 being silence if level is not None: - level = 0.0 if level >= 127 else 10 ** (-level / 20) + level = 0.0 if level >= _SILENT_LEVEL else 10 ** (-level / 20) return cls(timestamp, source, rtp_timestamp, level) #: Alias for :attr:`rtp_timestamp` - rtpTimestamp = alias('rtp_timestamp') + rtpTimestamp: ClassVar[Alias[int]] = alias('rtp_timestamp') #: Alias for :attr:`audio_level` - audioLevel = alias('audio_level') + audioLevel: ClassVar[Alias[float | None]] = alias('audio_level') @dataclass(frozen=True) class RTCRtpSynchronizationSource(RTCRtpContributingSource): - """A synchronization source (SSRC) of the media an :obj:`webrtc.RTCRtpReceiver` received in the last - 10 seconds. See :obj:`webrtc.RTCRtpContributingSource` for its members.""" + """A synchronization source (SSRC) of the media an :obj:`webrtc.RTCRtpReceiver` received in the last 10 s. + + See :obj:`webrtc.RTCRtpContributingSource` for its members. + """ diff --git a/python-webrtc/python/webrtc/models/rtp_transceiver_init.py b/python-webrtc/python/webrtc/models/rtp_transceiver_init.py index 177a42a..f41ee49 100644 --- a/python-webrtc/python/webrtc/models/rtp_transceiver_init.py +++ b/python-webrtc/python/webrtc/models/rtp_transceiver_init.py @@ -5,7 +5,11 @@ # that can be found in the LICENSE.md file in the root of the project. # -from typing import TYPE_CHECKING, List, Optional +"""The options of a new transceiver.""" + +from __future__ import annotations + +from typing import TYPE_CHECKING from webrtc import WebRTCObject, wrtc @@ -29,10 +33,10 @@ class RtpTransceiverInit(WebRTCObject): def __init__( self, - direction: Optional['webrtc.TransceiverDirection'] = None, - send_encodings: Optional[List['webrtc.RTCRtpEncodingParameters']] = None, - streams: Optional[List['webrtc.MediaStream']] = None, - ): + direction: webrtc.TransceiverDirection | None = None, + send_encodings: list[webrtc.RTCRtpEncodingParameters] | None = None, + streams: list[webrtc.MediaStream] | None = None, + ) -> None: super().__init__() self.__send_encodings = [] @@ -46,36 +50,42 @@ def __init__( self.streams = streams @property - def direction(self) -> 'webrtc.TransceiverDirection': + def direction(self) -> webrtc.TransceiverDirection: """:obj:`webrtc.TransceiverDirection`: The new transceiver's preferred directionality. + This value is used to initialize the new :obj:`webrtc.RTCRtpTransceiver` object's - :attr:`webrtc.RTCRtpTransceiver.direction` property.""" + :attr:`webrtc.RTCRtpTransceiver.direction` property. + """ return self._native_obj.direction @direction.setter - def direction(self, value: 'webrtc.TransceiverDirection'): + def direction(self, value: webrtc.TransceiverDirection) -> None: self._native_obj.direction = value @property - def send_encodings(self) -> List['webrtc.RTCRtpEncodingParameters']: - """:obj:`list` of :obj:`webrtc.RTCRtpEncodingParameters`: A list of encodings to allow when sending RTP media - from the :obj:`webrtc.RTCRtpSender`, one per simulcast layer.""" + def send_encodings(self) -> list[webrtc.RTCRtpEncodingParameters]: + """:obj:`list` of :obj:`webrtc.RTCRtpEncodingParameters`: The encodings of the sender. + + The encodings to allow when sending RTP media from the :obj:`webrtc.RTCRtpSender`, one per simulcast layer. + """ return list(self.__send_encodings) @send_encodings.setter - def send_encodings(self, value: List['webrtc.RTCRtpEncodingParameters']): + def send_encodings(self, value: list[webrtc.RTCRtpEncodingParameters]) -> None: self.__send_encodings = list(value) self._native_obj.sendEncodings = [param._to_native() for param in value] @property - def streams(self) -> List['webrtc.MediaStream']: - """:obj:`list` of :obj:`webrtc.MediaStream`: A list of :obj:`webrtc.MediaStream` objects to add to the - transceiver's :obj:`webrtc.RTCRtpReceiver`; when the remote peer's :obj:`webrtc.RTCPeerConnection`'s track - event occurs, these are the streams that will be specified by that event.""" + def streams(self) -> list[webrtc.MediaStream]: + """:obj:`list` of :obj:`webrtc.MediaStream`: The streams of the track of the sender. + + When the remote peer's :obj:`webrtc.RTCPeerConnection`'s track event occurs, these are the streams that will be + specified by that event. + """ return self.__original_streams @streams.setter - def streams(self, value: List['webrtc.MediaStream']): + def streams(self, value: list[webrtc.MediaStream]) -> None: self.__original_streams = value self._native_obj.streamIds = [stream.id for stream in value] diff --git a/python-webrtc/python/webrtc/models/video_frame.py b/python-webrtc/python/webrtc/models/video_frame.py index a0e2659..1606b2f 100644 --- a/python-webrtc/python/webrtc/models/video_frame.py +++ b/python-webrtc/python/webrtc/models/video_frame.py @@ -7,23 +7,36 @@ """VideoFrame of WebCodecs (https://developer.mozilla.org/en-US/docs/Web/API/VideoFrame) and its dictionaries.""" +from __future__ import annotations + import asyncio import math import warnings from dataclasses import dataclass, fields -from typing import Any, Dict, List, NamedTuple, Optional, Tuple, Union +from enum import Enum +from typing import TYPE_CHECKING, Any, ClassVar, NamedTuple, TypeVar from webrtc import ( AlphaOption, InvalidStateError, NotSupportedError, + RTCException, VideoColorPrimaries, VideoMatrixCoefficients, VideoPixelFormat, VideoTransferCharacteristics, wrtc, ) -from webrtc.utils.names import alias, snake_case +from webrtc.models.closable import Closable +from webrtc.utils.names import Alias, alias, snake_case + +if TYPE_CHECKING: + from collections.abc import Iterable + + from typing_extensions import Buffer + +_EnumT = TypeVar('_EnumT', bound=Enum) +_InitT = TypeVar('_InitT') _MAX_UNSIGNED_LONG = 2**32 - 1 _RGB_FORMATS = (VideoPixelFormat.RGBA, VideoPixelFormat.RGBX, VideoPixelFormat.BGRA, VideoPixelFormat.BGRX) @@ -90,12 +103,12 @@ class VideoColorSpace: full_range (:obj:`bool`, optional): Whether the samples use the full range of their bits. """ - primaries: Optional[VideoColorPrimaries] = None - transfer: Optional[VideoTransferCharacteristics] = None - matrix: Optional[VideoMatrixCoefficients] = None - full_range: Optional[bool] = None + primaries: VideoColorPrimaries | None = None + transfer: VideoTransferCharacteristics | None = None + matrix: VideoMatrixCoefficients | None = None + full_range: bool | None = None - def to_json(self) -> Dict[str, Any]: + def to_json(self) -> dict[str, Any]: """Returns the members as a dictionary with the camelCase names, like ``toJSON()``.""" return { 'primaries': self.primaries, @@ -105,23 +118,23 @@ def to_json(self) -> Dict[str, Any]: } #: Alias for :attr:`full_range` - fullRange = alias('full_range') + fullRange: ClassVar[Alias[bool | None]] = alias('full_range') #: Alias for :meth:`to_json` - toJSON = to_json + toJSON: ClassVar = to_json _REC709 = VideoColorSpace( - VideoColorPrimaries.bt709, VideoTransferCharacteristics.bt709, VideoMatrixCoefficients.bt709, False + VideoColorPrimaries.bt709, VideoTransferCharacteristics.bt709, VideoMatrixCoefficients.bt709, full_range=False ) _SRGB = VideoColorSpace( - VideoColorPrimaries.bt709, VideoTransferCharacteristics.iec61966_2_1, VideoMatrixCoefficients.rgb, True + VideoColorPrimaries.bt709, VideoTransferCharacteristics.iec61966_2_1, VideoMatrixCoefficients.rgb, full_range=True ) # libwebrtc frames carry no color space: its software codecs (VP8, VP9, AV1) use BT.601 unless told otherwise _REC601 = VideoColorSpace( VideoColorPrimaries.smpte170m, VideoTransferCharacteristics.smpte170m, VideoMatrixCoefficients.smpte170m, - False, + full_range=False, ) @@ -133,10 +146,10 @@ class VideoFrameMetadata: rtp_timestamp (:obj:`int`, optional): The RTP timestamp of a frame received from a remote peer. """ - rtp_timestamp: Optional[int] = None + rtp_timestamp: int | None = None #: Alias for :attr:`rtp_timestamp` - rtpTimestamp = alias('rtp_timestamp') + rtpTimestamp: ClassVar[Alias[int | None]] = alias('rtp_timestamp') @dataclass @@ -163,27 +176,27 @@ class VideoFrameBufferInit: coded_width: int coded_height: int timestamp: int - duration: Optional[int] = None - layout: Optional[List[PlaneLayout]] = None - visible_rect: Optional[DOMRectReadOnly] = None + duration: int | None = None + layout: list[PlaneLayout] | None = None + visible_rect: DOMRectReadOnly | None = None rotation: float = 0 flip: bool = False - display_width: Optional[int] = None - display_height: Optional[int] = None - color_space: Optional[VideoColorSpace] = None + display_width: int | None = None + display_height: int | None = None + color_space: VideoColorSpace | None = None #: Alias for :attr:`coded_width` - codedWidth = alias('coded_width') + codedWidth: ClassVar[Alias[int]] = alias('coded_width') #: Alias for :attr:`coded_height` - codedHeight = alias('coded_height') + codedHeight: ClassVar[Alias[int]] = alias('coded_height') #: Alias for :attr:`visible_rect` - visibleRect = alias('visible_rect') + visibleRect: ClassVar[Alias[DOMRectReadOnly | None]] = alias('visible_rect') #: Alias for :attr:`display_width` - displayWidth = alias('display_width') + displayWidth: ClassVar[Alias[int | None]] = alias('display_width') #: Alias for :attr:`display_height` - displayHeight = alias('display_height') + displayHeight: ClassVar[Alias[int | None]] = alias('display_height') #: Alias for :attr:`color_space` - colorSpace = alias('color_space') + colorSpace: ClassVar[Alias[VideoColorSpace | None]] = alias('color_space') @dataclass @@ -201,21 +214,21 @@ class VideoFrameInit: display_height (:obj:`int`, optional): The height to show the frame at, with ``display_width``. """ - timestamp: Optional[int] = None - duration: Optional[int] = None + timestamp: int | None = None + duration: int | None = None alpha: AlphaOption = AlphaOption.keep - visible_rect: Optional[DOMRectReadOnly] = None + visible_rect: DOMRectReadOnly | None = None rotation: float = 0 flip: bool = False - display_width: Optional[int] = None - display_height: Optional[int] = None + display_width: int | None = None + display_height: int | None = None #: Alias for :attr:`visible_rect` - visibleRect = alias('visible_rect') + visibleRect: ClassVar[Alias[DOMRectReadOnly | None]] = alias('visible_rect') #: Alias for :attr:`display_width` - displayWidth = alias('display_width') + displayWidth: ClassVar[Alias[int | None]] = alias('display_width') #: Alias for :attr:`display_height` - displayHeight = alias('display_height') + displayHeight: ClassVar[Alias[int | None]] = alias('display_height') @dataclass @@ -229,9 +242,9 @@ class VideoFrameCopyToOptions: ``RGBA``, ``RGBX``, ``BGRA`` and ``BGRX``. """ - rect: Optional[DOMRectReadOnly] = None - layout: Optional[List[PlaneLayout]] = None - format: Optional[VideoPixelFormat] = None + rect: DOMRectReadOnly | None = None + layout: list[PlaneLayout] | None = None + format: VideoPixelFormat | None = None class _Plane(NamedTuple): @@ -239,19 +252,19 @@ class _Plane(NamedTuple): subsampling_x: int subsampling_y: int - def rows(self, rect: DOMRectReadOnly) -> Tuple[int, int]: - """The first row of the rect in the plane and the number of rows""" + def rows(self, rect: DOMRectReadOnly) -> tuple[int, int]: + """The first row of the rect in the plane and the number of rows.""" top = int(rect.y) // self.subsampling_y return top, -(-int(rect.y + rect.height) // self.subsampling_y) - top - def columns(self, rect: DOMRectReadOnly) -> Tuple[int, int]: - """The first byte of the rect in a row of the plane and the number of bytes""" + def columns(self, rect: DOMRectReadOnly) -> tuple[int, int]: + """The first byte of the rect in a row of the plane and the number of bytes.""" left = int(rect.x) // self.subsampling_x width = -(-int(rect.x + rect.width) // self.subsampling_x) - left return left * self.sample_bytes, width * self.sample_bytes -def _planes(format: VideoPixelFormat) -> List[_Plane]: +def _planes(format: VideoPixelFormat) -> list[_Plane]: if format in _RGB_FORMATS: return [_Plane(4, 1, 1)] if format == VideoPixelFormat.NV12: @@ -266,27 +279,28 @@ def _planes(format: VideoPixelFormat) -> List[_Plane]: def _has_alpha(format: VideoPixelFormat) -> bool: - return format in (VideoPixelFormat.RGBA, VideoPixelFormat.BGRA) or format.value[4:5] == 'A' + return format in {VideoPixelFormat.RGBA, VideoPixelFormat.BGRA} or format.value[4:5] == 'A' def _without_alpha(format: VideoPixelFormat) -> VideoPixelFormat: - if format in (VideoPixelFormat.RGBA, VideoPixelFormat.BGRA): + if format in {VideoPixelFormat.RGBA, VideoPixelFormat.BGRA}: return VideoPixelFormat(format.value[:3] + 'X') return VideoPixelFormat(format.value[:4] + format.value[5:]) -def _enum(cls: type, value: Any) -> Any: +def _enum(cls: type[_EnumT], value: object) -> _EnumT: try: return cls(value) except ValueError: - raise TypeError(f'{value!r} is not a {cls.__name__}') from None + msg = f'{value!r} is not a {cls.__name__}' + raise TypeError(msg) from None -def _optional_enum(cls: type, value: Any) -> Any: +def _optional_enum(cls: type[_EnumT], value: object) -> _EnumT | None: return None if value is None else _enum(cls, value) -def _is_buffer(value: Any) -> bool: +def _is_buffer(value: object) -> bool: try: memoryview(value) except TypeError: @@ -294,82 +308,94 @@ def _is_buffer(value: Any) -> bool: return True -def _buffer_size(data: Any) -> int: +def _buffer_size(data: Buffer) -> int: return memoryview(data).nbytes -def _dimension(value: Any, name: str) -> int: +def _dimension(value: object, name: str) -> int: if isinstance(value, bool) or not isinstance(value, int) or not 0 <= value <= _MAX_UNSIGNED_LONG: - raise TypeError(f'{name} must be an unsigned 32-bit integer, not {value!r}') + msg = f'{name} must be an unsigned 32-bit integer, not {value!r}' + raise TypeError(msg) return value -def _display_size(init: Any) -> Optional[Tuple[int, int]]: - """The display size of an init, if given""" +def _display_size(init: VideoFrameBufferInit | VideoFrameInit) -> tuple[int, int] | None: + """The display size of an init, if given.""" if (init.display_width is None) != (init.display_height is None): - raise TypeError('display_width and display_height go together') + msg = 'display_width and display_height go together' + raise TypeError(msg) if init.display_width is None: return None width, height = _dimension(init.display_width, 'display_width'), _dimension(init.display_height, 'display_height') if width == 0 or height == 0: - raise TypeError('The display size must be positive') + msg = 'The display size must be positive' + raise TypeError(msg) return width, height def _is_sideways(rotation: int) -> bool: - return rotation in (90, 270) + return rotation in {90, 270} -def _oriented(width: int, height: int, rotation: int) -> Tuple[int, int]: - """The size as shown after the rotation""" +def _oriented(width: int, height: int, rotation: int) -> tuple[int, int]: + """The size as shown after the rotation.""" return (height, width) if _is_sideways(rotation) else (width, height) -def _rect(value: Any) -> Optional[DOMRectReadOnly]: +def _rect(value: object) -> DOMRectReadOnly | None: if value is None or isinstance(value, DOMRectReadOnly): return value if isinstance(value, dict): - return DOMRectReadOnly(**{k: v for k, v in value.items() if k in ('x', 'y', 'width', 'height')}) - raise TypeError(f'{value!r} is not a DOMRectReadOnly') + return DOMRectReadOnly(**{k: v for k, v in value.items() if k in {'x', 'y', 'width', 'height'}}) + msg = f'{value!r} is not a DOMRectReadOnly' + raise TypeError(msg) -def _layout(value: Any) -> Optional[List[PlaneLayout]]: +def _layout(value: Iterable[PlaneLayout | dict[str, int]] | None) -> list[PlaneLayout] | None: if value is None: return None layout = [] - for plane in value: - if isinstance(plane, dict): - plane = PlaneLayout(plane['offset'], plane['stride']) + for item in value: + plane = PlaneLayout(item['offset'], item['stride']) if isinstance(item, dict) else item layout.append(PlaneLayout(_dimension(plane.offset, 'offset'), _dimension(plane.stride, 'stride'))) return layout def _rotation(value: float) -> int: - """The nearest multiple of 90, ties rounded up, from 0 to 270""" + """The nearest multiple of 90, ties rounded up, from 0 to 270.""" if not math.isfinite(value): - raise TypeError(f'The rotation must be finite, not {value!r}') + msg = f'The rotation must be finite, not {value!r}' + raise TypeError(msg) return int(math.floor(value / 90 + 0.5) * 90) % 360 +def _checked_rect(rect: DOMRectReadOnly, coded_size: tuple[int, int]) -> DOMRectReadOnly: + """A rect given for a frame, with integer coordinates.""" + if min(rect.width, rect.height) <= 0 or min(rect.x, rect.y) < 0: + msg = 'The rect must have a positive size and offset' + raise TypeError(msg) + coded_width, coded_height = coded_size + if rect.x + rect.width > coded_width or rect.y + rect.height > coded_height: + msg = 'The rect must be inside of the coded size' + raise TypeError(msg) + if any(not float(v).is_integer() for v in (rect.x, rect.y, rect.width, rect.height)): + msg = 'The rect must have integer coordinates' + raise TypeError(msg) + return DOMRectReadOnly(int(rect.x), int(rect.y), int(rect.width), int(rect.height)) + + def _parse_visible_rect( default: DOMRectReadOnly, - override: Optional[DOMRectReadOnly], - coded_width: int, - coded_height: int, + override: DOMRectReadOnly | None, + *, + coded_size: tuple[int, int], format: VideoPixelFormat, ) -> DOMRectReadOnly: - rect = default - if override is not None: - if override.width <= 0 or override.height <= 0 or override.x < 0 or override.y < 0: - raise TypeError('The rect must have a positive size and offset') - if override.x + override.width > coded_width or override.y + override.height > coded_height: - raise TypeError('The rect must be inside of the coded size') - if any(not float(v).is_integer() for v in (override.x, override.y, override.width, override.height)): - raise TypeError('The rect must have integer coordinates') - rect = DOMRectReadOnly(int(override.x), int(override.y), int(override.width), int(override.height)) + rect = default if override is None else _checked_rect(override, coded_size) for plane in _planes(format): if rect.x % plane.subsampling_x or rect.y % plane.subsampling_y: - raise TypeError(f'The rect must be aligned to the subsampling of {format.value}') + msg = f'The rect must be aligned to the subsampling of {format.value}' + raise TypeError(msg) return rect @@ -381,46 +407,58 @@ class _PlaneCopy(NamedTuple): offset: int stride: int + @property + def end(self) -> int: + """The byte after the plane in the buffer.""" + return self.offset + self.stride * self.height + class _CopyPlan(NamedTuple): format: VideoPixelFormat rect: DOMRectReadOnly size: int - planes: List[_PlaneCopy] + planes: list[_PlaneCopy] def _compute_layout( - rect: DOMRectReadOnly, format: VideoPixelFormat, layout: Optional[List[PlaneLayout]] -) -> Tuple[int, List[_PlaneCopy]]: - """Compute Layout and Allocation Size: the size of the buffer and where each plane goes in it""" + rect: DOMRectReadOnly, format: VideoPixelFormat, layout: list[PlaneLayout] | None +) -> tuple[int, list[_PlaneCopy]]: + """Compute Layout and Allocation Size: the size of the buffer and where each plane goes in it.""" planes = _planes(format) if layout is not None and len(layout) != len(planes): - raise TypeError(f'The layout must have {len(planes)} planes for {format.value}') + msg = f'The layout must have {len(planes)} planes for {format.value}' + raise TypeError(msg) allocation_size = 0 - copies: List[_PlaneCopy] = [] - ends: List[int] = [] + copies: list[_PlaneCopy] = [] for index, plane in enumerate(planes): - top, height = plane.rows(rect) - left_bytes, width_bytes = plane.columns(rect) - if layout is not None: - if layout[index].stride < width_bytes: - raise TypeError(f'The stride of plane {index} is smaller than its rows') - offset, stride = layout[index].offset, layout[index].stride - else: - offset, stride = allocation_size, width_bytes - end = offset + stride * height - if end > _MAX_UNSIGNED_LONG: - raise TypeError('The planes are too large') - for earlier, copy in enumerate(copies): - if copy.offset < end and offset < ends[earlier]: - raise TypeError(f'Planes {earlier} and {index} overlap') - ends.append(end) - allocation_size = max(allocation_size, end) - copies.append(_PlaneCopy(left_bytes, top, width_bytes, height, offset, stride)) + copy = _plane_copy(plane, rect, index, layout=layout, next_offset=allocation_size) + if copy.end > _MAX_UNSIGNED_LONG: + msg = 'The planes are too large' + raise TypeError(msg) + for earlier, other in enumerate(copies): + if other.offset < copy.end and copy.offset < other.end: + msg = f'Planes {earlier} and {index} overlap' + raise TypeError(msg) + allocation_size = max(allocation_size, copy.end) + copies.append(copy) return allocation_size, copies -def _color_space(value: Any) -> Optional[VideoColorSpace]: +def _plane_copy( + plane: _Plane, rect: DOMRectReadOnly, index: int, *, layout: list[PlaneLayout] | None, next_offset: int +) -> _PlaneCopy: + """Where a plane of a rect goes in a buffer: as the layout says, or packed at the next offset.""" + top, height = plane.rows(rect) + left_bytes, width_bytes = plane.columns(rect) + if layout is None: + return _PlaneCopy(left_bytes, top, width_bytes, height, next_offset, width_bytes) + if layout[index].stride < width_bytes: + msg = f'The stride of plane {index} is smaller than its rows' + raise TypeError(msg) + return _PlaneCopy(left_bytes, top, width_bytes, height, layout[index].offset, layout[index].stride) + + +def _color_space(value: object) -> VideoColorSpace | None: if value is None: return None if isinstance(value, dict): @@ -431,7 +469,8 @@ def _color_space(value: Any) -> Optional[VideoColorSpace]: full_range=value.get('full_range', value.get('fullRange')), ) if not isinstance(value, VideoColorSpace): - raise TypeError(f'{value!r} is not a VideoColorSpace') + msg = f'{value!r} is not a VideoColorSpace' + raise TypeError(msg) return VideoColorSpace( _optional_enum(VideoColorPrimaries, value.primaries), _optional_enum(VideoTransferCharacteristics, value.transfer), @@ -440,12 +479,13 @@ def _color_space(value: Any) -> Optional[VideoColorSpace]: ) -def _init_from(init: Any, cls: type, options: Dict[str, Any]): - """The init of a constructor: a dataclass, a dictionary (with snake_case or camelCase names), or keywords""" +def _init_from(init: object, cls: type[_InitT], options: dict[str, object]) -> _InitT: + """The init of a constructor: a dataclass, a dictionary (with snake_case or camelCase names), or keywords.""" if init is None: init = options elif options: - raise TypeError('Pass either an init or keyword arguments') + msg = 'Pass either an init or keyword arguments' + raise TypeError(msg) if isinstance(init, cls): return init if isinstance(init, dict): @@ -454,16 +494,61 @@ def _init_from(init: Any, cls: type, options: Dict[str, Any]): for key, value in init.items(): name = snake_case(key) if name not in names: - raise TypeError(f'{cls.__name__} has no member {key!r}') + msg = f'{cls.__name__} has no member {key!r}' + raise TypeError(msg) kwargs[name] = value try: return cls(**kwargs) except TypeError as e: - raise TypeError(f'Invalid {cls.__name__}: {e}') from None - raise TypeError(f'{init!r} is not a {cls.__name__}') + msg = f'Invalid {cls.__name__}: {e}' + raise TypeError(msg) from None + msg = f'{init!r} is not a {cls.__name__}' + raise TypeError(msg) + + +def _coded_size(init: VideoFrameBufferInit) -> tuple[int, int]: + size = _dimension(init.coded_width, 'coded_width'), _dimension(init.coded_height, 'coded_height') + if 0 in size: + msg = 'The coded size must be positive' + raise TypeError(msg) + return size + + +def _visible_resource( + data: Buffer, init: VideoFrameBufferInit, format: VideoPixelFormat, *, coded_size: tuple[int, int] +) -> tuple[wrtc.VideoFrameBuffer, tuple[int, int]]: + """The pixels of the visible rect of a buffer, which becomes the whole frame, and their size.""" + coded = DOMRectReadOnly(0, 0, *coded_size) + rect = _parse_visible_rect(coded, _rect(init.visible_rect), coded_size=coded_size, format=format) + allocation_size, copies = _compute_layout(coded, format, _layout(init.layout)) + if _buffer_size(data) < allocation_size: + msg = f'The data must be at least {allocation_size} bytes for this format and size' + raise TypeError(msg) + layout = [ + (copy.offset + plane.rows(rect)[0] * copy.stride + plane.columns(rect)[0], copy.stride) + for copy, plane in zip(copies, _planes(format)) + ] + size = int(rect.width), int(rect.height) + return wrtc.VideoFrameBuffer.fromData(format.value, *size, data, layout), size + + +class _Geometry(NamedTuple): + """How the pixels of a frame are shown.""" + + visible_rect: DOMRectReadOnly + display: tuple[int, int] + rotation: int + flip: bool + + +class _FrameInfo(NamedTuple): + timestamp: int + duration: int | None + color_space: VideoColorSpace + metadata: VideoFrameMetadata -class VideoFrame: +class VideoFrame(Closable): """A frame of video: its pixels and metadata (https://developer.mozilla.org/en-US/docs/Web/API/VideoFrame). A frame holds its pixels until :meth:`close`, which frames read from a track should be once used: a frame @@ -475,159 +560,134 @@ class VideoFrame: (required) or a frame. A dictionary of its members, or keyword arguments, can be passed instead. Raises: - :obj:`TypeError`: If the init isn't valid, or the buffer is too small for it. - :obj:`webrtc.InvalidStateError`: If the source frame is closed. + TypeError: If the init isn't valid, or the buffer is too small for it. + webrtc.InvalidStateError: If the source frame is closed. Example:: frame = webrtc.VideoFrame(i420, format='I420', coded_width=640, coded_height=480, timestamp=0) """ - def __init__(self, source: Any, init: Any = None, **options): + def __init__( + self, + source: Buffer | VideoFrame, + init: VideoFrameBufferInit | VideoFrameInit | dict[str, object] | None = None, + **options: object, + ) -> None: self._resource = None if isinstance(source, VideoFrame): self._init_from_frame(source, _init_from(init, VideoFrameInit, options)) elif _is_buffer(source): self._init_from_buffer(source, _init_from(init, VideoFrameBufferInit, options)) else: - raise TypeError(f'A VideoFrame is created from a buffer or a VideoFrame, not {type(source).__name__}') + msg = f'A VideoFrame is created from a buffer or a VideoFrame, not {type(source).__name__}' + raise TypeError(msg) - def _init_from_buffer(self, data: Any, init: VideoFrameBufferInit) -> None: + def _init_from_buffer(self, data: Buffer, init: VideoFrameBufferInit) -> None: format = _enum(VideoPixelFormat, init.format) - coded_width = _dimension(init.coded_width, 'coded_width') - coded_height = _dimension(init.coded_height, 'coded_height') - if coded_width == 0 or coded_height == 0: - raise TypeError('The coded size must be positive') + coded_size = _coded_size(init) display = _display_size(init) if not isinstance(init.timestamp, int) or isinstance(init.timestamp, bool): - raise TypeError('The timestamp is an integer of microseconds') - - coded = DOMRectReadOnly(0, 0, coded_width, coded_height) - rect = _parse_visible_rect(coded, _rect(init.visible_rect), coded_width, coded_height, format) - allocation_size, copies = _compute_layout(coded, format, _layout(init.layout)) - if _buffer_size(data) < allocation_size: - raise TypeError(f'The data must be at least {allocation_size} bytes for this format and size') - # only the visible rect is copied, which becomes the whole frame - layout = [] - for copy, plane in zip(copies, _planes(format)): - top, left_bytes = plane.rows(rect)[0], plane.columns(rect)[0] - layout.append((copy.offset + top * copy.stride + left_bytes, copy.stride)) - width, height = int(rect.width), int(rect.height) - resource = wrtc.VideoFrameBuffer.fromData(format.value, width, height, data, layout) + msg = 'The timestamp is an integer of microseconds' + raise TypeError(msg) + resource, (width, height) = _visible_resource(data, init, format, coded_size=coded_size) rotation = _rotation(init.rotation) color_space = _color_space(init.color_space) or (_SRGB if format in _RGB_FORMATS else _REC709) self._set( resource, format, - DOMRectReadOnly(0, 0, width, height), - display or _oriented(width, height, rotation), - rotation, - bool(init.flip), - init.timestamp, - init.duration, - color_space, + geometry=_Geometry( + DOMRectReadOnly(0, 0, width, height), + display or _oriented(width, height, rotation), + rotation, + flip=bool(init.flip), + ), + info=_FrameInfo(init.timestamp, init.duration, color_space, VideoFrameMetadata()), ) - def _init_from_frame(self, other: 'VideoFrame', init: VideoFrameInit) -> None: + def _init_from_frame(self, other: VideoFrame, init: VideoFrameInit) -> None: if other._resource is None: - raise InvalidStateError('The frame is closed') + msg = 'The frame is closed' + raise InvalidStateError(msg) display = _display_size(init) resource, format = other._resource, other._format if AlphaOption(init.alpha) == AlphaOption.discard and _has_alpha(format): resource, format = resource.withoutAlpha(), _without_alpha(format) override = _rect(init.visible_rect) - rect = _parse_visible_rect(other._visible_rect, override, other.coded_width, other.coded_height, format) + coded_size = (other.coded_width, other.coded_height) + rect = _parse_visible_rect(other._visible_rect, override, coded_size=coded_size, format=format) applied = _rotation(init.rotation) rotation = (other._rotation + (360 - applied if other._flip else applied)) % 360 - flip = other._flip != bool(init.flip) - - if display is None and override is not None: - # keep the scale of the source's visible rect to its display size - shown_width, shown_height = _oriented(*other._display, other._rotation) - width_scale = shown_width / other._visible_rect.width - height_scale = shown_height / other._visible_rect.height - width, height = round(rect.width * width_scale), round(rect.height * height_scale) - if width == 0 or height == 0: - raise TypeError('The display size would be zero') - display = _oriented(width, height, rotation) - elif display is None: - display = other._display - if _is_sideways(rotation) != _is_sideways(other._rotation): - display = (display[1], display[0]) + if display is None: + display = other._display_for(rect, rotation, scaled=override is not None) self._set( resource, format, - rect, - display, - rotation, - flip, - init.timestamp if init.timestamp is not None else other._timestamp, - init.duration if init.duration is not None else other._duration, - other._color_space, - other._metadata, + geometry=_Geometry(rect, display, rotation, flip=other._flip != bool(init.flip)), + info=_FrameInfo( + init.timestamp if init.timestamp is not None else other._timestamp, + init.duration if init.duration is not None else other._duration, + other._color_space, + other._metadata, + ), ) + def _display_for(self, rect: DOMRectReadOnly, rotation: int, *, scaled: bool) -> tuple[int, int]: + """The display size of a frame of these pixels with another visible rect and rotation.""" + if scaled: + # keep the scale of the visible rect to the display size + shown_width, shown_height = _oriented(*self._display, self._rotation) + width = round(rect.width * (shown_width / self._visible_rect.width)) + height = round(rect.height * (shown_height / self._visible_rect.height)) + if width == 0 or height == 0: + msg = 'The display size would be zero' + raise TypeError(msg) + return _oriented(width, height, rotation) + if _is_sideways(rotation) != _is_sideways(self._rotation): + return self._display[1], self._display[0] + return self._display + def _set( - self, - resource: 'wrtc.VideoFrameBuffer', - format: VideoPixelFormat, - rect: DOMRectReadOnly, - display: Tuple[int, int], - rotation: int, - flip: bool, - timestamp: int, - duration: Optional[int], - color_space: VideoColorSpace, - metadata: Optional[VideoFrameMetadata] = None, + self, resource: wrtc.VideoFrameBuffer, format: VideoPixelFormat, *, geometry: _Geometry, info: _FrameInfo ) -> None: self._resource = resource self._format = format - self._visible_rect = rect - self._display = display - self._rotation = rotation - self._flip = flip - self._timestamp = timestamp - self._duration = duration - self._color_space = color_space - self._metadata = metadata or VideoFrameMetadata() + self._visible_rect, self._display, self._rotation, self._flip = geometry + self._timestamp, self._duration, self._color_space, self._metadata = info @classmethod - def _from_native( - cls, resource: 'wrtc.VideoFrameBuffer', timestamp: int, rotation: int = 0, rtp_timestamp: Optional[int] = None - ) -> 'VideoFrame': - """A frame of a track""" + def _from_native(cls, native: tuple[wrtc.VideoFrameBuffer, int, int, int]) -> VideoFrame: + """A frame of a track: the pixels, timestamp, rotation and RTP timestamp (0 if unknown).""" + resource, timestamp, rotation, rtp_timestamp = native frame = cls.__new__(cls) width, height = resource.width, resource.height frame._set( resource, VideoPixelFormat(resource.format), - DOMRectReadOnly(0, 0, width, height), - _oriented(width, height, rotation), - rotation, - False, - timestamp, - None, - _REC601, - VideoFrameMetadata(rtp_timestamp), + geometry=_Geometry( + DOMRectReadOnly(0, 0, width, height), _oriented(width, height, rotation), rotation, flip=False + ), + info=_FrameInfo(timestamp, None, _REC601, VideoFrameMetadata(rtp_timestamp or None)), ) return frame - def _take_resource(self) -> 'wrtc.VideoFrameBuffer': - """The pixels, for a generator, which closes the frame""" + def _take_resource(self) -> wrtc.VideoFrameBuffer: + """The pixels, for a generator, which closes the frame.""" if self._resource is None: - raise InvalidStateError('The frame is closed') + msg = 'The frame is closed' + raise InvalidStateError(msg) resource = self._resource self._resource = None return resource - def __del__(self): + def __del__(self) -> None: if getattr(self, '_resource', None) is not None: warnings.warn('A VideoFrame was garbage collected without being closed', ResourceWarning, stacklevel=2) @property - def format(self) -> Optional[VideoPixelFormat]: + def format(self) -> VideoPixelFormat | None: """:obj:`webrtc.VideoPixelFormat`, optional: The layout of the pixels, :obj:`None` once closed.""" return self._format if self._resource is not None else None @@ -642,14 +702,14 @@ def coded_height(self) -> int: return self._resource.height if self._resource is not None else 0 @property - def coded_rect(self) -> Optional[DOMRectReadOnly]: + def coded_rect(self) -> DOMRectReadOnly | None: """:obj:`DOMRectReadOnly`, optional: The rect of all the pixels, :obj:`None` once closed.""" if self._resource is None: return None return DOMRectReadOnly(0, 0, self._resource.width, self._resource.height) @property - def visible_rect(self) -> Optional[DOMRectReadOnly]: + def visible_rect(self) -> DOMRectReadOnly | None: """:obj:`DOMRectReadOnly`, optional: The part of the pixels to show, :obj:`None` once closed.""" return self._visible_rect if self._resource is not None else None @@ -679,7 +739,7 @@ def timestamp(self) -> int: return self._timestamp @property - def duration(self) -> Optional[int]: + def duration(self) -> int | None: """:obj:`int`, optional: The duration in microseconds.""" return self._duration @@ -692,41 +752,43 @@ def metadata(self) -> VideoFrameMetadata: """Returns what else is known of the frame, like the RTP timestamp of a received frame. Raises: - :obj:`webrtc.InvalidStateError`: If the frame is closed. + webrtc.InvalidStateError: If the frame is closed. """ if self._resource is None: - raise InvalidStateError('The frame is closed') + msg = 'The frame is closed' + raise InvalidStateError(msg) return VideoFrameMetadata(self._metadata.rtp_timestamp) - def _plan_copy(self, options: Any) -> _CopyPlan: + def _plan_copy(self, options: VideoFrameCopyToOptions | dict[str, object] | None) -> _CopyPlan: if self._resource is None: - raise InvalidStateError('The frame is closed') + msg = 'The frame is closed' + raise InvalidStateError(msg) options = _init_from(options, VideoFrameCopyToOptions, {}) format = self._format if options.format is not None: format = _enum(VideoPixelFormat, options.format) if format != self._format and format not in _RGB_FORMATS: - raise NotSupportedError(f'Frames are converted to RGB formats only, not {format.value}') - rect = _parse_visible_rect( - self._visible_rect, _rect(options.rect), self.coded_width, self.coded_height, self._format - ) + msg = f'Frames are converted to RGB formats only, not {format.value}' + raise NotSupportedError(msg) + coded_size = (self.coded_width, self.coded_height) + rect = _parse_visible_rect(self._visible_rect, _rect(options.rect), coded_size=coded_size, format=self._format) size, planes = _compute_layout(rect, format, _layout(options.layout)) return _CopyPlan(format, rect, size, planes) - def allocation_size(self, options: Any = None) -> int: + def allocation_size(self, options: VideoFrameCopyToOptions | dict[str, object] | None = None) -> int: """Returns how many bytes :meth:`copy_to` needs. + Raises :obj:`webrtc.InvalidStateError` if the frame is closed, :obj:`TypeError` if the options aren't valid + and :obj:`webrtc.NotSupportedError` if the frame can't be converted to the format. + Args: options (:obj:`VideoFrameCopyToOptions`, optional): How the frame is copied. - - Raises: - :obj:`webrtc.InvalidStateError`: If the frame is closed. - :obj:`TypeError`: If the options aren't valid. - :obj:`webrtc.NotSupportedError`: If the frame can't be converted to the format. """ return self._plan_copy(options).size - def copy_to(self, destination: Union[bytearray, memoryview], options: Any = None) -> asyncio.Future: + def copy_to( + self, destination: bytearray | memoryview, options: VideoFrameCopyToOptions | dict[str, object] | None = None + ) -> asyncio.Future[list[PlaneLayout]]: """Copies the pixels into a buffer. Args: @@ -740,14 +802,17 @@ def copy_to(self, destination: Union[bytearray, memoryview], options: Any = None future = asyncio.get_running_loop().create_future() try: future.set_result(self._copy_to(destination, options)) - except Exception as e: + except (TypeError, ValueError, LookupError, RuntimeError, RTCException) as e: future.set_exception(e) return future - def _copy_to(self, destination: Any, options: Any) -> List[PlaneLayout]: + def _copy_to( + self, destination: bytearray | memoryview, options: VideoFrameCopyToOptions | dict[str, object] | None + ) -> list[PlaneLayout]: plan = self._plan_copy(options) if _buffer_size(destination) < plan.size: - raise TypeError(f'The destination must be at least {plan.size} bytes') + msg = f'The destination must be at least {plan.size} bytes' + raise TypeError(msg) if plan.format == self._format: self._resource.copyPlanes(destination, [tuple(plane) for plane in plan.planes]) else: @@ -767,14 +832,15 @@ def _copy_to(self, destination: Any, options: Any) -> List[PlaneLayout]: ) return [PlaneLayout(plane.offset, plane.stride) for plane in plan.planes] - def clone(self) -> 'VideoFrame': + def clone(self) -> VideoFrame: """Returns another frame of the same pixels, which is closed separately. Raises: - :obj:`webrtc.InvalidStateError`: If the frame is closed. + webrtc.InvalidStateError: If the frame is closed. """ if self._resource is None: - raise InvalidStateError('The frame is closed') + msg = 'The frame is closed' + raise InvalidStateError(msg) frame = VideoFrame.__new__(VideoFrame) frame.__dict__.update(self.__dict__) return frame @@ -783,13 +849,7 @@ def close(self) -> None: """Releases the pixels. Closing a closed frame does nothing.""" self._resource = None - def __enter__(self) -> 'VideoFrame': - return self - - def __exit__(self, *exc_info) -> None: - self.close() - - def __repr__(self): + def __repr__(self) -> str: if self._resource is None: return '' return ( diff --git a/python-webrtc/python/webrtc/streams.py b/python-webrtc/python/webrtc/streams.py index 3fadf18..26410d9 100644 --- a/python-webrtc/python/webrtc/streams.py +++ b/python-webrtc/python/webrtc/streams.py @@ -5,26 +5,33 @@ # that can be found in the LICENSE.md file in the root of the project. # -"""The part of WHATWG Streams (https://streams.spec.whatwg.org) that media processing uses: readable, writable and -transform streams of objects. Methods return futures like the promises of the specification, so a read or a write -is requested when it's called, not when it's awaited. They need a running asyncio event loop.""" +"""The part of WHATWG Streams (https://streams.spec.whatwg.org) that media processing uses. + +Readable, writable and transform streams of objects. Methods return futures like the promises of the specification, +so a read or a write is requested when it's called, not when it's awaited. They need a running asyncio event loop. +""" + +from __future__ import annotations import asyncio import collections import inspect from dataclasses import dataclass -from typing import Any, AsyncIterator, Callable, Deque, Optional, Set, Tuple +from typing import TYPE_CHECKING, Any, Callable, NamedTuple, Protocol + +if TYPE_CHECKING: + from collections.abc import AsyncIterator __all__ = [ 'ReadableStream', - 'ReadableStreamDefaultReader', 'ReadableStreamDefaultController', + 'ReadableStreamDefaultReader', 'ReadableStreamReadResult', - 'WritableStream', - 'WritableStreamDefaultWriter', - 'WritableStreamDefaultController', 'TransformStream', 'TransformStreamDefaultController', + 'WritableStream', + 'WritableStreamDefaultController', + 'WritableStreamDefaultWriter', ] @@ -32,21 +39,22 @@ def _loop() -> asyncio.AbstractEventLoop: try: return asyncio.get_running_loop() except RuntimeError: - raise RuntimeError('streams are used from a running asyncio event loop') from None + msg = 'streams are used from a running asyncio event loop' + raise RuntimeError(msg) from None def _pending() -> asyncio.Future: return _loop().create_future() -def _resolved(value: Any = None) -> asyncio.Future: +def _resolved(value: object = None) -> asyncio.Future: future = _pending() future.set_result(value) return future # pipes running, see ReadableStream.pipe_to -_running_pipes: Set[asyncio.Future] = set() +_running_pipes: set[asyncio.Future] = set() def _rejected(error: BaseException) -> asyncio.Future: @@ -61,7 +69,7 @@ def _handled(future: asyncio.Future) -> asyncio.Future: return future -def _settle(future: Optional[asyncio.Future], value: Any = None, error: Optional[BaseException] = None) -> None: +def _settle(future: asyncio.Future | None, value: object = None, error: BaseException | None = None) -> None: if future is None or future.done(): return if error is not None: @@ -78,38 +86,38 @@ def _reject(future: asyncio.Future, error: BaseException) -> asyncio.Future: return future -def _error_or_default(error: Optional[BaseException]) -> BaseException: +def _error_or_default(error: BaseException | None) -> BaseException: return error if error is not None else TypeError('The stream errored') -def _reason_error(reason: Any) -> BaseException: +def _reason_error(reason: object) -> BaseException: return reason if isinstance(reason, BaseException) else TypeError(str(reason)) -def _member(obj: Any, name: str) -> Any: - """A method of an underlying source, sink or transformer: an object, or a dictionary as in browsers""" +def _member(obj: object, name: str) -> Callable[..., object] | None: + """A method of an underlying source, sink or transformer: an object, or a dictionary as in browsers.""" if isinstance(obj, dict): return obj.get(name) return getattr(obj, name, None) if obj is not None else None -def _call(obj: Any, name: str, *args) -> Any: +def _call(obj: object, name: str, *args: object) -> object: """Calls a method of an underlying source, sink or transformer, if it has one.""" method = _member(obj, name) return method(*args) if method is not None else None -async def _await(result: Any) -> Any: +async def _await(result: object) -> object: return await result if inspect.isawaitable(result) else result -def _then(result: Any, on_done: Callable[[], None], on_error: Callable[[BaseException], None]) -> None: +def _then(result: object, on_done: Callable[[], None], on_error: Callable[[BaseException], None]) -> None: """Runs a callback once the result of an algorithm (a value or an awaitable) settles.""" if not inspect.isawaitable(result): on_done() return - def done(future: asyncio.Future): + def done(future: asyncio.Future) -> None: error = asyncio.CancelledError() if future.cancelled() else future.exception() if error is not None: on_error(error) @@ -119,13 +127,44 @@ def done(future: asyncio.Future): asyncio.ensure_future(result).add_done_callback(done) -def _future_of(result: Any) -> asyncio.Future: +def _future_of(result: object) -> asyncio.Future: """Returns a future settled like the result of an algorithm (a value or an awaitable).""" future = _pending() _then(result, lambda: _settle(future), lambda e: _settle(future, error=e)) return future +def _run( + obj: object, + name: str, + *args: object, + on_done: Callable[[], None], + on_error: Callable[[BaseException], None], +) -> None: + """Calls a method of an underlying source or sink and a callback once its result settles, or once it raises.""" + try: + result = _call(obj, name, *args) + except Exception as e: + # anything the method raises errors the stream, as a rejection in the specification + on_error(e) + return + _then(result, on_done, on_error) + + +class _ReadableWritablePair(Protocol): + @property + def readable(self) -> ReadableStream: ... + + @property + def writable(self) -> WritableStream: ... + + +class _PipeOptions(NamedTuple): + prevent_close: bool + prevent_abort: bool + prevent_cancel: bool + + @dataclass class ReadableStreamReadResult: """The result of :meth:`ReadableStreamDefaultReader.read`. @@ -142,18 +181,18 @@ class ReadableStreamReadResult: class ReadableStreamDefaultController: """Lets an underlying source enqueue chunks, close or error its stream.""" - def __init__(self, stream: 'ReadableStream', source: Any, high_water_mark: float): + def __init__(self, stream: ReadableStream, source: object, high_water_mark: float) -> None: self._stream = stream self._source = source self._high_water_mark = high_water_mark - self._queue: Deque[Any] = collections.deque() + self._queue: collections.deque[Any] = collections.deque() self._close_requested = False self._started = False self._pulling = False self._pull_again = False @property - def desired_size(self) -> Optional[float]: + def desired_size(self) -> float | None: """:obj:`float`, optional: How many chunks the queue can take until it's full, :obj:`None` if errored.""" state = self._stream._state if state == 'errored': @@ -162,17 +201,17 @@ def desired_size(self) -> Optional[float]: return 0 return self._high_water_mark - len(self._queue) - def enqueue(self, chunk: Any) -> None: + def enqueue(self, chunk: object) -> None: """Enqueues a chunk, which fulfills a pending read if there's one. Raises: - :obj:`TypeError`: If the stream is closed or closing. + TypeError: If the stream is closed or closing. """ if self._close_requested or self._stream._state != 'readable': - raise TypeError('The stream is closed or closing') - reader = self._stream._reader - if reader is not None and reader._read_requests: - _settle(reader._read_requests.popleft(), ReadableStreamReadResult(chunk, False)) + msg = 'The stream is closed or closing' + raise TypeError(msg) + if self._stream._has_read_requests(): + _settle(self._stream._reader._read_requests.popleft(), ReadableStreamReadResult(chunk, done=False)) else: self._queue.append(chunk) self._call_pull_if_needed() @@ -181,15 +220,16 @@ def close(self) -> None: """Closes the stream once its queue is read. Raises: - :obj:`TypeError`: If the stream is closed or closing. + TypeError: If the stream is closed or closing. """ if self._close_requested or self._stream._state != 'readable': - raise TypeError('The stream is closed or closing') + msg = 'The stream is closed or closing' + raise TypeError(msg) self._close_requested = True if not self._queue: self._stream._close() - def error(self, error: Optional[BaseException] = None) -> None: + def error(self, error: BaseException | None = None) -> None: """Errors the stream: pending and later reads fail with the error.""" if self._stream._state != 'readable': return @@ -197,7 +237,7 @@ def error(self, error: Optional[BaseException] = None) -> None: self._stream._error(_error_or_default(error)) def _start(self) -> None: - def started(): + def started() -> None: self._started = True self._call_pull_if_needed() @@ -207,7 +247,7 @@ def _should_call_pull(self) -> bool: stream = self._stream if not self._started or self._close_requested or stream._state != 'readable': return False - if stream._reader is not None and stream._reader._read_requests: + if stream._has_read_requests(): return True return self.desired_size > 0 @@ -219,18 +259,13 @@ def _call_pull_if_needed(self) -> None: return self._pulling = True - def pulled(): + def pulled() -> None: self._pulling = False if self._pull_again: self._pull_again = False self._call_pull_if_needed() - try: - result = _call(self._source, 'pull', self) - except Exception as e: - self.error(e) - return - _then(result, pulled, self.error) + _run(self._source, 'pull', self, on_done=pulled, on_error=self.error) def _read(self, request: asyncio.Future) -> None: if self._queue: @@ -239,12 +274,12 @@ def _read(self, request: asyncio.Future) -> None: self._stream._close() else: self._call_pull_if_needed() - _settle(request, ReadableStreamReadResult(chunk, False)) + _settle(request, ReadableStreamReadResult(chunk, done=False)) else: self._stream._reader._read_requests.append(request) self._call_pull_if_needed() - def _cancel(self, reason: Any) -> Any: + def _cancel(self, reason: object) -> object: self._queue.clear() return _call(self._source, 'cancel', reason) @@ -261,10 +296,10 @@ class ReadableStream: high_water_mark (:obj:`float`, optional): How many chunks are queued ahead of reads, 1 by default. """ - def __init__(self, underlying_source: Any = None, high_water_mark: float = 1): + def __init__(self, underlying_source: object = None, high_water_mark: float = 1) -> None: self._state = 'readable' - self._stored_error: Optional[BaseException] = None - self._reader: Optional[ReadableStreamDefaultReader] = None + self._stored_error: BaseException | None = None + self._reader: ReadableStreamDefaultReader | None = None self._controller = ReadableStreamDefaultController(self, underlying_source, high_water_mark) self._controller._start() @@ -273,15 +308,14 @@ def locked(self) -> bool: """:obj:`bool`: Whether a reader holds the stream.""" return self._reader is not None - def get_reader(self) -> 'ReadableStreamDefaultReader': + def get_reader(self) -> ReadableStreamDefaultReader: """Returns a reader, which holds the stream until it's released. - Raises: - :obj:`TypeError`: If the stream is locked. + Raises :obj:`TypeError` if the stream is locked. """ return ReadableStreamDefaultReader(self) - def cancel(self, reason: Any = None) -> asyncio.Future: + def cancel(self, reason: object = None) -> asyncio.Future: """Cancels the stream: its source stops and its chunks are dropped. Returns: @@ -293,7 +327,8 @@ def cancel(self, reason: Any = None) -> asyncio.Future: def pipe_to( self, - destination: 'WritableStream', + destination: WritableStream, + *, prevent_close: bool = False, prevent_abort: bool = False, prevent_cancel: bool = False, @@ -313,7 +348,8 @@ def pipe_to( return _rejected(TypeError('A stream is locked')) reader = self.get_reader() writer = destination.get_writer() - pipe = asyncio.ensure_future(self._pipe(reader, writer, prevent_close, prevent_abort, prevent_cancel)) + options = _PipeOptions(prevent_close, prevent_abort, prevent_cancel) + pipe = asyncio.ensure_future(self._pipe(reader, writer, options)) # kept until done, as in browsers: asyncio keeps tasks weakly _running_pipes.add(pipe) pipe.add_done_callback(_running_pipes.discard) @@ -321,39 +357,22 @@ def pipe_to( @staticmethod async def _pipe( - reader: 'ReadableStreamDefaultReader', - writer: 'WritableStreamDefaultWriter', - prevent_close: bool, - prevent_abort: bool, - prevent_cancel: bool, + reader: ReadableStreamDefaultReader, writer: WritableStreamDefaultWriter, options: _PipeOptions ) -> None: try: - while True: - await writer.ready - result = await reader.read() - if result.done: - if not prevent_close: - await writer.close() - return - # writes aren't awaited, like in the specification - _handled(writer.write(result.value)) + await _pipe_chunks(reader, writer, prevent_close=options.prevent_close) except GeneratorExit: # closed, at exit: nothing can be awaited anymore raise except BaseException as e: - if writer._stream._state in ('erroring', 'errored'): - if not prevent_cancel: - await _handled(reader.cancel(e)) - elif not prevent_abort: - await _handled(writer.abort(e)) + await _stop_pipe(reader, writer, e, options=options) raise finally: reader.release_lock() writer.release_lock() - def pipe_through(self, transform: Any, **options) -> 'ReadableStream': - """Pipes the stream into the writable side of a transform (like :obj:`TransformStream`) and returns its - readable side. + def pipe_through(self, transform: _ReadableWritablePair, **options: bool) -> ReadableStream: + """Pipes the stream into the writable side of a transform (like :obj:`TransformStream`). Args: transform: An object with ``writable`` and ``readable`` streams. @@ -365,14 +384,19 @@ def pipe_through(self, transform: Any, **options) -> 'ReadableStream': _handled(self.pipe_to(transform.writable, **options)) return transform.readable - def values(self, prevent_cancel: bool = False) -> AsyncIterator[Any]: - """Iterates over the chunks, like ``async for``. Stopping early cancels the stream once the iterator is - finalized, right away with :func:`contextlib.aclosing`. + def values(self, *, prevent_cancel: bool = False) -> AsyncIterator[Any]: + """Iterates over the chunks, like ``async for``. + + Stopping early cancels the stream once the iterator is finalized, right away with + :func:`contextlib.aclosing`. Args: prevent_cancel (:obj:`bool`, optional): Whether the stream is left open when the iteration stops early. + + Returns: + An asynchronous iterator of the chunks. """ - return _iterate(self.get_reader(), prevent_cancel) + return _iterate(self.get_reader(), prevent_cancel=prevent_cancel) def __aiter__(self) -> AsyncIterator[Any]: return self.values() @@ -382,7 +406,7 @@ def _close(self) -> None: return self._state = 'closed' if self._reader is not None: - self._reader._settle_read_requests(ReadableStreamReadResult(None, True)) + self._reader._settle_read_requests(ReadableStreamReadResult(None, done=True)) _settle(self._reader._closed) def _error(self, error: BaseException) -> None: @@ -394,7 +418,10 @@ def _error(self, error: BaseException) -> None: self._reader._settle_read_requests(error=error) _settle(self._reader._closed, error=error) - def _cancel(self, reason: Any) -> asyncio.Future: + def _has_read_requests(self) -> bool: + return self._reader is not None and bool(self._reader._read_requests) + + def _cancel(self, reason: object) -> asyncio.Future: if self._state == 'closed': return _resolved() if self._state == 'errored': @@ -417,14 +444,15 @@ class ReadableStreamDefaultReader: stream (:obj:`ReadableStream`): The stream. Raises: - :obj:`TypeError`: If the stream is locked. + TypeError: If the stream is locked. """ - def __init__(self, stream: ReadableStream): + def __init__(self, stream: ReadableStream) -> None: if stream.locked: - raise TypeError('The stream is locked') - self._stream: Optional[ReadableStream] = stream - self._read_requests: Deque[asyncio.Future] = collections.deque() + msg = 'The stream is locked' + raise TypeError(msg) + self._stream: ReadableStream | None = stream + self._read_requests: collections.deque[asyncio.Future] = collections.deque() self._closed = _handled(_pending()) stream._reader = self if stream._state == 'closed': @@ -448,14 +476,14 @@ def read(self) -> asyncio.Future: if stream is None: return _rejected(TypeError('The reader is released')) if stream._state == 'closed': - return _resolved(ReadableStreamReadResult(None, True)) + return _resolved(ReadableStreamReadResult(None, done=True)) if stream._state == 'errored': return _rejected(stream._stored_error) request = _pending() stream._controller._read(request) return request - def cancel(self, reason: Any = None) -> asyncio.Future: + def cancel(self, reason: object = None) -> asyncio.Future: """Cancels the stream (see :meth:`ReadableStream.cancel`).""" if self._stream is None: return _rejected(TypeError('The reader is released')) @@ -475,7 +503,7 @@ def release_lock(self) -> None: stream._reader = None self._stream = None - def _settle_read_requests(self, value: Any = None, error: Optional[BaseException] = None) -> None: + def _settle_read_requests(self, value: object = None, error: BaseException | None = None) -> None: while self._read_requests: _settle(self._read_requests.popleft(), value, error) @@ -483,7 +511,36 @@ def _settle_read_requests(self, value: Any = None, error: Optional[BaseException releaseLock = release_lock -async def _iterate(reader: ReadableStreamDefaultReader, prevent_cancel: bool) -> AsyncIterator[Any]: +async def _pipe_chunks( + reader: ReadableStreamDefaultReader, writer: WritableStreamDefaultWriter, *, prevent_close: bool +) -> None: + while True: + await writer.ready + result = await reader.read() + if result.done: + if not prevent_close: + await writer.close() + return + # writes aren't awaited, like in the specification + _handled(writer.write(result.value)) + + +async def _stop_pipe( + reader: ReadableStreamDefaultReader, + writer: WritableStreamDefaultWriter, + error: BaseException, + *, + options: _PipeOptions, +) -> None: + """Cancels the source or aborts the destination of a pipe that failed, unless the options prevent it.""" + if writer._stream._state in {'erroring', 'errored'}: + if not options.prevent_cancel: + await _handled(reader.cancel(error)) + elif not options.prevent_abort: + await _handled(writer.abort(error)) + + +async def _iterate(reader: ReadableStreamDefaultReader, *, prevent_cancel: bool) -> AsyncIterator[Any]: done = False try: while True: @@ -507,16 +564,16 @@ async def _iterate(reader: ReadableStreamDefaultReader, prevent_cancel: bool) -> class WritableStreamDefaultController: """Lets an underlying sink error its stream.""" - def __init__(self, stream: 'WritableStream', sink: Any, high_water_mark: float): + def __init__(self, stream: WritableStream, sink: object, high_water_mark: float) -> None: self._stream = stream self._sink = sink self._high_water_mark = high_water_mark # (chunk, future) of the writes, then (_CLOSE, future) - self._queue: Deque[Tuple[Any, asyncio.Future]] = collections.deque() + self._queue: collections.deque[tuple[object, asyncio.Future]] = collections.deque() self._started = False self._in_flight = False - def error(self, error: Optional[BaseException] = None) -> None: + def error(self, error: BaseException | None = None) -> None: """Errors the stream: pending and later writes fail with the error.""" if self._stream._state == 'writable': self._stream._start_erroring(_error_or_default(error)) @@ -525,17 +582,17 @@ def _desired_size(self) -> float: return self._high_water_mark - sum(1 for chunk, _ in self._queue if chunk is not _CLOSE) def _start(self) -> None: - def started(): + def started() -> None: self._started = True self._advance() - def failed(error: BaseException): + def failed(error: BaseException) -> None: self._started = True self._stream._deal_with_rejection(error) _then(_call(self._sink, 'start', self), started, failed) - def _write(self, chunk: Any, future: asyncio.Future) -> None: + def _write(self, chunk: object, future: asyncio.Future) -> None: self._queue.append((chunk, future)) # advance first: a sink done right away leaves no backpressure to signal self._advance() @@ -557,32 +614,31 @@ def _advance(self) -> None: chunk, future = self._queue[0] self._in_flight = True - def finish(): - self._in_flight = False - self._queue.popleft() - - def succeeded(): - finish() - _settle(future) - if chunk is _CLOSE: - stream._state = 'closed' - if stream._writer is not None: - _settle(stream._writer._closed) - else: - stream._update_backpressure() - self._advance() - - def failed(error: BaseException): - finish() - _settle(future, error=error) + def failed(error: BaseException) -> None: + self._settle_in_flight(future, error) stream._deal_with_rejection(error) - try: - result = _call(self._sink, 'close') if chunk is _CLOSE else _call(self._sink, 'write', chunk, self) - except Exception as e: - failed(e) - return - _then(result, succeeded, failed) + if chunk is _CLOSE: + _run(self._sink, 'close', on_done=lambda: self._closed(future), on_error=failed) + else: + _run(self._sink, 'write', chunk, self, on_done=lambda: self._written(future), on_error=failed) + + def _settle_in_flight(self, future: asyncio.Future, error: BaseException | None = None) -> None: + self._in_flight = False + self._queue.popleft() + _settle(future, error=error) + + def _written(self, future: asyncio.Future) -> None: + self._settle_in_flight(future) + self._stream._update_backpressure() + self._advance() + + def _closed(self, future: asyncio.Future) -> None: + self._settle_in_flight(future) + stream = self._stream + stream._state = 'closed' + if stream._writer is not None: + _settle(stream._writer._closed) def _reject_queue(self, error: BaseException) -> None: """Rejects the queued requests but the one in flight.""" @@ -604,10 +660,10 @@ class WritableStream: 1 by default. """ - def __init__(self, underlying_sink: Any = None, high_water_mark: float = 1): + def __init__(self, underlying_sink: object = None, high_water_mark: float = 1) -> None: self._state = 'writable' - self._stored_error: Optional[BaseException] = None - self._writer: Optional[WritableStreamDefaultWriter] = None + self._stored_error: BaseException | None = None + self._writer: WritableStreamDefaultWriter | None = None self._close_requested = False self._controller = WritableStreamDefaultController(self, underlying_sink, high_water_mark) self._controller._start() @@ -617,11 +673,10 @@ def locked(self) -> bool: """:obj:`bool`: Whether a writer holds the stream.""" return self._writer is not None - def get_writer(self) -> 'WritableStreamDefaultWriter': + def get_writer(self) -> WritableStreamDefaultWriter: """Returns a writer, which holds the stream until it's released. - Raises: - :obj:`TypeError`: If the stream is locked. + Raises :obj:`TypeError` if the stream is locked. """ return WritableStreamDefaultWriter(self) @@ -635,7 +690,7 @@ def close(self) -> asyncio.Future: return _rejected(TypeError('The stream is locked')) return self._close() - def abort(self, reason: Any = None) -> asyncio.Future: + def abort(self, reason: object = None) -> asyncio.Future: """Aborts the stream: queued chunks are dropped and the sink is aborted. Returns: @@ -646,15 +701,15 @@ def abort(self, reason: Any = None) -> asyncio.Future: return self._abort(reason) def _close(self) -> asyncio.Future: - if self._state in ('closed', 'errored') or self._close_requested: + if self._state in {'closed', 'errored'} or self._close_requested: return _rejected(TypeError('The stream is closed or closing')) self._close_requested = True future = _pending() self._controller._close(future) return future - def _abort(self, reason: Any) -> asyncio.Future: - if self._state in ('closed', 'errored'): + def _abort(self, reason: object) -> asyncio.Future: + if self._state in {'closed', 'errored'}: return _resolved() error = reason if isinstance(reason, BaseException) else asyncio.CancelledError(reason) self._controller._reject_queue(error) @@ -712,13 +767,14 @@ class WritableStreamDefaultWriter: stream (:obj:`WritableStream`): The stream. Raises: - :obj:`TypeError`: If the stream is locked. + TypeError: If the stream is locked. """ - def __init__(self, stream: WritableStream): + def __init__(self, stream: WritableStream) -> None: if stream.locked: - raise TypeError('The stream is locked') - self._stream: Optional[WritableStream] = stream + msg = 'The stream is locked' + raise TypeError(msg) + self._stream: WritableStream | None = stream stream._writer = self self._closed = _handled(_pending()) self._ready = _handled(_pending()) @@ -743,22 +799,23 @@ def ready(self) -> asyncio.Future: return self._ready @property - def desired_size(self) -> Optional[float]: + def desired_size(self) -> float | None: """:obj:`float`, optional: How many chunks can be written until the queue is full. Raises: - :obj:`TypeError`: If the writer is released. + TypeError: If the writer is released. """ stream = self._stream if stream is None: - raise TypeError('The writer is released') - if stream._state in ('errored', 'erroring'): + msg = 'The writer is released' + raise TypeError(msg) + if stream._state in {'errored', 'erroring'}: return None if stream._state == 'closed': return 0 return stream._controller._desired_size() - def write(self, chunk: Any) -> asyncio.Future: + def write(self, chunk: object) -> asyncio.Future: """Writes a chunk. Returns: @@ -767,7 +824,7 @@ def write(self, chunk: Any) -> asyncio.Future: stream = self._stream if stream is None: return _rejected(TypeError('The writer is released')) - if stream._state in ('errored', 'erroring'): + if stream._state in {'errored', 'erroring'}: return _rejected(stream._stored_error) if stream._close_requested or stream._state == 'closed': return _rejected(TypeError('The stream is closed or closing')) @@ -781,7 +838,7 @@ def close(self) -> asyncio.Future: return _rejected(TypeError('The writer is released')) return self._stream._close() - def abort(self, reason: Any = None) -> asyncio.Future: + def abort(self, reason: object = None) -> asyncio.Future: """Aborts the stream (see :meth:`WritableStream.abort`).""" if self._stream is None: return _rejected(TypeError('The writer is released')) @@ -807,19 +864,19 @@ def release_lock(self) -> None: class TransformStreamDefaultController: """Lets a transformer enqueue chunks to the readable side, error or terminate its stream.""" - def __init__(self, stream: 'TransformStream'): + def __init__(self, stream: TransformStream) -> None: self._stream = stream @property - def desired_size(self) -> Optional[float]: + def desired_size(self) -> float | None: """:obj:`float`, optional: The desired size of the readable side.""" return self._stream._readable._controller.desired_size - def enqueue(self, chunk: Any) -> None: + def enqueue(self, chunk: object) -> None: """Enqueues a chunk to the readable side.""" self._stream._readable._controller.enqueue(chunk) - def error(self, error: Optional[BaseException] = None) -> None: + def error(self, error: BaseException | None = None) -> None: """Errors both sides.""" error = _error_or_default(error) self._stream._readable._controller.error(error) @@ -837,8 +894,9 @@ def terminate(self) -> None: class TransformStream: - """A pair of streams where what's written is transformed and read - (https://developer.mozilla.org/en-US/docs/Web/API/TransformStream). + """A pair of streams where what's written is transformed and read. + + See https://developer.mozilla.org/en-US/docs/Web/API/TransformStream. Args: transformer (optional): An object with optional ``start(controller)``, ``transform(chunk, controller)`` and @@ -846,11 +904,11 @@ class TransformStream: unchanged without ``transform``. """ - def __init__(self, transformer: Any = None): + def __init__(self, transformer: object = None) -> None: self._transformer = transformer self._controller = TransformStreamDefaultController(self) # settled by a pull of the readable side, which relieves backpressure - self._pull_waiter: Optional[asyncio.Future] = None + self._pull_waiter: asyncio.Future | None = None self._readable = ReadableStream(_TransformSource(self), high_water_mark=0) self._writable = WritableStream(_TransformSink(self), high_water_mark=1) _call(transformer, 'start', self._controller) @@ -867,17 +925,17 @@ def writable(self) -> WritableStream: class _TransformSink: - def __init__(self, stream: TransformStream): + def __init__(self, stream: TransformStream) -> None: self._stream = stream - async def write(self, chunk, controller): + async def write(self, chunk: object, _controller: WritableStreamDefaultController) -> None: stream = self._stream readable = stream._readable # waits for the readable side to have room while ( readable._state == 'readable' and readable._controller.desired_size <= 0 - and not (readable._reader is not None and readable._reader._read_requests) + and not readable._has_read_requests() ): stream._pull_waiter = _pending() await stream._pull_waiter @@ -887,22 +945,22 @@ async def write(self, chunk, controller): else: await _await(transform(chunk, stream._controller)) - async def close(self): + async def close(self) -> None: stream = self._stream await _await(_call(stream._transformer, 'flush', stream._controller)) if stream._readable._state == 'readable' and not stream._readable._controller._close_requested: stream._readable._controller.close() - def abort(self, reason): + def abort(self, reason: object) -> None: self._stream._readable._controller.error(_reason_error(reason)) class _TransformSource: - def __init__(self, stream: TransformStream): + def __init__(self, stream: TransformStream) -> None: self._stream = stream - def pull(self, controller): + def pull(self, _controller: ReadableStreamDefaultController) -> None: _settle(self._stream._pull_waiter) - def cancel(self, reason): + def cancel(self, reason: object) -> None: self._stream._writable._controller.error(_reason_error(reason)) diff --git a/python-webrtc/python/webrtc/utils/__init__.py b/python-webrtc/python/webrtc/utils/__init__.py index 548a80a..d4ed397 100644 --- a/python-webrtc/python/webrtc/utils/__init__.py +++ b/python-webrtc/python/webrtc/utils/__init__.py @@ -4,3 +4,5 @@ # Use of this source code is governed by a BSD-style license # that can be found in the LICENSE.md file in the root of the project. # + +"""Internals shared by the modules of the package.""" diff --git a/python-webrtc/python/webrtc/utils/events.py b/python-webrtc/python/webrtc/utils/events.py index 3d08298..a00c216 100644 --- a/python-webrtc/python/webrtc/utils/events.py +++ b/python-webrtc/python/webrtc/utils/events.py @@ -7,45 +7,46 @@ """Events of WebRTC objects. libwebrtc threads only schedule them: handlers run on their event loop.""" +from __future__ import annotations + import asyncio import inspect -from typing import TYPE_CHECKING, Any, Callable, Dict, List, Optional, Tuple +from typing import Callable, NamedTuple, TypeVar, overload +import webrtc from webrtc.utils.task_queue import TaskQueue -if TYPE_CHECKING: - import webrtc +Handler = Callable[['webrtc.Event'], object] +_H = TypeVar('_H', bound=Handler) -Handler = Callable[['webrtc.Event'], Any] +#: The tasks of the coroutine handlers, referenced until they're done +_handler_tasks: set[asyncio.Future[object]] = set() -def _running_loop() -> Optional[asyncio.AbstractEventLoop]: +def _running_loop() -> asyncio.AbstractEventLoop | None: try: return asyncio.get_running_loop() except RuntimeError: return None -class _Registration: - __slots__ = ('handler', 'loop', 'once') - - def __init__(self, handler: Handler, loop: asyncio.AbstractEventLoop, once: bool): - self.handler = handler - self.loop = loop - self.once = once +class _Registration(NamedTuple): + handler: Handler + loop: asyncio.AbstractEventLoop + once: bool class _Listeners: """Handlers of one native object, called by it with the name and the native arguments of an event.""" - def __init__(self, target: 'EventTarget'): + def __init__(self, target: EventTarget) -> None: self.target = target - self.registrations: Dict[str, List[_Registration]] = {} + self.registrations: dict[str, list[_Registration]] = {} # the loop of the first handler, which delivers every event, even without handlers for it, # as events also update what the object shows (see EventTarget._on_event) - self.primary_loop: Optional[asyncio.AbstractEventLoop] = None + self.primary_loop: asyncio.AbstractEventLoop | None = None - def __call__(self, name: str, *args): + def __call__(self, name: str, *args: object) -> None: # a libwebrtc thread, with the GIL held: only schedule registrations = self.__dict__.get('registrations') if registrations is None: @@ -65,14 +66,14 @@ def __call__(self, name: str, *args): if not loop.is_closed(): TaskQueue.of(loop).post(self.deliver, loop, name, args) - def ensure_primary_loop(self) -> Optional[asyncio.AbstractEventLoop]: + def ensure_primary_loop(self) -> asyncio.AbstractEventLoop | None: """Makes the running loop the primary one if there's none yet. Returns the running loop, if any.""" loop = _running_loop() if loop is not None and (self.primary_loop is None or self.primary_loop.is_closed()): self.primary_loop = loop return loop - def deliver(self, loop: asyncio.AbstractEventLoop, name: str, args: Tuple): + def deliver(self, loop: asyncio.AbstractEventLoop, name: str, args: tuple[object, ...]) -> None: if loop is self.primary_loop: self.target._on_event(name, *args) registrations = [r for r in self.registrations.get(name, ()) if r.loop is loop] @@ -88,23 +89,28 @@ def deliver(self, loop: asyncio.AbstractEventLoop, name: str, args: Tuple): try: result = registration.handler(event) if inspect.isawaitable(result): - asyncio.ensure_future(result, loop=loop) + task = asyncio.ensure_future(result, loop=loop) + _handler_tasks.add(task) + task.add_done_callback(_handler_tasks.discard) except Exception as e: - loop.call_exception_handler( - {'message': f'Exception in {name!r} event handler', 'exception': e, 'event': event} - ) + loop.call_exception_handler({ + 'message': f'Exception in {name!r} event handler', + 'exception': e, + 'event': event, + }) - def add(self, name: str, handler: Handler, once: bool): + def add(self, name: str, handler: Handler, *, once: bool) -> None: loop = self.ensure_primary_loop() if loop is None: - raise RuntimeError('event handlers must be registered from a running asyncio event loop') + msg = 'event handlers must be registered from a running asyncio event loop' + raise RuntimeError(msg) registrations = self.registrations.setdefault(name, []) if any(r.handler == handler for r in registrations): return # like addEventListener, a handler is registered once registrations.append(_Registration(handler, loop, once)) - def remove(self, name: str, handler: Optional[Handler]) -> None: + def remove(self, name: str, handler: Handler | None) -> None: if handler is None: self.registrations.pop(name, None) return @@ -124,13 +130,14 @@ class EventTarget: async def on_candidate(event): await signaling.send(event.candidate) + pc.on('track', lambda event: print(event.track)) """ #: Names of the events the object emits - _events: Tuple[str, ...] = () + _events: tuple[str, ...] = () - def _listeners(self, create: bool) -> Optional[_Listeners]: + def _listeners(self, *, create: bool) -> _Listeners | None: native = self._native_obj listeners = native._listeners if listeners is None and create: @@ -141,41 +148,53 @@ def _listeners(self, create: bool) -> Optional[_Listeners]: return listeners def _attach(self) -> None: - """Delivers the events of the object to the running loop from now on, even without handlers, as they also - update what the object shows (like its state, see :meth:`_on_event`). Does nothing outside of a loop.""" + """Delivers the events of the object to the running loop from now on, even without handlers. + + Events also update what the object shows (like its state, see :meth:`_on_event`). Does nothing outside of + a loop. + """ if _running_loop() is not None: self._listeners(create=True).ensure_primary_loop() - def _check_event(self, name: str): + def _check_event(self, name: str) -> None: if name not in self._events: - raise ValueError(f'{type(self).__name__} has no event {name!r}, its events are: {", ".join(self._events)}') + msg = f'{type(self).__name__} has no event {name!r}, its events are: {", ".join(self._events)}' + raise ValueError(msg) - def _add(self, name: str, handler: Optional[Handler], once: bool): + def _add(self, name: str, handler: _H | None, *, once: bool) -> _H | Callable[[_H], _H]: self._check_event(name) if handler is None: - return lambda func: self._add(name, func, once) - self._listeners(create=True).add(name, handler, once) + return lambda func: self._add_handler(name, func, once=once) + return self._add_handler(name, handler, once=once) + + def _add_handler(self, name: str, handler: _H, *, once: bool) -> _H: + self._listeners(create=True).add(name, handler, once=once) return handler - def _dispatch(self, name: str, *args) -> None: + def _dispatch(self, name: str, *args: object) -> None: """Delivers an event to the handlers on the running loop right away, from an event being delivered.""" listeners = self._listeners(create=False) if listeners is not None: listeners.deliver(asyncio.get_running_loop(), name, args) - def _on_event(self, name: str, *args) -> None: + def _on_event(self, name: str, *args: object) -> None: """Called for every event on the loop of the first handler, before the handlers of the event. Names starting with ``_`` (like ``'_sent'``) are internal events of the native object: they only reach this - method, never handlers.""" + method, never handlers. + """ - def _create_event(self, name: str, *args): + def _create_event(self, name: str, *_args: object) -> webrtc.Event | None: """Creates the event object from the native arguments of an event, or returns :obj:`None` to drop it.""" - from webrtc import Event + return webrtc.Event(name, self) + + @overload + def on(self, name: str, handler: None = None) -> Callable[[_H], _H]: ... - return Event(name, self) + @overload + def on(self, name: str, handler: _H) -> _H: ... - def on(self, name: str, handler: Optional[Handler] = None): + def on(self, name: str, handler: _H | None = None) -> _H | Callable[[_H], _H]: """Registers a handler of an event. Can be used as a decorator. Args: @@ -187,12 +206,18 @@ def on(self, name: str, handler: Optional[Handler] = None): :obj:`callable`: The handler, or a decorator registering it. Raises: - :obj:`ValueError`: If the object has no such event. - :obj:`RuntimeError`: If called outside of a running event loop. + ValueError: If the object has no such event. + RuntimeError: If called outside of a running event loop. """ return self._add(name, handler, once=False) - def once(self, name: str, handler: Optional[Handler] = None): + @overload + def once(self, name: str, handler: None = None) -> Callable[[_H], _H]: ... + + @overload + def once(self, name: str, handler: _H) -> _H: ... + + def once(self, name: str, handler: _H | None = None) -> _H | Callable[[_H], _H]: """Registers a handler that is removed after it's called for the first time. Can be used as a decorator. Args: @@ -203,12 +228,12 @@ def once(self, name: str, handler: Optional[Handler] = None): :obj:`callable`: The handler, or a decorator registering it. Raises: - :obj:`ValueError`: If the object has no such event. - :obj:`RuntimeError`: If called outside of a running event loop. + ValueError: If the object has no such event. + RuntimeError: If called outside of a running event loop. """ return self._add(name, handler, once=True) - def off(self, name: str, handler: Optional[Handler] = None) -> None: + def off(self, name: str, handler: Handler | None = None) -> None: """Removes a handler of an event, or every handler of the event. Args: @@ -216,7 +241,7 @@ def off(self, name: str, handler: Optional[Handler] = None) -> None: handler (:obj:`callable`, optional): The handler to remove. If omitted, all handlers of the event are. Raises: - :obj:`ValueError`: If the object has no such event. + ValueError: If the object has no such event. """ self._check_event(name) listeners = self._listeners(create=False) diff --git a/python-webrtc/python/webrtc/utils/names.py b/python-webrtc/python/webrtc/utils/names.py index d0c2a32..2dfbe95 100644 --- a/python-webrtc/python/webrtc/utils/names.py +++ b/python-webrtc/python/webrtc/utils/names.py @@ -7,10 +7,16 @@ """The camelCase names of the WebRTC specification next to the snake_case ones of the library.""" +from __future__ import annotations + import re -from typing import Any +from typing import TYPE_CHECKING, Any, Generic, TypeVar, overload + +if TYPE_CHECKING: + from collections.abc import Iterable, Mapping _CAMEL = re.compile(r'_([a-z0-9])') +_T = TypeVar('_T') def camel_case(name: str) -> str: @@ -23,19 +29,46 @@ def snake_case(name: str) -> str: return ''.join(f'_{c.lower()}' if c.isupper() else c for c in name) -class alias: +def members(value: Mapping[str, Any], names: Iterable[str]) -> dict[str, Any]: + """The members of a dictionary, with snake_case or camelCase names: unknown ones are ignored, as in WebIDL.""" + wanted = set(names) + return {snake_case(name): member for name, member in value.items() if snake_case(name) in wanted} + + +class Alias(Generic[_T]): """A camelCase alias of an instance attribute, like a dataclass field: ``maxBitrate = alias('max_bitrate')``. It isn't a field itself, so ``__init__``, ``repr()``, ``==`` and :obj:`dataclasses.asdict` don't see it. Setting it sets the attribute, which a frozen dataclass doesn't allow. + + Args: + name (:obj:`str`): The name of the attribute. """ - def __init__(self, name: str): + def __init__(self, name: str) -> None: self.name = name self.__doc__ = f'Alias for :attr:`{name}`' - def __get__(self, obj: Any, owner: Any = None) -> Any: + @overload + def __get__(self, obj: None, owner: type | None = None) -> Alias[_T]: ... + + @overload + def __get__(self, obj: object, owner: type | None = None) -> _T: ... + + def __get__(self, obj: object, owner: type | None = None) -> Alias[_T] | _T: return self if obj is None else getattr(obj, self.name) - def __set__(self, obj: Any, value: Any) -> None: + def __set__(self, obj: object, value: _T) -> None: setattr(obj, self.name, value) + + +def alias(name: str) -> Alias[Any]: + """The :obj:`Alias` of an attribute, like :func:`dataclasses.field` for a field. + + Args: + name (:obj:`str`): The name of the attribute. + + Returns: + :obj:`Alias`: The alias. + """ + return Alias(name) diff --git a/python-webrtc/python/webrtc/utils/native_calls.py b/python-webrtc/python/webrtc/utils/native_calls.py index ea66cfd..c59b482 100644 --- a/python-webrtc/python/webrtc/utils/native_calls.py +++ b/python-webrtc/python/webrtc/utils/native_calls.py @@ -5,13 +5,30 @@ # that can be found in the LICENSE.md file in the root of the project. # +"""Awaiting the native methods that report their result with callbacks.""" + +from __future__ import annotations + import asyncio -from typing import Any, Callable +from typing import TYPE_CHECKING, Callable, TypeVar from webrtc.utils.task_queue import TaskQueue +if TYPE_CHECKING: + from typing_extensions import Concatenate, ParamSpec + + import wrtc + + _P = ParamSpec('_P') + +_T = TypeVar('_T') + -async def call_native(method: Callable, *args) -> Any: +async def call_native( + method: Callable[Concatenate[Callable[[_T], None], Callable[[wrtc.RTCCallbackException], None], _P], None], + *args: _P.args, + **kwargs: _P.kwargs, +) -> _T: """Calls a native method taking success and failure callbacks, called from a libwebrtc thread, and awaits them. The result goes through the task queue of the loop, so the code awaiting it runs after the handlers of the events @@ -20,6 +37,7 @@ async def call_native(method: Callable, *args) -> Any: Args: method (:obj:`callable`): The native method, called as ``method(on_success, on_failure, *args)``. *args: Its arguments. + **kwargs: Its keyword arguments. Returns: The result passed to ``on_success``, if any. @@ -30,7 +48,7 @@ async def call_native(method: Callable, *args) -> Any: loop = asyncio.get_running_loop() future = loop.create_future() - def settle(result: Any, error: Any) -> None: + def settle(result: _T | None, error: wrtc.RTCCallbackException | None) -> None: # the caller may have been canceled meanwhile if future.done(): return @@ -40,11 +58,11 @@ def settle(result: Any, error: Any) -> None: future.set_result(result) # libwebrtc threads, with the GIL held: only schedule - def on_success(result: Any = None) -> None: + def on_success(result: _T | None = None) -> None: TaskQueue.of(loop).post(settle, result, None, resumes=True, after_ready=True) - def on_failure(error: Any) -> None: + def on_failure(error: wrtc.RTCCallbackException) -> None: TaskQueue.of(loop).post(settle, None, error, resumes=True, after_ready=True) - method(on_success, on_failure, *args) + method(on_success, on_failure, *args, **kwargs) return await future diff --git a/python-webrtc/python/webrtc/utils/operations.py b/python-webrtc/python/webrtc/utils/operations.py index ad07cf7..366fad2 100644 --- a/python-webrtc/python/webrtc/utils/operations.py +++ b/python-webrtc/python/webrtc/utils/operations.py @@ -15,23 +15,29 @@ ... # the native call """ +from __future__ import annotations + import asyncio import contextlib -from typing import AsyncIterator, Callable, Optional +from typing import TYPE_CHECKING, Callable + +if TYPE_CHECKING: + from collections.abc import AsyncIterator class OperationsChain: - """Runs the operations of a connection one after another, as the specification requires. With none running, - an operation starts right away, so its checks fail in the task that called it. + """Runs the operations of a connection one after another, as the specification requires. + + With none running, an operation starts right away, so its checks fail in the task that called it. Args: on_empty (:obj:`callable`): Called when the last operation ends. """ - def __init__(self, on_empty: Callable[[], None]): + def __init__(self, on_empty: Callable[[], None]) -> None: self._on_empty = on_empty #: The operation that ends last - self._last: Optional[asyncio.Future] = None + self._last: asyncio.Future[None] | None = None @property def busy(self) -> bool: @@ -55,6 +61,8 @@ async def operation(self) -> AsyncIterator[None]: async def later() -> None: - """Operations of a connection take effect in a later task than the code that started them: a track added right - after set_remote_description() is added before the description is applied.""" + """Waits for a later task, where the operations of a connection take effect. + + A track added right after set_remote_description() is added before the description is applied. + """ await asyncio.sleep(0) diff --git a/python-webrtc/python/webrtc/utils/task_queue.py b/python-webrtc/python/webrtc/utils/task_queue.py index 916a3fb..1d7286f 100644 --- a/python-webrtc/python/webrtc/utils/task_queue.py +++ b/python-webrtc/python/webrtc/utils/task_queue.py @@ -5,23 +5,32 @@ # that can be found in the LICENSE.md file in the root of the project. # +"""The queue of the callbacks libwebrtc threads post to an event loop.""" + +from __future__ import annotations + import asyncio import collections import threading import weakref -from typing import Callable, NamedTuple, Tuple +from typing import Callable, NamedTuple class _Item(NamedTuple): - callback: Callable - args: Tuple + callback: Callable[..., object] + args: tuple[object, ...] resumes: bool after_ready: bool class TaskQueue: - """Runs callbacks posted from any thread on an event loop in order, each with what it schedules with - ``call_soon`` before the next one (like browser tasks and microtasks).""" + """Runs callbacks posted from any thread on an event loop in order, like browser tasks. + + Each one runs with what it schedules with ``call_soon`` before the next one, like microtasks. + + Args: + loop (:obj:`asyncio.AbstractEventLoop`): The loop. + """ #: How many times a callback waits for other ready callbacks of the loop MAX_DEFERRALS = 100 @@ -29,24 +38,24 @@ class TaskQueue: MAX_BATCH = 100 #: The queues of the loops without a ``__dict__`` (like uvloop's), the others keep their own - _queues: 'weakref.WeakKeyDictionary[asyncio.AbstractEventLoop, TaskQueue]' = weakref.WeakKeyDictionary() + _queues: weakref.WeakKeyDictionary[asyncio.AbstractEventLoop, TaskQueue] = weakref.WeakKeyDictionary() _ATTRIBUTE = '_webrtc_task_queue' - def __init__(self, loop: asyncio.AbstractEventLoop): + def __init__(self, loop: asyncio.AbstractEventLoop) -> None: # weakly: a queue in _queues referencing its loop would keep it forever self._loop_ref = weakref.ref(loop) - self._items = collections.deque() + self._items: collections.deque[_Item] = collections.deque() self._lock = threading.Lock() self._scheduled = False # whether the last callback resumed code (like a coroutine awaiting a result) that runs before the next one self._resumed = False @property - def _loop(self) -> asyncio.AbstractEventLoop: + def _loop(self) -> asyncio.AbstractEventLoop | None: return self._loop_ref() @classmethod - def of(cls, loop: asyncio.AbstractEventLoop) -> 'TaskQueue': + def of(cls, loop: asyncio.AbstractEventLoop) -> TaskQueue: """Returns the queue of a loop. Args: @@ -73,9 +82,17 @@ def of(cls, loop: asyncio.AbstractEventLoop) -> 'TaskQueue': return queue @classmethod - def post_to_running(cls, callback: Callable, *args, **kwargs) -> bool: + def post_to_running( + cls, callback: Callable[..., object], *args: object, resumes: bool = False, after_ready: bool = False + ) -> bool: """Posts a callback to the queue of the running loop (see :meth:`post`). + Args: + callback (:obj:`callable`): The callback. + *args: Its arguments. + resumes (:obj:`bool`, optional): See :meth:`post`. + after_ready (:obj:`bool`, optional): See :meth:`post`. + Returns: :obj:`bool`: Whether it was posted: :obj:`False` outside of a running loop. """ @@ -83,10 +100,12 @@ def post_to_running(cls, callback: Callable, *args, **kwargs) -> bool: loop = asyncio.get_running_loop() except RuntimeError: return False - cls.of(loop).post(callback, *args, **kwargs) + cls.of(loop).post(callback, *args, resumes=resumes, after_ready=after_ready) return True - def post(self, callback: Callable, *args, resumes: bool = False, after_ready: bool = False) -> None: + def post( + self, callback: Callable[..., object], *args: object, resumes: bool = False, after_ready: bool = False + ) -> None: """Schedules a callback, from any thread. Callbacks posted to a closed loop are dropped. Args: @@ -106,31 +125,33 @@ def post(self, callback: Callable, *args, resumes: bool = False, after_ready: bo return self._scheduled = True loop = self._loop + # a gone or closed loop won't run what's queued, nor release it + if loop is None: + self._items.clear() + return try: - if loop is None: - raise RuntimeError('the loop is gone') loop.call_soon_threadsafe(self._run) - except RuntimeError: # the loop is closed: nothing will run what's queued, nor release it + except RuntimeError: self._items.clear() def _others_ready(self) -> bool: - """Whether the loop has callbacks ready that are like microtasks: the steps of coroutines and callbacks - scheduled with ``call_soon``, rather than timers, the loop's own ones or the ones of this queue""" + """Whether the loop has microtask-like callbacks ready: coroutine steps and ``call_soon`` callbacks.""" + # rather than timers, the loop's own callbacks or the ones of this queue # CPython internals (loop._ready, handle._callback): other loops (like uvloop) never report any ready = getattr(self._loop, '_ready', None) or () for handle in ready: callback = getattr(handle, '_callback', None) - if isinstance(handle, asyncio.TimerHandle) or getattr(callback, '__self__', None) in (self, self._loop): + if isinstance(handle, asyncio.TimerHandle) or getattr(callback, '__self__', None) in {self, self._loop}: continue return True return False - def _defers(self, deferred: int, resumed: bool) -> bool: - """Whether a callback waits for the code the previous one resumed, or for the microtasks the loop has - ready. Deferring is bounded, so a busy loop can't starve the queue.""" + def _defers(self, deferred: int, *, resumed: bool) -> bool: + """Whether a callback waits for the code the previous one resumed, or for the ready microtasks.""" + # bounded, so a busy loop can't starve the queue return (resumed or self._others_ready()) and deferred < self.MAX_DEFERRALS - def _settle(self, deferred: int = 0): + def _settle(self, deferred: int = 0) -> None: # the code a callback resumed runs until the loop has nothing else ready (a coroutine continues over # several iterations): from then on, callbacks don't wait for it anymore if self._defers(deferred, resumed=False): @@ -138,36 +159,14 @@ def _settle(self, deferred: int = 0): else: self._resumed = False - def _run(self, deferred: int = 0): - if self._defers(deferred, self._resumed): + def _run(self, deferred: int = 0) -> None: + if self._defers(deferred, resumed=self._resumed): self._loop.call_soon(self._run, deferred + 1) return - # without the internals _others_ready() reads, one callback runs per iteration of the loop - can_inspect_ready = hasattr(self._loop, '_ready') more = True try: - for _ in range(self.MAX_BATCH): - with self._lock: - item = self._items.popleft() - resumes = item.resumes - if resumes: - self._resumed = True - self._loop.call_soon(self._settle) - item.callback(*item.args) - # released outside of the lock: the destructor of a native object may wait for a libwebrtc thread, - # which may be posting here - item = None - with self._lock: - more = bool(self._items) - self._scheduled = more - if not more: - return - after_ready = self._items[0].after_ready - # the next callback runs right away, unless this one resumed code, the next one waits for the ready - # callbacks, or this one scheduled some (they run first) - if resumes or after_ready or not can_inspect_ready or self._others_ready(): - break + more = self._run_batch() except BaseException: with self._lock: more = bool(self._items) @@ -176,3 +175,37 @@ def _run(self, deferred: int = 0): finally: if more: self._loop.call_soon(self._run) + + def _run_batch(self) -> bool: + """Runs callbacks until one has to wait for the loop. Returns whether any are left.""" + for _ in range(self.MAX_BATCH): + resumes = self._run_next() + with self._lock: + more = bool(self._items) + self._scheduled = more + if not more: + return False + after_ready = self._items[0].after_ready + if self._yields(resumes=resumes, after_ready=after_ready): + return True + return True + + def _run_next(self) -> bool: + """Runs the next callback. Returns whether it resumes code.""" + with self._lock: + item = self._items.popleft() + if item.resumes: + self._resumed = True + self._loop.call_soon(self._settle) + item.callback(*item.args) + # the item is released on return, outside of the lock: the destructor of a native object may wait for + # a libwebrtc thread, which may be posting here + return item.resumes + + def _yields(self, *, resumes: bool, after_ready: bool) -> bool: + """Whether the next callback waits for the loop rather than running right away.""" + # it waits for the code this one resumed, for the ready callbacks, or for the ones this one scheduled (they + # run first); without the internals _others_ready() reads, one callback runs per iteration of the loop + if resumes or after_ready: + return True + return not hasattr(self._loop, '_ready') or self._others_ready() diff --git a/tests/chaos.py b/tests/chaos.py index 7f8be27..0575245 100644 --- a/tests/chaos.py +++ b/tests/chaos.py @@ -5,92 +5,117 @@ # that can be found in the LICENSE.md file in the root of the project. # -"""Random sequences of API calls, to find crashes, deadlocks and leaks. Every step is printed before it runs, so the -output of a crash names the sequence; the same seed replays it. +"""Random sequences of API calls, to find crashes, deadlocks and leaks. + +Every step is logged before it runs, so the output of a crash names the sequence; the same seed replays it. python -m tests.chaos --seed 7 --steps 500 """ +from __future__ import annotations + import argparse import asyncio import gc +import logging import random import sys import threading import time +from typing import TYPE_CHECKING, ClassVar, TypeVar import webrtc import wrtc from tests.helpers import connect +if TYPE_CHECKING: + from collections.abc import Callable + +log = logging.getLogger('chaos') + #: How long one step may take: longer is a deadlock STEP_TIMEOUT = 20 +#: What misusing the API raises +MISUSE = (webrtc.PythonWebRTCExceptionBase, ValueError, TypeError, RuntimeError) + +T = TypeVar('T') + -class Chaos: - def __init__(self, seed: int): +class State: + """The objects the steps work on, and what the steps share.""" + + def __init__(self, seed: int) -> None: self.random = random.Random(seed) - self.connections = [] - self.channels = [] - self.tracks = [] - self.processors = [] - self.generators = [] - self.frames = [] - self.tasks = [] - - def pick(self, pool): + self.connections: list[webrtc.RTCPeerConnection] = [] + self.channels: list[webrtc.RTCDataChannel] = [] + self.tracks: list[webrtc.MediaStreamTrack] = [] + self.processors: list[tuple[webrtc.MediaStreamTrackProcessor, webrtc.ReadableStreamDefaultReader]] = [] + self.generators: list[tuple[webrtc.WritableStreamDefaultWriter, str]] = [] + self.frames: list[webrtc.VideoFrame] = [] + self.tasks: list[threading.Thread | asyncio.Future[None]] = [] + + def pick(self, pool: list[T]) -> T | None: return self.random.choice(pool) if pool else None - def drop(self, pool): + def drop(self, pool: list[object]) -> None: if pool: pool.pop(self.random.randrange(len(pool))) - def handler(self): - """A handler doing something to a random object: closing, raising, referencing (a cycle), collecting""" - target = self.pick(self.connections + self.channels + self.tracks) + def handler(self) -> Callable[[webrtc.Event], str | None]: + """A handler doing something to a random object: closing, raising, referencing (a cycle), collecting.""" + target = self.pick([*self.connections, *self.channels, *self.tracks]) action = self.random.randrange(5) - def handle(event): + def handle(_event: webrtc.Event) -> str | None: if action == 0 and target is not None: target.close() if hasattr(target, 'close') else target.stop() elif action == 1: - raise RuntimeError('a handler raises') + msg = 'a handler raises' + raise RuntimeError(msg) elif action == 2: gc.collect() elif action == 3: return repr(target) + return None return handle - # the steps, each with the objects it works on - async def new_connection(self): +def steps(cls: type) -> list[str]: + return [name for name in vars(cls) if not name.startswith('_')] + + +class ConnectionSteps(State): + """Steps of connections and channels.""" + + async def new_connection(self) -> None: self.connections.append(webrtc.RTCPeerConnection()) - async def close_connection(self): + async def close_connection(self) -> None: pc = self.pick(self.connections) if pc: pc.close() - async def drop_connection(self): + async def drop_connection(self) -> None: self.drop(self.connections) - async def connect_two(self): + async def connect_two(self) -> None: if len(self.connections) >= 2: a, b = self.random.sample(self.connections, 2) await connect(a, b, timeout=5) - async def add_track(self): + async def add_track(self) -> None: pc, track = self.pick(self.connections), self.pick(self.tracks) if pc and track: pc.add_track(track) - async def remove_track(self): + async def remove_track(self) -> None: pc = self.pick(self.connections) if pc and pc.get_senders(): pc.remove_track(self.random.choice(pc.get_senders())) - async def add_transceiver(self): + async def add_transceiver(self) -> None: pc = self.pick(self.connections) if pc: transceiver = pc.add_transceiver(self.random.choice(['audio', 'video'])) @@ -99,72 +124,91 @@ async def add_transceiver(self): elif self.random.random() < 0.3: transceiver.direction = self.random.choice(list(webrtc.TransceiverDirection)[:4]) - async def negotiate(self): + async def negotiate(self) -> None: pc = self.pick(self.connections) if pc: await pc.set_local_description() - async def create_channel(self): + async def create_channel(self) -> None: pc = self.pick(self.connections) if pc: channel = pc.create_data_channel(f'chaos{self.random.randrange(1000)}') channel.on(self.random.choice(['open', 'message', 'close']), self.handler()) self.channels.append(channel) - async def send(self): + async def send(self) -> None: channel = self.pick(self.channels) if channel: channel.send(self.random.choice(['text', b'\x00' * self.random.randrange(70000), bytearray(10)])) - async def close_channel(self): + async def close_channel(self) -> None: channel = self.pick(self.channels) if channel: channel.close() - async def stats(self): + async def stats(self) -> None: pc = self.pick(self.connections) if pc: await pc.get_stats() - async def restart_ice(self): + async def restart_ice(self) -> None: pc = self.pick(self.connections) if pc: pc.restart_ice() - async def handle_connection_event(self): + async def handle_connection_event(self) -> None: pc = self.pick(self.connections) if pc: pc.on(self.random.choice(['connectionstatechange', 'icecandidate', 'track', 'datachannel']), self.handler()) - async def get_user_media(self): + async def replace_track(self) -> None: + pc = self.pick(self.connections) + if pc and pc.get_senders(): + await self.random.choice(pc.get_senders()).replace_track(self.pick([*self.tracks, None])) + + async def set_parameters(self) -> None: + pc = self.pick(self.connections) + if pc and pc.get_senders(): + sender = self.random.choice(pc.get_senders()) + parameters = sender.get_parameters() + for encoding in parameters.encodings: + encoding.active = self.random.random() < 0.8 + encoding.max_bitrate = self.random.choice([None, 30000, 2**31]) + await sender.set_parameters(parameters) + + +class MediaSteps(State): + """Steps of tracks, processors, generators and frames.""" + + async def get_user_media(self) -> None: self.tracks.extend(webrtc.get_user_media(audio=True, video=True).get_tracks()) - async def stop_track(self): + async def stop_track(self) -> None: track = self.pick(self.tracks) if track: track.stop() - async def clone_track(self): + async def clone_track(self) -> None: track = self.pick(self.tracks) if track: self.tracks.append(track.clone()) - async def toggle_track(self): + async def toggle_track(self) -> None: track = self.pick(self.tracks) if track: track.enabled = not track.enabled - async def drop_track(self): + async def drop_track(self) -> None: self.drop(self.tracks) - async def new_processor(self): + async def new_processor(self) -> None: track = self.pick(self.tracks) if track: processor = webrtc.MediaStreamTrackProcessor(track, max_buffer_size=self.random.randrange(4)) track.on('ended', self.handler()) self.processors.append((processor, processor.readable.get_reader())) - async def read(self): + async def read(self) -> None: if self.processors: _, reader = self.pick(self.processors) try: @@ -174,15 +218,15 @@ async def read(self): if not result.done: result.value.close() - async def cancel_processor(self): + async def cancel_processor(self) -> None: if self.processors: _, reader = self.pick(self.processors) await reader.cancel() - async def drop_processor(self): + async def drop_processor(self) -> None: self.drop(self.processors) - async def new_generator(self): + async def new_generator(self) -> None: if self.random.random() < 0.5: generator = webrtc.VideoTrackGenerator() self.generators.append((generator.writable.get_writer(), 'video')) @@ -192,7 +236,7 @@ async def new_generator(self): self.generators.append((generator.writable.get_writer(), 'audio')) self.tracks.append(generator) - async def write(self): + async def write(self) -> None: if not self.generators: return writer, kind = self.pick(self.generators) @@ -214,15 +258,15 @@ async def write(self): ) await writer.write(chunk) - async def close_generator(self): + async def close_generator(self) -> None: if self.generators: writer, _ = self.pick(self.generators) await writer.close() - async def drop_generator(self): + async def drop_generator(self) -> None: self.drop(self.generators) - async def frame(self): + async def frame(self) -> None: fmt = self.random.choice(list(webrtc.VideoPixelFormat)) width, height = self.random.randrange(1, 40), self.random.randrange(1, 40) frame = webrtc.VideoFrame( @@ -232,27 +276,12 @@ async def frame(self): await frame.copy_to(bytearray(frame.allocation_size(options)), options) self.frames.append(frame) - async def use_frame(self): + async def use_frame(self) -> None: frame = self.pick(self.frames) if frame: self.random.choice([frame.close, lambda: self.frames.append(frame.clone())])() - async def replace_track(self): - pc = self.pick(self.connections) - if pc and pc.get_senders(): - await self.random.choice(pc.get_senders()).replace_track(self.pick(self.tracks + [None])) - - async def set_parameters(self): - pc = self.pick(self.connections) - if pc and pc.get_senders(): - sender = self.random.choice(pc.get_senders()) - parameters = sender.get_parameters() - for encoding in parameters.encodings: - encoding.active = self.random.random() < 0.8 - encoding.max_bitrate = self.random.choice([None, 30000, 2**31]) - await sender.set_parameters(parameters) - - async def pipe(self): + async def pipe(self) -> None: track = self.pick([track for track in self.tracks if track.kind == 'video']) if track: processor = webrtc.MediaStreamTrackProcessor(track) @@ -260,24 +289,28 @@ async def pipe(self): self.tracks.append(generator.track) self.tasks.append(processor.readable.pipe_through(webrtc.TransformStream()).pipe_to(generator.writable)) - async def constraints(self): + async def constraints(self) -> None: track = self.pick(self.tracks) if track: track.get_settings() track.apply_constraints(self.random.choice([{'width': 320}, {'frame_rate': 5}, {'width': {'exact': 7}}])) - async def stream(self): + async def stream(self) -> None: tracks = self.random.sample(self.tracks, min(len(self.tracks), 2)) stream = webrtc.MediaStream(tracks) if tracks and self.random.random() < 0.5: stream.remove_track(tracks[0]) stream.get_tracks() - async def reader_thread(self): - """A thread reading the objects while the loop goes on""" + +class LoopSteps(State): + """Steps of threads, the garbage collector and the loop.""" + + async def reader_thread(self) -> None: + """A thread reading the objects while the loop goes on.""" connections, tracks = list(self.connections), list(self.tracks) - def read(): + def read() -> None: for _ in range(50): for pc in connections: _ = pc.connection_state, pc.get_transceivers(), pc.sctp @@ -288,30 +321,24 @@ def read(): thread.start() self.tasks.append(thread) - async def collect(self): + @staticmethod + async def collect() -> None: gc.collect() - async def pause(self): + async def pause(self) -> None: await asyncio.sleep(self.random.random() * 0.05) - STEPS = [name for name in dir() if not name.startswith('_') and name not in ('pick', 'drop', 'handler')] - async def run(self, steps: int): - loop = asyncio.get_running_loop() +class Chaos(ConnectionSteps, MediaSteps, LoopSteps): + """Every step, run in a random sequence.""" + + STEPS: ClassVar[list[str]] = sorted([*steps(ConnectionSteps), *steps(MediaSteps), *steps(LoopSteps)]) + + async def run(self, count: int) -> None: # handlers raise on purpose - loop.set_exception_handler(lambda loop, context: None) - for index in range(steps): - name = self.random.choice(self.STEPS) - print(f'{index} {name}', flush=True) - started = time.monotonic() - try: - await asyncio.wait_for(getattr(self, name)(), STEP_TIMEOUT) - except asyncio.TimeoutError: - if time.monotonic() - started >= STEP_TIMEOUT: - print(f'step {index} {name} is stuck', flush=True) - sys.exit(3) - except Exception as e: # noqa: BLE001 (misuse is expected, crashes and deadlocks aren't) - print(f' {type(e).__name__}: {str(e)[:80]}', flush=True) + asyncio.get_running_loop().set_exception_handler(lambda _loop, _context: None) + for index in range(count): + await self.step(index, self.random.choice(self.STEPS)) for pc in self.connections: pc.close() for track in self.tracks: @@ -320,16 +347,31 @@ async def run(self, steps: int): if isinstance(task, threading.Thread): task.join(STEP_TIMEOUT) if task.is_alive(): - print('a reader thread is stuck', flush=True) + log.info('a reader thread is stuck') sys.exit(3) - -def main(): + async def step(self, index: int, name: str) -> None: + log.info('%d %s', index, name) + started = time.monotonic() + try: + await asyncio.wait_for(getattr(self, name)(), STEP_TIMEOUT) + # a step timing out, or connect_two's wait (a builtin TimeoutError before 3.11) + except (asyncio.TimeoutError, TimeoutError): + if time.monotonic() - started >= STEP_TIMEOUT: + log.info('step %d %s is stuck', index, name) + sys.exit(3) + # misuse is expected, crashes and deadlocks aren't + except MISUSE as e: + log.info(' %s: %s', type(e).__name__, str(e)[:80]) + + +def main() -> None: parser = argparse.ArgumentParser() parser.add_argument('--seed', type=int, default=0) parser.add_argument('--steps', type=int, default=300) args = parser.parse_args() - print(f'seed {args.seed}, {args.steps} steps', flush=True) + logging.basicConfig(stream=sys.stdout, format='%(message)s', level=logging.INFO) + log.info('seed %d, %d steps', args.seed, args.steps) asyncio.run(Chaos(args.seed).run(args.steps)) # the last references may be released on helper threads deadline = time.monotonic() + 1 @@ -337,7 +379,7 @@ def main(): gc.collect() time.sleep(0.05) alive = {name: count for name, count in wrtc._alive().items() if count} - print(f'done, {wrtc._alive_factories()} factories alive, native objects alive: {alive}', flush=True) + log.info('done, %d factories alive, native objects alive: %s', wrtc._alive_factories(), alive) if __name__ == '__main__': diff --git a/tests/conftest.py b/tests/conftest.py index 6298dc5..b70082a 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -5,21 +5,29 @@ # that can be found in the LICENSE.md file in the root of the project. # +from __future__ import annotations + import gc import threading +from typing import TYPE_CHECKING import pytest import webrtc from webrtc.utils import events +if TYPE_CHECKING: + from collections.abc import Iterator + + from tests.helpers import CreatePC + -def pytest_addoption(parser): +def pytest_addoption(parser: pytest.Parser) -> None: parser.addoption('--gc-on-emit', action='store_true', help='collect garbage on events of libwebrtc threads') parser.addoption('--stress', action='store_true', help='run the long stress tests too') -def pytest_collection_modifyitems(config, items): +def pytest_collection_modifyitems(config: pytest.Config, items: list[pytest.Item]) -> None: if config.getoption('--stress'): return skip = pytest.mark.skip(reason='a long stress test, run with --stress') @@ -28,14 +36,14 @@ def pytest_collection_modifyitems(config, items): item.add_marker(skip) -def pytest_configure(config): +def pytest_configure(config: pytest.Config) -> None: config.addinivalue_line('markers', 'stress: a long stress test, run with --stress') if not config.getoption('--gc-on-emit'): return # the collector runs on libwebrtc threads, as it may whenever they emit: whatever it releases must not block them emit = events._Listeners.__call__ - def collecting_emit(self, name, *args): + def collecting_emit(self: events._Listeners, name: str, *args: object) -> None: if threading.current_thread() is not threading.main_thread(): gc.collect() emit(self, name, *args) @@ -44,25 +52,20 @@ def collecting_emit(self, name, *args): @pytest.fixture -def rtc_peer_connection(request): +def rtc_peer_connection() -> Iterator[webrtc.RTCPeerConnection]: pc = webrtc.RTCPeerConnection() - - def close_pc(): - pc.close() - - request.addfinalizer(close_pc) - - return pc + yield pc + pc.close() pc = caller = callee = callee2 = rtc_peer_connection @pytest.fixture -def create_pc(request): - """Creates connections with a configuration, closed after the test""" +def create_pc(request: pytest.FixtureRequest) -> CreatePC: + """Creates connections with a configuration, closed after the test.""" - def create(configuration=None): + def create(configuration: webrtc.RTCConfiguration | None = None) -> webrtc.RTCPeerConnection: pc = webrtc.RTCPeerConnection(configuration) request.addfinalizer(pc.close) return pc @@ -70,10 +73,10 @@ def create(configuration=None): return create -def get_stream(constraints, request): +def get_stream(constraints: dict[str, bool], request: pytest.FixtureRequest) -> webrtc.MediaStream: stream = webrtc.get_user_media(**constraints) - def stop_tracks(): + def stop_tracks() -> None: for track in stream.get_tracks(): track.stop() @@ -83,7 +86,7 @@ def stop_tracks(): @pytest.fixture -def audio_stream(request): +def audio_stream(request: pytest.FixtureRequest) -> webrtc.MediaStream: return get_stream({'audio': True}, request) @@ -91,5 +94,5 @@ def audio_stream(request): @pytest.fixture -def video_stream(request): +def video_stream(request: pytest.FixtureRequest) -> webrtc.MediaStream: return get_stream({'audio': False, 'video': True}, request) diff --git a/tests/fuzz/__init__.py b/tests/fuzz/__init__.py new file mode 100644 index 0000000..f13f15d --- /dev/null +++ b/tests/fuzz/__init__.py @@ -0,0 +1,6 @@ +# +# Copyright 2026 Ilya (Marshal) . All rights reserved. +# +# Use of this source code is governed by a BSD-style license +# that can be found in the LICENSE.md file in the root of the project. +# diff --git a/tests/fuzz/fuzz_audio_data.py b/tests/fuzz/fuzz_audio_data.py index 1b83c12..b31d182 100644 --- a/tests/fuzz/fuzz_audio_data.py +++ b/tests/fuzz/fuzz_audio_data.py @@ -7,12 +7,14 @@ """Fuzzes AudioData: creation, and copy_to with every conversion of format and layout.""" +from __future__ import annotations + import sys import atheris with atheris.instrument_imports(): - from _input import Input + from inputs import Input import webrtc @@ -22,7 +24,7 @@ def check_identity(audio: webrtc.AudioData, data: bytes) -> None: - """Interleaved samples copied out in their own format are the same bytes""" + """Interleaved samples copied out in their own format are the same bytes.""" if audio.format.value.endswith('-planar'): return out = bytearray(audio.allocation_size({'plane_index': 0})) @@ -53,8 +55,12 @@ def test_one_input(data: bytes) -> None: except EXPECTED: return check_identity(audio, bytes(source)) + exercise(inp, audio) + + +def exercise(inp: Input, audio: webrtc.AudioData) -> None: for _ in range(inp.small(4)): - options = {'plane_index': inp.integer(8)} + options: dict[str, object] = {'plane_index': inp.integer(8)} if inp.flag(): options['frame_offset'] = inp.integer(512) if inp.flag(): diff --git a/tests/fuzz/fuzz_generator.py b/tests/fuzz/fuzz_generator.py index bcbe479..000c4f5 100644 --- a/tests/fuzz/fuzz_generator.py +++ b/tests/fuzz/fuzz_generator.py @@ -11,19 +11,21 @@ take. The remote tracks are read by processors, so received frames go through the native path too. """ +from __future__ import annotations + import asyncio -import os +import pathlib import sys import atheris with atheris.instrument_imports(): - from _input import Input + from inputs import Input import webrtc -sys.path.insert(0, os.path.join(os.path.dirname(__file__), '..', '..')) -from tests.helpers import connect # noqa: E402 +sys.path.insert(0, str(pathlib.Path(__file__).parent.parent.parent)) +from tests.helpers import connect EXPECTED = (TypeError, ValueError, BufferError, webrtc.NotSupportedError, webrtc.InvalidStateError) PIXEL_FORMATS = list(webrtc.VideoPixelFormat) @@ -36,7 +38,7 @@ class Session: - """A connected pair sending a generator of each kind, whose writers are replaced once they fail""" + """A connected pair sending a generator of each kind, whose writers are replaced once they fail.""" async def start(self) -> None: self.caller, self.callee = webrtc.RTCPeerConnection(), webrtc.RTCPeerConnection() @@ -49,7 +51,7 @@ async def start(self) -> None: self.writers[kind] = generator.writable.get_writer() await connect(self.caller, self.callee) - async def write(self, kind: str, chunk) -> None: + async def write(self, kind: str, chunk: webrtc.AudioData | webrtc.VideoFrame) -> None: try: await self.writers[kind].write(chunk) except EXPECTED: @@ -61,12 +63,12 @@ async def write(self, kind: str, chunk) -> None: await asyncio.sleep(0) -def audio_data(inp: Input): +def audio_data(inp: Input) -> webrtc.AudioData: format = inp.choice(SAMPLE_FORMATS) channels = inp.small(20) if inp.flag() else inp.integer(20) rate = inp.choice(RATES) if inp.flag() else inp.number(400000) frames = inp.small(8000) if inp.flag() else inp.integer(8000) - if not all(isinstance(v, int) and not isinstance(v, bool) and 0 < v for v in (channels, frames)): + if not all(isinstance(v, int) and not isinstance(v, bool) and v > 0 for v in (channels, frames)): channels, frames = 1, 480 size = min(frames * channels * SAMPLE_BYTES[format.value.split('-')[0]], 1 << 20) return webrtc.AudioData( @@ -79,7 +81,7 @@ def audio_data(inp: Input): ) -def video_frame(inp: Input): +def video_frame(inp: Input) -> webrtc.VideoFrame: format = inp.choice(PIXEL_FORMATS) width = inp.small(64) + 1 if inp.flag() else inp.choice([1, 2, 3, 15, 16, 17, 639, 640, 1920, 4096]) height = inp.small(64) + 1 if inp.flag() else inp.choice([1, 2, 3, 15, 16, 17, 479, 480, 1080, 4096]) diff --git a/tests/fuzz/fuzz_native_buffers.py b/tests/fuzz/fuzz_native_buffers.py index bcc1b23..d6df1bd 100644 --- a/tests/fuzz/fuzz_native_buffers.py +++ b/tests/fuzz/fuzz_native_buffers.py @@ -10,12 +10,15 @@ The native module must stay memory safe on its own: it may raise, but never read or write out of a buffer. """ +from __future__ import annotations + +import contextlib import sys import atheris with atheris.instrument_imports(): - from _input import Input + from inputs import Input from webrtc import wrtc @@ -30,7 +33,7 @@ EXPECTED = (TypeError, ValueError, RuntimeError, BufferError) -def frame(inp: Input): +def frame(inp: Input) -> wrtc.VideoFrameBuffer: width, height = inp.unsigned(40), inp.unsigned(40) layout = [(inp.unsigned(1 << 12), inp.unsigned(256)) for _ in range(inp.small(4))] return wrtc.VideoFrameBuffer.fromData( @@ -78,10 +81,8 @@ def audio(inp: Input) -> None: def test_one_input(data: bytes) -> None: inp = Input(data) - try: + with contextlib.suppress(EXPECTED): video(inp) if inp.flag() else audio(inp) - except EXPECTED: - pass def main() -> None: diff --git a/tests/fuzz/fuzz_video_frame.py b/tests/fuzz/fuzz_video_frame.py index cd81929..90810c1 100644 --- a/tests/fuzz/fuzz_video_frame.py +++ b/tests/fuzz/fuzz_video_frame.py @@ -7,12 +7,15 @@ """Fuzzes VideoFrame: creation from a buffer, copy_to with conversions and layouts, and frames of frames.""" +from __future__ import annotations + +import asyncio import sys import atheris with atheris.instrument_imports(): - from _input import Input + from inputs import Buffer, Input import webrtc @@ -20,17 +23,27 @@ # what the specification lets these raise, and BufferError for read-only buffers; anything else is a bug EXPECTED = (TypeError, ValueError, BufferError, webrtc.NotSupportedError, webrtc.InvalidStateError) +loop = asyncio.new_event_loop() + + +async def _copy_to(frame: webrtc.VideoFrame, destination: Buffer, options: dict[str, object] | None) -> None: + await frame.copy_to(destination, options) + -def rect(inp: Input): +def copy_to(frame: webrtc.VideoFrame, destination: Buffer, options: dict[str, object] | None) -> None: + loop.run_until_complete(_copy_to(frame, destination, options)) + + +def rect(inp: Input) -> dict[str, float]: return {'x': inp.number(), 'y': inp.number(), 'width': inp.number(), 'height': inp.number()} -def layout(inp: Input): +def layout(inp: Input) -> list[dict[str, object]]: return [{'offset': inp.integer(4096), 'stride': inp.integer(256)} for _ in range(inp.small(4))] -def copy_options(inp: Input): - options = {} +def copy_options(inp: Input) -> dict[str, object]: + options: dict[str, object] = {} if inp.flag(): options['rect'] = rect(inp) if inp.flag(): @@ -40,23 +53,27 @@ def copy_options(inp: Input): return options +def frame_of_frame(inp: Input, frame: webrtc.VideoFrame) -> webrtc.VideoFrame: + init: dict[str, object] = {'visible_rect': rect(inp)} if inp.flag() else {} + if inp.flag(): + init['alpha'] = inp.choice(['keep', 'discard']) + if inp.flag(): + init['rotation'] = inp.number(360) + init['flip'] = inp.flag() + if inp.flag(): + init['display_width'], init['display_height'] = inp.integer(), inp.integer() + return webrtc.VideoFrame(frame, init) + + def exercise(inp: Input, frame: webrtc.VideoFrame) -> None: for _ in range(inp.small(4)): action = inp.small(5) if action == 0: options = copy_options(inp) size = frame.allocation_size(options) - frame._copy_to(inp.destination(min(size, 1 << 20)), options) + copy_to(frame, inp.destination(min(size, 1 << 20)), options) elif action == 1: - init = {'visible_rect': rect(inp)} if inp.flag() else {} - if inp.flag(): - init['alpha'] = inp.choice(['keep', 'discard']) - if inp.flag(): - init['rotation'] = inp.number(360) - init['flip'] = inp.flag() - if inp.flag(): - init['display_width'], init['display_height'] = inp.integer(), inp.integer() - frame = webrtc.VideoFrame(frame, init) + frame = frame_of_frame(inp, frame) elif action == 2: frame = frame.clone() elif action == 3: @@ -65,17 +82,18 @@ def exercise(inp: Input, frame: webrtc.VideoFrame) -> None: frame.metadata() -def check_identity(inp: Input, format, width: int, height: int, data) -> None: - """A packed frame copied out as it is gives the same bytes""" +def check_identity(format: webrtc.VideoPixelFormat, size: tuple[int, int], data: Buffer) -> None: + """A packed frame copied out as it is gives the same bytes.""" + width, height = size init = {'format': format, 'coded_width': width, 'coded_height': height, 'timestamp': 0} size = webrtc.VideoFrame(bytes(1 << 16), init).allocation_size() if width * height <= 1024 else 0 if not size: return - data = bytes(data)[:size] + bytes(max(0, size - len(data))) - with webrtc.VideoFrame(data, init) as frame: + packed = bytes(data)[:size] + bytes(max(0, size - len(data))) + with webrtc.VideoFrame(packed, init) as frame: out = bytearray(size) - frame._copy_to(out, None) - assert bytes(out) == data, f'{format} {width}x{height} copied out differently' + copy_to(frame, out, None) + assert bytes(out) == packed, f'{format} {width}x{height} copied out differently' def test_one_input(data: bytes) -> None: @@ -83,10 +101,15 @@ def test_one_input(data: bytes) -> None: format = inp.choice(FORMATS) if inp.flag(): width, height = inp.small(32) + 1, inp.small(32) + 1 - check_identity(inp, format, width, height, inp.buffer(width * height * 8)) + check_identity(format, (width, height), inp.buffer(width * height * 8)) return width, height = inp.integer(), inp.integer() - init = {'format': format, 'coded_width': width, 'coded_height': height, 'timestamp': inp.integer()} + init: dict[str, object] = { + 'format': format, + 'coded_width': width, + 'coded_height': height, + 'timestamp': inp.integer(), + } if inp.flag(): init['layout'] = layout(inp) if inp.flag(): diff --git a/tests/fuzz/_input.py b/tests/fuzz/inputs.py similarity index 66% rename from tests/fuzz/_input.py rename to tests/fuzz/inputs.py index 871ebf2..07d7001 100644 --- a/tests/fuzz/_input.py +++ b/tests/fuzz/inputs.py @@ -7,29 +7,48 @@ """Values for the fuzz targets, drawn from the fuzzer's bytes and biased to the edges where checks go wrong.""" +from __future__ import annotations + +from typing import TYPE_CHECKING, Callable, TypeVar, Union + import atheris +if TYPE_CHECKING: + from collections.abc import Iterable + +T = TypeVar('T') +Buffer = Union[bytes, bytearray, memoryview] + # where sizes overflow: 16, 24 (the native limit of a frame side), 31, 32 and 64 bits EDGES = [0, 1, 2, 3, 4, 7, 8, 15, 16, 255, 256, 2**16 - 1, 2**16, 2**24, 2**24 + 1, 2**31 - 1, 2**31] EDGES += [2**32 - 1, 2**32, 2**32 + 1, 2**63 - 1, 2**63, 2**64 - 1, 2**64, 2**65] FLOATS = [0.0, -0.0, 0.5, 1.5, -1.0, 1e-300, 1e300, float('inf'), float('-inf'), float('nan')] +# the source buffers of Input.buffer: a copy, strided views and a 2D view +VIEWS: list[Callable[[bytes], Buffer]] = [ + bytearray, + lambda data: memoryview(bytearray(data * 2))[::2], + lambda data: memoryview(bytearray(data))[::-1], + lambda data: memoryview(bytearray(data)).cast('B', [len(data), 1]) if data else memoryview(bytearray()), +] class Input: - def __init__(self, data: bytes): + """The fuzzer's bytes as values.""" + + def __init__(self, data: bytes) -> None: self._fdp = atheris.FuzzedDataProvider(data) def flag(self) -> bool: return self._fdp.ConsumeBool() - def choice(self, values): + def choice(self, values: Iterable[T]) -> T: return self._fdp.PickValueInList(list(values)) def small(self, limit: int = 64) -> int: return self._fdp.ConsumeIntInRange(0, limit) def unsigned(self, limit: int = 64) -> int: - """Mostly small, sometimes an edge""" + """Mostly small, sometimes an edge.""" mode = self._fdp.ConsumeIntInRange(0, 7) if mode == 0: return self.choice(EDGES) @@ -37,8 +56,8 @@ def unsigned(self, limit: int = 64) -> int: return self._fdp.ConsumeIntInRange(0, 2**64 - 1) return self.small(limit) - def integer(self, limit: int = 64): - """An unsigned value, a negative one, or something that isn't an integer""" + def integer(self, limit: int = 64) -> object: + """An unsigned value, a negative one, or something that isn't an integer.""" mode = self._fdp.ConsumeIntInRange(0, 15) if mode == 0: return -self.unsigned(limit) - 1 @@ -48,7 +67,7 @@ def integer(self, limit: int = 64): return self.choice([None, True, '1', b'1']) return self.unsigned(limit) - def number(self, limit: int = 64): + def number(self, limit: int = 64) -> float: mode = self._fdp.ConsumeIntInRange(0, 3) if mode == 0: return self.choice(FLOATS) @@ -56,26 +75,18 @@ def number(self, limit: int = 64): return self.small(limit) + self.choice([0, 0, 0.5, 1e-9]) return self.small(limit) - def maybe(self, make): + def maybe(self, make: Callable[[], T]) -> T | None: return make() if self.flag() else None - def buffer(self, size: int): - """size bytes as bytes, a bytearray, or a view of one, which may be strided""" + def buffer(self, size: int) -> Buffer: + """Size bytes as bytes, a bytearray, or a view of one, which may be strided.""" data = self._fdp.ConsumeBytes(size) data += bytes(size - len(data)) kind = self._fdp.ConsumeIntInRange(0, 7) - if kind == 0: - return bytearray(data) - if kind == 1: - return memoryview(bytearray(data * 2))[::2] - if kind == 2: - return memoryview(bytearray(data))[::-1] - if kind == 3: - return memoryview(bytearray(data)).cast('B', [size, 1]) if size else memoryview(bytearray()) - return data + return VIEWS[kind](data) if kind < len(VIEWS) else data - def destination(self, size: int): - """A writable buffer of about size bytes""" + def destination(self, size: int) -> Buffer: + """A writable buffer of about size bytes.""" size = max(0, size + self.choice([0, 0, 0, -1, 1, 64])) kind = self._fdp.ConsumeIntInRange(0, 7) if kind == 0: diff --git a/tests/helpers.py b/tests/helpers.py index fe8241e..4b696ae 100644 --- a/tests/helpers.py +++ b/tests/helpers.py @@ -5,39 +5,49 @@ # that can be found in the LICENSE.md file in the root of the project. # +from __future__ import annotations + import asyncio import contextlib import ctypes import inspect import os +import pathlib import subprocess import sys import textwrap +from typing import TYPE_CHECKING, Callable import pytest import webrtc import wrtc +if TYPE_CHECKING: + from collections.abc import AsyncIterator, Awaitable + +#: The fixture creating connections with a configuration +CreatePC = Callable[..., webrtc.RTCPeerConnection] + -async def exchange_offer(caller, callee): +async def exchange_offer(caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection) -> None: offer = await caller.create_offer() await caller.set_local_description(offer) await callee.set_remote_description(offer) -async def exchange_answer(caller, callee): +async def exchange_answer(caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection) -> None: answer = await callee.create_answer() await callee.set_local_description(answer) await caller.set_remote_description(answer) -async def exchange_offer_answer(caller, callee): +async def exchange_offer_answer(caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection) -> None: await exchange_offer(caller, callee) await exchange_answer(caller, callee) -async def generate_answer(offer): +async def generate_answer(offer: webrtc.RTCSessionDescriptionInit) -> webrtc.RTCSessionDescriptionInit: pc = webrtc.RTCPeerConnection() await pc.set_remote_description(offer) @@ -52,32 +62,45 @@ async def generate_answer(offer): QUIET_PERIOD = 0.3 -async def _called(function): +async def _called(function: Callable[[], object]) -> object: result = function() return await result if inspect.isawaitable(result) else result -async def wait_until(predicate, what, timeout=10): - """Polls until the predicate (a function or a coroutine function) is true, or raises TimeoutError naming what - it waited for""" +async def wait_until(predicate: Callable[[], object], what: str, timeout: float = 10) -> None: + """Polls until the predicate (a function or a coroutine function) is true. + + Raises: + TimeoutError: The predicate isn't true in time, the message names what it waited for. + """ loop = asyncio.get_running_loop() deadline = loop.time() + timeout while not await _called(predicate): if loop.time() > deadline: - raise TimeoutError(f'Timed out waiting for {what}') + msg = f'Timed out waiting for {what}' + raise TimeoutError(msg) await asyncio.sleep(0.05) -async def wait_for_ice_gathering_complete(pc, timeout=10): +async def wait_for_ice_gathering_complete(pc: webrtc.RTCPeerConnection, timeout: float = 10) -> None: await wait_until(lambda: pc.ice_gathering_state == webrtc.RTCIceGatheringState.complete, 'ICE gathering', timeout) -def wait_for_event(target, name, timeout=10, predicate=None): - """Registers for the next event of a type (the next one the predicate accepts, if given) right away, and returns - an awaitable of it""" +def wait_for_event( + target: webrtc.EventTarget, + name: str, + timeout: float = 10, + *, + predicate: Callable[[webrtc.Event], bool] | None = None, +) -> Awaitable[webrtc.Event]: + """Registers for the next event of a type (the next one the predicate accepts, if given) right away. + + Returns: + An awaitable of the event. + """ future = asyncio.get_running_loop().create_future() - def on_event(event): + def on_event(event: webrtc.Event) -> None: if not future.done() and (predicate is None or predicate(event)): target.off(name, on_event) future.set_result(event) @@ -86,10 +109,12 @@ def on_event(event): return asyncio.wait_for(future, timeout) -async def next_task(): - """Lets the current task end, like awaiting a timer: what a task keeps until it ends (like the parameters of a - sender) expires, and callbacks posted meanwhile run. asyncio.sleep(0) doesn't end it: the code it resumes is like - a microtask of the same task.""" +async def next_task() -> None: + """Lets the current task end, like awaiting a timer. + + What a task keeps until it ends (like the parameters of a sender) expires, and callbacks posted meanwhile run. + asyncio.sleep(0) doesn't end it: the code it resumes is like a microtask of the same task. + """ loop = asyncio.get_running_loop() timer = loop.create_future() loop.call_later(0, timer.set_result, None) @@ -97,10 +122,10 @@ async def next_task(): # tasks adding candidates, kept until done (the loop only keeps a weak reference to a task) -_adding_candidates = set() +_adding_candidates: set[asyncio.Future[None]] = set() -async def _add_ice_candidate(pc, candidate): +async def _add_ice_candidate(pc: webrtc.RTCPeerConnection, candidate: webrtc.RTCIceCandidate) -> None: try: await pc.add_ice_candidate(candidate) except webrtc.InvalidStateError: @@ -109,11 +134,11 @@ async def _add_ice_candidate(pc, candidate): raise -def exchange_ice_candidates(caller, callee): - """Trickles the candidates of each connection to the other one""" +def exchange_ice_candidates(caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection) -> None: + """Trickles the candidates of each connection to the other one.""" for pc, other in ((caller, callee), (callee, caller)): - def on_candidate(event, other=other): + def on_candidate(event: webrtc.RTCPeerConnectionIceEvent, other: webrtc.RTCPeerConnection = other) -> None: if event.candidate is not None: task = asyncio.ensure_future(_add_ice_candidate(other, event.candidate)) _adding_candidates.add(task) @@ -122,33 +147,55 @@ def on_candidate(event, other=other): pc.on('icecandidate', on_candidate) -async def connect(caller, callee, timeout=10): - """Negotiates and waits until both connections are connected""" +async def connect(caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection, timeout: float = 10) -> None: + """Negotiates and waits until both connections are connected.""" exchange_ice_candidates(caller, callee) await exchange_offer_answer(caller, callee) - def connected(): + def connected() -> bool: return all(pc.connection_state == webrtc.RTCPeerConnectionState.connected for pc in (caller, callee)) await wait_until(connected, 'both connections to connect', timeout) -async def wait_until_unmuted(track, timeout=10): - """Waits until media arrives on a remote track, which may take longer than connecting""" +async def wait_until_unmuted(track: webrtc.MediaStreamTrack, timeout: float = 10) -> None: + """Waits until media arrives on a remote track, which may take longer than connecting.""" if track.muted: await wait_for_event(track, 'unmute', timeout) -async def connect_track(caller, callee, track, timeout=10): - """Sends a track from the caller, connects, and returns the remote track of the callee""" +def capture_mode(track: webrtc.MediaStreamTrack) -> tuple[int | None, int | None, float | None]: + """The width, height and frame rate a camera track captures at.""" + settings = track.get_settings() + return settings.width, settings.height, settings.frame_rate + + +async def connect_track( + caller: webrtc.RTCPeerConnection, + callee: webrtc.RTCPeerConnection, + track: webrtc.MediaStreamTrack, + *, + timeout: float = 10, +) -> webrtc.MediaStreamTrack: + """Sends a track from the caller, connects, and returns the remote track of the callee.""" caller.add_track(track) track_event = wait_for_event(callee, 'track', timeout) await connect(caller, callee, timeout) - return (await track_event).track - - -async def write_video(generator, data, width, height, stop, interval=1 / 30): - """Writes I420 frames of the data to a generator, one an interval, until stopped""" + event = await track_event + assert isinstance(event, webrtc.RTCTrackEvent) + return event.track + + +async def write_video( + generator: webrtc.VideoTrackGenerator, + data: bytes, + size: tuple[int, int], + *, + stop: asyncio.Event, + interval: float = 1 / 30, +) -> None: + """Writes I420 frames of the data and the size to a generator, one an interval, until stopped.""" + width, height = size writer = generator.writable.get_writer() timestamp = 0 while not stop.is_set(): @@ -160,8 +207,8 @@ async def write_video(generator, data, width, height, stop, interval=1 / 30): @contextlib.asynccontextmanager -async def writing(write, *args, **kwargs): - """Runs write(*args, stop=stop, **kwargs) in a task for the block, then stops it""" +async def writing(write: Callable[..., Awaitable[None]], *args: object, **kwargs: object) -> AsyncIterator[None]: + """Runs write(*args, stop=stop, **kwargs) in a task for the block, then stops it.""" stop = asyncio.Event() task = asyncio.ensure_future(write(*args, stop=stop, **kwargs)) try: @@ -172,14 +219,15 @@ async def writing(write, *args, **kwargs): # the memory of sanitizers (ASan quarantine, TSan shadow) hides leaks from resident memory -skip_if_sanitized = pytest.mark.skipif(wrtc._sanitized, reason='resident memory says nothing under sanitizers') +SANITIZED: bool = wrtc._sanitized +skip_if_sanitized = pytest.mark.skipif(SANITIZED, reason='resident memory says nothing under sanitizers') -def rss_bytes(): - """The resident memory of the process, in bytes""" +def rss_bytes() -> int: + """The resident memory of the process, in bytes.""" if sys.platform.startswith('linux'): - with open('/proc/self/statm') as statm: - return int(statm.read().split()[1]) * os.sysconf('SC_PAGE_SIZE') + statm = pathlib.Path('/proc/self/statm').read_text(encoding='utf-8') + return int(statm.split()[1]) * os.sysconf('SC_PAGE_SIZE') if sys.platform == 'win32': class Counters(ctypes.Structure): @@ -203,15 +251,15 @@ class Counters(ctypes.Structure): ctypes.windll.psapi.GetProcessMemoryInfo(process, ctypes.byref(counters), counters.cb) return counters.WorkingSetSize # macOS and other BSDs - return int(subprocess.check_output(['ps', '-o', 'rss=', '-p', str(os.getpid())])) * 1024 + return int(subprocess.check_output(['/bin/ps', '-o', 'rss=', '-p', str(os.getpid())])) * 1024 #: The root of the project, which has the tests package (pytest may run from elsewhere, like cibuildwheel) -ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) +ROOT = pathlib.Path(pathlib.Path(pathlib.Path(__file__).resolve()).parent).parent -def run_isolated(script, timeout=60): - """Runs a script in its own process, so a crash or a deadlock fails the test only; returns its output""" +def run_isolated(script: str, timeout: float = 60) -> str: + """Runs a script in its own process, so a crash or a deadlock fails the test only; returns its output.""" # a crash tells where: the Python stacks of every thread, and glibc's fatal errors, written to a tty otherwise env = {**os.environ, 'PYTHONFAULTHANDLER': '1', 'LIBC_FATAL_STDERR_': '1'} result = subprocess.run( @@ -221,6 +269,7 @@ def run_isolated(script, timeout=60): timeout=timeout, cwd=ROOT, env=env, + check=False, ) assert result.returncode == 0, f'exit code {result.returncode}:\n{result.stderr[-6000:]}' return result.stdout diff --git a/tests/rtc_peer_connection/test_add_track.py b/tests/rtc_peer_connection/test_add_track.py index b96d316..c5e6d23 100644 --- a/tests/rtc_peer_connection/test_add_track.py +++ b/tests/rtc_peer_connection/test_add_track.py @@ -5,6 +5,8 @@ # that can be found in the LICENSE.md file in the root of the project. # +from __future__ import annotations + import asyncio import pytest @@ -13,8 +15,8 @@ from tests.helpers import exchange_offer_answer, wait_for_ice_gathering_complete -def test_1(pc, audio_stream): - """addTrack when pc is closed should throw PythonWebRTCException with invalid state""" +def test_1(pc: webrtc.RTCPeerConnection, audio_stream: webrtc.MediaStream) -> None: + """AddTrack when pc is closed should throw PythonWebRTCException with invalid state.""" track, *_ = audio_stream.get_audio_tracks() pc.close() @@ -22,15 +24,15 @@ def test_1(pc, audio_stream): pc.add_track(track, audio_stream) -def test_2(pc, audio_stream): - """add_track with single track argument and no stream should succeed""" +def test_2(pc: webrtc.RTCPeerConnection, audio_stream: webrtc.MediaStream) -> None: + """add_track with single track argument and no stream should succeed.""" track, *_ = audio_stream.get_tracks() sender = pc.add_track(track) assert isinstance(sender, webrtc.RTCRtpSender), 'Expect sender to be instance of RTCRtpSender' - assert track == sender.track, 'Expect sender\'s track to be the added track' + assert track == sender.track, "Expect sender's track to be the added track" transceivers = pc.get_transceivers() assert len(transceivers) == 1, 'Expect only one transceiver with sender added' @@ -46,19 +48,19 @@ def test_2(pc, audio_stream): assert [receiver] == pc.get_receivers(), 'Expect only one receiver associated with transceiver added' -def test_3(pc, audio_stream): - """add_track with single track argument and single stream should succeed""" +def test_3(pc: webrtc.RTCPeerConnection, audio_stream: webrtc.MediaStream) -> None: + """add_track with single track argument and single stream should succeed.""" track, *_ = audio_stream.get_tracks() sender = pc.add_track(track, audio_stream) assert isinstance(sender, webrtc.RTCRtpSender), 'Expect sender to be instance of RTCRtpSender' - assert sender.track == track, 'Expect sender\'s track to be the added track' + assert sender.track == track, "Expect sender's track to be the added track" -def test_4(pc, audio_stream): - """add_track with single track argument and multiple streams should succeed""" +def test_4(pc: webrtc.RTCPeerConnection, audio_stream: webrtc.MediaStream) -> None: + """add_track with single track argument and multiple streams should succeed.""" track, *_ = audio_stream.get_tracks() stream2 = audio_stream.clone() @@ -67,11 +69,11 @@ def test_4(pc, audio_stream): assert isinstance(sender, webrtc.RTCRtpSender), 'Expect sender to be instance of RTCRtpSender' - assert sender.track == track, 'Expect sender\'s track to be the added track' + assert sender.track == track, "Expect sender's track to be the added track" -def test_5(pc, audio_stream): - """Adding the same track multiple times should throw RTCException""" +def test_5(pc: webrtc.RTCPeerConnection, audio_stream: webrtc.MediaStream) -> None: + """Adding the same track multiple times should throw RTCException.""" track, *_ = audio_stream.get_tracks() pc.add_track(track, audio_stream) @@ -80,8 +82,8 @@ def test_5(pc, audio_stream): pc.add_track(track, audio_stream) -def test_6(pc, audio_stream): - """add_track with existing sender with None track, same kind, and recvonly direction should reuse sender""" +def test_6(pc: webrtc.RTCPeerConnection, audio_stream: webrtc.MediaStream) -> None: + """add_track with existing sender with None track, same kind, and recvonly direction should reuse sender.""" init = webrtc.RtpTransceiverInit(direction=webrtc.TransceiverDirection.recvonly) transceiver = pc.add_transceiver(webrtc.MediaType.audio, init) @@ -97,8 +99,8 @@ def test_6(pc, audio_stream): assert [sender] == pc.get_senders() -def test_7(pc, audio_stream): - """add_track with existing sender that has not been used to send should reuse the sender""" +def test_7(pc: webrtc.RTCPeerConnection, audio_stream: webrtc.MediaStream) -> None: + """add_track with existing sender that has not been used to send should reuse the sender.""" transceiver = pc.add_transceiver(webrtc.MediaType.audio) assert transceiver.sender.track is None assert transceiver.direction == webrtc.TransceiverDirection.sendrecv @@ -111,8 +113,10 @@ def test_7(pc, audio_stream): @pytest.mark.asyncio -async def test_8(caller, callee, audio_stream): - """add_track with existing sender that has been used to send should create new sender""" +async def test_8( + caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection, audio_stream: webrtc.MediaStream +) -> None: + """add_track with existing sender that has been used to send should create new sender.""" track, *_ = audio_stream.get_tracks() transceiver = caller.add_transceiver(track) @@ -135,8 +139,8 @@ async def test_8(caller, callee, audio_stream): assert sender != transceiver.sender -def test_9(pc, audio_stream): - """add_track with existing recvonly sender with null track of a different kind should create new sender""" +def test_9(pc: webrtc.RTCPeerConnection, audio_stream: webrtc.MediaStream) -> None: + """add_track with existing recvonly sender with null track of a different kind should create new sender.""" init = webrtc.RtpTransceiverInit(direction=webrtc.TransceiverDirection.recvonly) transceiver = pc.add_transceiver(webrtc.MediaType.video, init) @@ -153,12 +157,18 @@ def test_9(pc, audio_stream): assert len(senders) == 2, 'Expect 2 senders added to connection' assert sender in senders, 'Expect senders list to include sender' - assert transceiver.sender in senders, 'Expect senders list to include first transceiver\'s sender' + assert transceiver.sender in senders, "Expect senders list to include first transceiver's sender" @pytest.mark.asyncio -async def test_10(caller, callee, audio_stream, audio_stream2): - """Adding more tracks does not generate more candidates if bundled""" +async def test_10( + caller: webrtc.RTCPeerConnection, + callee: webrtc.RTCPeerConnection, + audio_stream: webrtc.MediaStream, + *, + audio_stream2: webrtc.MediaStream, +) -> None: + """Adding more tracks does not generate more candidates if bundled.""" track, *_ = audio_stream.get_tracks() transceiver = caller.add_transceiver(track) @@ -170,12 +180,14 @@ async def test_10(caller, callee, audio_stream, audio_stream2): await wait_for_ice_gathering_complete(callee) second_track, *_ = audio_stream2.get_tracks() - - # TODO onicecandidate event should not be occurred + candidates = [] + caller.on('icecandidate', lambda event: candidates.append(event.candidate)) caller.add_track(second_track) await exchange_offer_answer(caller, callee) + await asyncio.sleep(0.1) + assert not candidates, 'Expect no icecandidate events after adding a bundled track' first_transceiver, second_transceiver, *_ = caller.get_transceivers() assert first_transceiver.receiver.transport == second_transceiver.receiver.transport @@ -183,8 +195,10 @@ async def test_10(caller, callee, audio_stream, audio_stream2): @pytest.mark.asyncio -async def test_11(caller, callee, audio_stream): - """add_track while set_remote_description(offer) is pending should reuse the transceiver the offer creates""" +async def test_11( + caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection, audio_stream: webrtc.MediaStream +) -> None: + """add_track while set_remote_description(offer) is pending should reuse the transceiver the offer creates.""" track, *_ = audio_stream.get_tracks() caller.add_track(track) diff --git a/tests/rtc_peer_connection/test_add_transceiver.py b/tests/rtc_peer_connection/test_add_transceiver.py index 41ad659..81afd7a 100644 --- a/tests/rtc_peer_connection/test_add_transceiver.py +++ b/tests/rtc_peer_connection/test_add_transceiver.py @@ -5,21 +5,22 @@ # that can be found in the LICENSE.md file in the root of the project. # +from __future__ import annotations + import pytest import webrtc -def test_1(pc): - """add_transceiver with string argument as invalid kind should throw TypeError""" - +def test_1(pc: webrtc.RTCPeerConnection) -> None: + """add_transceiver with string argument as invalid kind should throw TypeError.""" assert hasattr(pc, 'add_transceiver') with pytest.raises(TypeError): pc.add_transceiver('invalid') -def _create_and_test_transceiver(pc, kind): +def _create_and_test_transceiver(pc: webrtc.RTCPeerConnection, kind: webrtc.MediaType) -> None: assert hasattr(pc, 'add_transceiver') transceiver = pc.add_transceiver(kind) @@ -32,7 +33,7 @@ def _create_and_test_transceiver(pc, kind): assert transceiver.current_direction is None assert [transceiver] == pc.get_transceivers(), ( - 'Expect added transceiver to be the only element in connection\'s list of transceivers' + "Expect added transceiver to be the only element in connection's list of transceivers" ) sender = transceiver.sender @@ -41,7 +42,7 @@ def _create_and_test_transceiver(pc, kind): assert sender.track is None - assert [sender] == pc.get_senders(), 'Expect added sender to be the only element in connection\'s list of senders' + assert [sender] == pc.get_senders(), "Expect added sender to be the only element in connection's list of senders" receiver = transceiver.receiver assert isinstance(receiver, webrtc.RTCRtpReceiver) @@ -53,37 +54,36 @@ def _create_and_test_transceiver(pc, kind): assert track.ready_state == webrtc.MediaStreamTrackState.live assert [receiver] == pc.get_receivers(), ( - 'Expect added receiver to be the only element in connection\'s list of receivers' + "Expect added receiver to be the only element in connection's list of receivers" ) -def test_2(pc): - """add_transceiver('audio') should return an audio transceiver""" +def test_2(pc: webrtc.RTCPeerConnection) -> None: + """add_transceiver('audio') should return an audio transceiver.""" _create_and_test_transceiver(pc, webrtc.MediaType.audio) -def test_3(pc): - """add_transceiver('video') should return a video transceiver""" +def test_3(pc: webrtc.RTCPeerConnection) -> None: + """add_transceiver('video') should return a video transceiver.""" _create_and_test_transceiver(pc, webrtc.MediaType.video) -def test_4(pc): - """add_transceiver with direction inactive should have result transceiver.direction be the same""" +def test_4(pc: webrtc.RTCPeerConnection) -> None: + """add_transceiver with direction inactive should have result transceiver.direction be the same.""" init = webrtc.RtpTransceiverInit(direction=webrtc.TransceiverDirection.inactive) transceiver = pc.add_transceiver(webrtc.MediaType.audio, init) assert transceiver.direction == webrtc.TransceiverDirection.inactive -def test_5(pc): - """add_transceiver with invalid direction should throw TypeError""" +def test_5(pc: webrtc.RTCPeerConnection) -> None: + """add_transceiver with invalid direction should throw TypeError.""" with pytest.raises(TypeError): - init = webrtc.RtpTransceiverInit(direction='invalid') - pc.add_transceiver(webrtc.MediaType.audio, init) + pc.add_transceiver(webrtc.MediaType.audio, {'direction': 'invalid'}) -def test_6(pc, audio_stream): - """add_transceiver(track) should have result with sender.track be given track""" +def test_6(pc: webrtc.RTCPeerConnection, audio_stream: webrtc.MediaStream) -> None: + """add_transceiver(track) should have result with sender.track be given track.""" track, *_ = audio_stream.get_tracks() transceiver = pc.add_transceiver(track) sender, receiver = transceiver.sender, transceiver.receiver @@ -99,24 +99,24 @@ def test_6(pc, audio_stream): 'Expect receiver.track to be instance of MediaStreamTrack' ) assert receiver_track.kind == webrtc.MediaType.audio, ( - 'receiver.track should have the same kind as added track\'s kind' + "receiver.track should have the same kind as added track's kind" ) assert receiver_track.ready_state == webrtc.MediaStreamTrackState.live assert [transceiver] == pc.get_transceivers(), ( - 'Expect added transceiver to be the only element in connection\'s list of transceivers' + "Expect added transceiver to be the only element in connection's list of transceivers" ) - assert [sender] == pc.get_senders(), 'Expect added sender to be the only element in connection\'s list of senders' + assert [sender] == pc.get_senders(), "Expect added sender to be the only element in connection's list of senders" assert [receiver] == pc.get_receivers(), ( - 'Expect added receiver to be the only element in connection\'s list of receivers' + "Expect added receiver to be the only element in connection's list of receivers" ) -def test_7(pc, audio_stream): - """add_transceiver(track) multiple times should create multiple transceivers""" +def test_7(pc: webrtc.RTCPeerConnection, audio_stream: webrtc.MediaStream) -> None: + """add_transceiver(track) multiple times should create multiple transceivers.""" track, *_ = audio_stream.get_tracks() transceiver1 = pc.add_transceiver(track) transceiver2 = pc.add_transceiver(track) @@ -144,51 +144,51 @@ def test_7(pc, audio_stream): @pytest.mark.parametrize('kind', [webrtc.MediaType.video, webrtc.MediaType.audio]) -def test_8(pc, kind): - """add_transceiver with rid containing invalid non-alphanumeric characters should throw ValueError""" - encodings = [webrtc.RtpEncodingParameters(rid="@Invalid!")] +def test_8(pc: webrtc.RTCPeerConnection, kind: webrtc.MediaType) -> None: + """add_transceiver with rid containing invalid non-alphanumeric characters should throw ValueError.""" + encodings = [webrtc.RTCRtpEncodingParameters(rid='@Invalid!')] init = webrtc.RtpTransceiverInit(send_encodings=encodings) - with pytest.raises(ValueError): + with pytest.raises(ValueError, match='is not a valid rid'): pc.add_transceiver(kind, init) @pytest.mark.parametrize('kind', [webrtc.MediaType.video, webrtc.MediaType.audio]) -def test_9(pc, kind): - """add_transceiver with rid longer than 16 characters should throw ValueError""" - encodings = [webrtc.RtpEncodingParameters(rid="a" * 17)] +def test_9(pc: webrtc.RTCPeerConnection, kind: webrtc.MediaType) -> None: + """add_transceiver with rid longer than 16 characters should throw ValueError.""" + encodings = [webrtc.RTCRtpEncodingParameters(rid='a' * 17)] init = webrtc.RtpTransceiverInit(send_encodings=encodings) - with pytest.raises(ValueError): + with pytest.raises(ValueError, match='is not a valid rid'): pc.add_transceiver(kind, init) @pytest.mark.parametrize('kind', [webrtc.MediaType.video, webrtc.MediaType.audio]) -def test_10(pc, kind): - """add_transceiver with valid rid value should succeed""" - encodings = [webrtc.RtpEncodingParameters(rid="foo")] +def test_10(pc: webrtc.RTCPeerConnection, kind: webrtc.MediaType) -> None: + """add_transceiver with valid rid value should succeed.""" + encodings = [webrtc.RTCRtpEncodingParameters(rid='foo')] init = webrtc.RtpTransceiverInit(send_encodings=encodings) pc.add_transceiver(kind, init) -def test_11(pc): - """add_transceiver with valid sendEncodings should succeed""" - # dtx and ptime from the original test aren't supported by RtpEncodingParameters - encodings = [webrtc.RtpEncodingParameters(active=False, max_bitrate=8, max_framerate=25, rid="foo")] +def test_11(pc: webrtc.RTCPeerConnection) -> None: + """add_transceiver with valid sendEncodings should succeed.""" + # dtx and ptime from the original test aren't supported by RTCRtpEncodingParameters + encodings = [webrtc.RTCRtpEncodingParameters(active=False, max_bitrate=8, max_framerate=25, rid='foo')] init = webrtc.RtpTransceiverInit(send_encodings=encodings) pc.add_transceiver(webrtc.MediaType.video, init) -def test_12(pc): - """add_transceiver with direction sendonly should have result transceiver.direction be the same""" +def test_12(pc: webrtc.RTCPeerConnection) -> None: + """add_transceiver with direction sendonly should have result transceiver.direction be the same.""" init = webrtc.RtpTransceiverInit(direction=webrtc.TransceiverDirection.sendonly) transceiver = pc.add_transceiver(webrtc.MediaType.audio, init) assert transceiver.direction == webrtc.TransceiverDirection.sendonly -def test_13(pc): - """add_transceiver with multiple rid values should succeed""" - encodings = [webrtc.RtpEncodingParameters(rid="a"), webrtc.RtpEncodingParameters(rid="b")] +def test_13(pc: webrtc.RTCPeerConnection) -> None: + """add_transceiver with multiple rid values should succeed.""" + encodings = [webrtc.RTCRtpEncodingParameters(rid='a'), webrtc.RTCRtpEncodingParameters(rid='b')] init = webrtc.RtpTransceiverInit(send_encodings=encodings) pc.add_transceiver(webrtc.MediaType.video, init) diff --git a/tests/rtc_rtp_transceiver/test_direction.py b/tests/rtc_rtp_transceiver/test_direction.py index d0f638d..fa35589 100644 --- a/tests/rtc_rtp_transceiver/test_direction.py +++ b/tests/rtc_rtp_transceiver/test_direction.py @@ -5,14 +5,16 @@ # that can be found in the LICENSE.md file in the root of the project. # +from __future__ import annotations + import pytest import webrtc from tests.helpers import generate_answer -def test_1(pc): - """setting direction should change transceiver.direction""" +def test_1(pc: webrtc.RTCPeerConnection) -> None: + """Setting direction should change transceiver.direction.""" transceiver = pc.add_transceiver(webrtc.MediaType.audio) assert transceiver.direction == webrtc.TransceiverDirection.sendrecv @@ -23,8 +25,8 @@ def test_1(pc): assert transceiver.current_direction is None, 'Expect transceiver.currentDirection to not change' -def test_2(pc): - """setting direction with same direction should have no effect""" +def test_2(pc: webrtc.RTCPeerConnection) -> None: + """Setting direction with same direction should have no effect.""" init = webrtc.RtpTransceiverInit(direction=webrtc.TransceiverDirection.sendonly) transceiver = pc.add_transceiver(webrtc.MediaType.audio, init) @@ -34,8 +36,8 @@ def test_2(pc): @pytest.mark.asyncio -async def test_3(pc): - """setting direction should change transceiver.direction independent of transceiver.currentDirection""" +async def test_3(pc: webrtc.RTCPeerConnection) -> None: + """Setting direction should change transceiver.direction independent of transceiver.currentDirection.""" init = webrtc.RtpTransceiverInit(direction=webrtc.TransceiverDirection.recvonly) transceiver = pc.add_transceiver(webrtc.MediaType.audio, init) diff --git a/tests/rtc_rtp_transceiver/test_stop.py b/tests/rtc_rtp_transceiver/test_stop.py index 809f404..c36e08d 100644 --- a/tests/rtc_rtp_transceiver/test_stop.py +++ b/tests/rtc_rtp_transceiver/test_stop.py @@ -5,6 +5,8 @@ # that can be found in the LICENSE.md file in the root of the project. # +from __future__ import annotations + import pytest import webrtc @@ -12,8 +14,8 @@ @pytest.mark.asyncio -async def test_1(pc): - """A transceiver added and stopped before the initial offer should not get an m-section in it""" +async def test_1(pc: webrtc.RTCPeerConnection) -> None: + """A transceiver added and stopped before the initial offer should not get an m-section in it.""" init = webrtc.RtpTransceiverInit(direction=webrtc.TransceiverDirection.sendonly) pc.add_transceiver(webrtc.MediaType.audio, init) pc.add_transceiver(webrtc.MediaType.video) @@ -21,12 +23,12 @@ async def test_1(pc): offer = await pc.create_offer() - assert "m=audio" not in offer.sdp, 'offer should not contain an audio m-section' - assert "m=video" in offer.sdp, 'offer should contain a video m-section' + assert 'm=audio' not in offer.sdp, 'offer should not contain an audio m-section' + assert 'm=video' in offer.sdp, 'offer should contain a video m-section' -def test_2(pc): - """A transceiver added and stopped should not crash when getting receiver's transport""" +def test_2(pc: webrtc.RTCPeerConnection) -> None: + """A transceiver added and stopped should not crash when getting receiver's transport.""" init = webrtc.RtpTransceiverInit(direction=webrtc.TransceiverDirection.sendonly) pc.add_transceiver(webrtc.MediaType.audio, init) pc.add_transceiver(webrtc.MediaType.video) @@ -37,8 +39,8 @@ def test_2(pc): @pytest.mark.asyncio -async def test_3(caller, callee): - """During renegotiation, a transceiver added and stopped should not get an m-section in the offer""" +async def test_3(caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection) -> None: + """During renegotiation, a transceiver added and stopped should not get an m-section in the offer.""" caller.add_transceiver(webrtc.MediaType.audio) await exchange_offer_answer(caller, callee) @@ -56,7 +58,9 @@ async def test_3(caller, callee): assert 'm=video' not in offer.sdp, 'offer should not contain a video m-section' -async def _test_inactive_m_section(caller, callee, direction): +async def _test_inactive_m_section( + caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection, direction: webrtc.TransceiverDirection +) -> None: caller.add_transceiver(webrtc.MediaType.audio) await exchange_offer_answer(caller, callee) @@ -70,20 +74,20 @@ async def _test_inactive_m_section(caller, callee, direction): @pytest.mark.asyncio -async def test_4(caller, callee): - """A stopped sendonly transceiver should generate an inactive m-section in the offer""" +async def test_4(caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection) -> None: + """A stopped sendonly transceiver should generate an inactive m-section in the offer.""" await _test_inactive_m_section(caller, callee, webrtc.TransceiverDirection.sendonly) @pytest.mark.asyncio -async def test_5(caller, callee): - """A stopped inactive transceiver should generate an inactive m-section in the offer""" +async def test_5(caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection) -> None: + """A stopped inactive transceiver should generate an inactive m-section in the offer.""" await _test_inactive_m_section(caller, callee, webrtc.TransceiverDirection.inactive) @pytest.mark.asyncio -async def test_6(caller, callee): - """If a transceiver is stopped locally, setting a locally generated answer should still work""" +async def test_6(caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection) -> None: + """If a transceiver is stopped locally, setting a locally generated answer should still work.""" caller.add_transceiver(webrtc.MediaType.audio) await exchange_offer_answer(caller, callee) @@ -94,8 +98,8 @@ async def test_6(caller, callee): @pytest.mark.asyncio -async def test_7(caller, callee): - """If a transceiver is stopped remotely, setting a locally generated answer should still work""" +async def test_7(caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection) -> None: + """If a transceiver is stopped remotely, setting a locally generated answer should still work.""" caller.add_transceiver(webrtc.MediaType.audio) await exchange_offer_answer(caller, callee) @@ -106,8 +110,8 @@ async def test_7(caller, callee): @pytest.mark.asyncio -async def test_8(caller, callee): - """If a transceiver is stopped, transceivers, senders and receivers should disappear after offer/answer""" +async def test_8(caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection) -> None: + """If a transceiver is stopped, transceivers, senders and receivers should disappear after offer/answer.""" caller.add_transceiver(webrtc.MediaType.audio) await exchange_offer_answer(caller, callee) @@ -130,8 +134,8 @@ async def test_8(caller, callee): @pytest.mark.asyncio -async def test_9(caller, callee): - """If a transceiver is stopped, transceivers should end up in state stopped""" +async def test_9(caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection) -> None: + """If a transceiver is stopped, transceivers should end up in state stopped.""" caller.add_transceiver(webrtc.MediaType.audio) await exchange_offer_answer(caller, callee) diff --git a/tests/test_audio_data.py b/tests/test_audio_data.py index 8d6ef05..6e45f27 100644 --- a/tests/test_audio_data.py +++ b/tests/test_audio_data.py @@ -7,6 +7,8 @@ """AudioData of WebCodecs: construction, copies and sample conversions.""" +from __future__ import annotations + import array import struct @@ -16,12 +18,14 @@ from webrtc import AudioSampleFormat -def f32(*values): +def f32(*values: float) -> bytes: return array.array('f', values).tobytes() -def audio_data(format='f32-planar', channels=2, frames=5, data=None, **init): - size = {'u8': 1, 's16': 2}.get(format.split('-')[0], 4) +def audio_data( + *, format: str = 'f32-planar', channels: int = 2, frames: int = 5, data: bytes | None = None, **init: object +) -> webrtc.AudioData: + size = {'u8': 1, 's16': 2}.get(format.split('-', maxsplit=1)[0], 4) return webrtc.AudioData( format=format, sample_rate=8000, @@ -33,19 +37,19 @@ def audio_data(format='f32-planar', channels=2, frames=5, data=None, **init): ) -def test_construct(): - """AudioData has the attributes of its init""" +def test_construct() -> None: + """AudioData has the attributes of its init.""" audio = audio_data(frames=100) assert audio.format == AudioSampleFormat.f32_planar assert (audio.sample_rate, audio.number_of_frames, audio.number_of_channels) == (8000, 100, 2) - assert audio.duration == 100 / 8000 * 1_000_000 + assert audio.duration == 12_500 assert audio.timestamp == 1234 assert audio.numberOfFrames == 100 audio.close() -def test_init_as_dictionary(): - """The init is also a dictionary with camelCase names""" +def test_init_as_dictionary() -> None: + """The init is also a dictionary with camelCase names.""" init = { 'format': 's16', 'sampleRate': 48000, @@ -68,14 +72,14 @@ def test_init_as_dictionary(): {'data': bytes(3)}, ], ) -def test_invalid_init(change): - """An invalid init, or data too small for it, is a TypeError""" +def test_invalid_init(change: dict[str, object]) -> None: + """An invalid init, or data too small for it, is a TypeError.""" with pytest.raises(TypeError): audio_data(**change) -def test_close_and_clone(): - """A closed data has no samples, a clone is closed separately""" +def test_close_and_clone() -> None: + """A closed data has no samples, a clone is closed separately.""" audio = audio_data() clone = audio.clone() audio.close() @@ -87,8 +91,8 @@ def test_close_and_clone(): clone.close() -def test_copy_frames_of_a_plane(): - """copyTo copies frame_count frames from frame_offset, of one plane""" +def test_copy_frames_of_a_plane() -> None: + """CopyTo copies frame_count frames from frame_offset, of one plane.""" audio = audio_data(data=f32(1, 2, 3, 4, 5, 6, 7, 8, 9, 10)) out = bytearray(12) options = webrtc.AudioDataCopyToOptions(plane_index=1, frame_offset=1, frame_count=3) @@ -98,8 +102,8 @@ def test_copy_frames_of_a_plane(): audio.close() -def test_copy_to_interleaved_and_planar(): - """Planar data copies to an interleaved format with every channel, and back one channel at a time""" +def test_copy_to_interleaved_and_planar() -> None: + """Planar data copies to an interleaved format with every channel, and back one channel at a time.""" audio = audio_data(data=f32(1, 2, 3, 4, 5, 6, 7, 8, 9, 10)) out = bytearray(40) audio.copy_to(out, {'planeIndex': 0, 'format': 'f32'}) @@ -121,16 +125,16 @@ def test_copy_to_interleaved_and_planar(): {'plane_index': 0, 'frame_offset': 1, 'frame_count': 5}, ], ) -def test_copy_ranges(options): - """Planes and frames that don't exist are a RangeError""" +def test_copy_ranges(options: dict[str, object]) -> None: + """Planes and frames that don't exist are a RangeError.""" audio = audio_data() with pytest.raises(webrtc.InvalidRangeError): audio.copy_to(bytearray(100), options) audio.close() -def test_destination_too_small(): - """A destination smaller than the copy is a RangeError""" +def test_destination_too_small() -> None: + """A destination smaller than the copy is a RangeError.""" audio = audio_data() with pytest.raises(webrtc.InvalidRangeError): audio.copy_to(bytearray(19), {'plane_index': 0}) @@ -148,8 +152,8 @@ def test_destination_too_small(): @pytest.mark.parametrize('source', VALUES) @pytest.mark.parametrize('destination', VALUES) -def test_sample_conversions(source, destination): - """Samples convert between types, scaled to their range""" +def test_sample_conversions(source: str, destination: str) -> None: + """Samples convert between types, scaled to their range.""" values, code = VALUES[source] audio = audio_data(format=source, channels=1, frames=4, data=array.array(code, values).tobytes()) expected, destination_code = VALUES[destination] @@ -163,8 +167,8 @@ def test_sample_conversions(source, destination): @pytest.mark.parametrize('destination', ['u8', 's16', 's32']) -def test_non_finite_f32_samples_convert(destination): - """NaN is silence and infinities are the extremes; converting NaN was undefined behavior (found by fuzzing)""" +def test_non_finite_f32_samples_convert(destination: str) -> None: + """NaN is silence and infinities are the extremes; converting NaN was undefined behavior (found by fuzzing).""" values = [float('nan'), float('inf'), float('-inf')] audio = audio_data(format='f32', channels=1, frames=3, data=array.array('f', values).tobytes()) silence, maximum, minimum = VALUES[destination][0][3], VALUES[destination][0][1], VALUES[destination][0][0] @@ -174,8 +178,8 @@ def test_non_finite_f32_samples_convert(destination): audio.close() -def test_s16_bytes_are_little_endian(): - """s16 samples are little endian, scaled by 1/32768 to f32""" +def test_s16_bytes_are_little_endian() -> None: + """s16 samples are little endian, scaled by 1/32768 to f32.""" audio = audio_data(format='s16', channels=1, frames=2, data=struct.pack('<2h', 1, -1)) out = bytearray(8) audio.copy_to(out, {'plane_index': 0, 'format': 'f32'}) diff --git a/tests/test_configuration.py b/tests/test_configuration.py index fe23e45..d3d58ef 100644 --- a/tests/test_configuration.py +++ b/tests/test_configuration.py @@ -7,17 +7,24 @@ """Configurations, certificates, descriptions, ICE candidates and errors: the models of a connection.""" +from __future__ import annotations + import asyncio +import base64 import dataclasses import time +from typing import TYPE_CHECKING import pytest import webrtc from tests.helpers import exchange_offer +if TYPE_CHECKING: + from tests.helpers import CreatePC + -def configuration(): +def configuration() -> webrtc.RTCConfiguration: return webrtc.RTCConfiguration( ice_servers=[ webrtc.RTCIceServer('stun:stun.example.org'), @@ -30,23 +37,24 @@ def configuration(): ) -def test_configuration_round_trip(create_pc): - """A connection reports the configuration it was created with""" +def test_configuration_round_trip(create_pc: CreatePC) -> None: + """A connection reports the configuration it was created with.""" got = create_pc(configuration()).get_configuration() assert [server.urls for server in got.ice_servers] == [ ['stun:stun.example.org'], ['turn:turn.example.org:3478?transport=tcp'], ] - assert got.ice_servers[1].username == 'user' and got.ice_servers[1].credential == 'pass' + assert got.ice_servers[1].username == 'user' + assert got.ice_servers[1].credential == 'pass' assert got.ice_transport_policy == webrtc.RTCIceTransportPolicy.relay assert got.bundle_policy == webrtc.RTCBundlePolicy.max_bundle assert got.ice_candidate_pool_size == 2 assert got.port_range == (40000, 40100) -def test_set_configuration_defaults(create_pc): - """Members set_configuration isn't given get their defaults, which may not change what can't be changed""" +def test_set_configuration_defaults(create_pc: CreatePC) -> None: + """Members set_configuration isn't given get their defaults, which may not change what can't be changed.""" pc = create_pc(configuration()) with pytest.raises(webrtc.InvalidModificationError): # the default bundle policy isn't the one of the connection @@ -55,8 +63,8 @@ def test_set_configuration_defaults(create_pc): assert pc.get_configuration().ice_servers == [] -def test_set_configuration_of_closed_connection(pc): - """A closed connection can't be configured""" +def test_set_configuration_of_closed_connection(pc: webrtc.RTCPeerConnection) -> None: + """A closed connection can't be configured.""" pc.close() with pytest.raises(webrtc.InvalidStateError): pc.set_configuration() @@ -77,48 +85,50 @@ def test_set_configuration_of_closed_connection(pc): 'stun:2001:db8::1', ], ) -def test_invalid_ice_server_urls(url): - """Malformed ICE server URLs are an InvalidSyntaxError""" +def test_invalid_ice_server_urls(url: str) -> None: + """Malformed ICE server URLs are an InvalidSyntaxError.""" with pytest.raises(webrtc.InvalidSyntaxError): webrtc.RTCPeerConnection(webrtc.RTCConfiguration(ice_servers=[webrtc.RTCIceServer(url, 'u', 'p')])) -def test_ice_server_needs_credentials(): - """A TURN server needs a username and a credential""" +def test_ice_server_needs_credentials() -> None: + """A TURN server needs a username and a credential.""" with pytest.raises(webrtc.InvalidAccessError): webrtc.RTCPeerConnection(webrtc.RTCConfiguration(ice_servers=[webrtc.RTCIceServer('turn:example.org')])) -def test_ice_server_needs_urls(): - """An ICE server needs at least one URL""" +def test_ice_server_needs_urls() -> None: + """An ICE server needs at least one URL.""" with pytest.raises(webrtc.InvalidSyntaxError): webrtc.RTCPeerConnection(webrtc.RTCConfiguration(ice_servers=[webrtc.RTCIceServer([])])) -def test_ice_candidate_pool_size_range(): - """The candidate pool has at most 255 candidates""" - with pytest.raises(ValueError): +def test_ice_candidate_pool_size_range() -> None: + """The candidate pool has at most 255 candidates.""" + with pytest.raises(ValueError, match='from 0 to 255'): webrtc.RTCPeerConnection(webrtc.RTCConfiguration(ice_candidate_pool_size=256)) -def test_ice_server_of_ipv6_address(create_pc): - """An IPv6 address of an ICE server is in brackets""" +def test_ice_server_of_ipv6_address(create_pc: CreatePC) -> None: + """An IPv6 address of an ICE server is in brackets.""" create_pc(webrtc.RTCConfiguration(ice_servers=[webrtc.RTCIceServer('stun:[2001:db8::1]:3478')])) -def test_oauth_ice_server(): - """An OAuth credential is an RTCOAuthCredential, which libwebrtc doesn't support""" +def test_oauth_ice_server() -> None: + """An OAuth credential is an RTCOAuthCredential, which libwebrtc doesn't support.""" server = {'urls': 'turns:turn.example.org', 'username': 'user', 'credential': 'cred', 'credential_type': 'oauth'} with pytest.raises(webrtc.InvalidAccessError): webrtc.RTCPeerConnection(webrtc.RTCConfiguration(ice_servers=[server])) - server['credential'] = webrtc.RTCOAuthCredential(mac_key='a2V5', access_token='dG9rZW4=') + server['credential'] = webrtc.RTCOAuthCredential( + mac_key=base64.b64encode(b'key').decode(), access_token=base64.b64encode(b'token').decode() + ) with pytest.raises(webrtc.NotSupportedError): webrtc.RTCPeerConnection(webrtc.RTCConfiguration(ice_servers=[server])) @pytest.mark.asyncio -async def test_always_negotiate_data_channels_and_header_encryption(create_pc): - """As configured, an offer always has a data section and encrypts header extensions; neither can change""" +async def test_always_negotiate_data_channels_and_header_encryption(create_pc: CreatePC) -> None: + """As configured, an offer always has a data section and encrypts header extensions; neither can change.""" pc = create_pc( webrtc.RTCConfiguration( always_negotiate_data_channels=True, @@ -140,20 +150,25 @@ async def test_always_negotiate_data_channels_and_header_encryption(create_pc): @pytest.mark.asyncio -async def test_generate_ecdsa_certificate(): - """An ECDSA certificate expires in the future and has a SHA-256 fingerprint""" +async def test_generate_ecdsa_certificate() -> None: + """An ECDSA certificate expires in the future and has a SHA-256 fingerprint.""" certificate = await webrtc.RTCPeerConnection.generate_certificate('ECDSA') - assert certificate.expires > time.time() * 1000 and not certificate.expired + assert certificate.expires > time.time() * 1000 + assert not certificate.expired (fingerprint,) = certificate.get_fingerprints() - assert fingerprint.algorithm == 'sha-256' and len(fingerprint.value.split(':')) == 32 + assert fingerprint.algorithm == 'sha-256' + assert len(fingerprint.value.split(':')) == 32 @pytest.mark.asyncio -async def test_generate_rsa_certificate(): - """An RSA certificate is generated from WebCrypto parameters""" - rsa = await webrtc.RTCCertificate.generate( - {'name': 'RSASSA-PKCS1-v1_5', 'modulus_length': 1024, 'public_exponent': 65537, 'hash': 'SHA-256'} - ) +async def test_generate_rsa_certificate() -> None: + """An RSA certificate is generated from WebCrypto parameters.""" + rsa = await webrtc.RTCCertificate.generate({ + 'name': 'RSASSA-PKCS1-v1_5', + 'modulus_length': 1024, + 'public_exponent': 65537, + 'hash': 'SHA-256', + }) assert not rsa.expired @@ -162,15 +177,15 @@ async def test_generate_rsa_certificate(): 'algorithm', ['nonsense', {'name': 'RSASSA-PKCS1-v1_5', 'modulusLength': 2048, 'publicExponent': 3, 'hash': 'SHA-1'}], ) -async def test_generate_unsupported_certificate(algorithm): - """Algorithms other than ECDSA and RSASSA-PKCS1-v1_5 with SHA-256 and the exponent 65537 aren't supported""" +async def test_generate_unsupported_certificate(algorithm: str | dict[str, object]) -> None: + """Algorithms other than ECDSA and RSASSA-PKCS1-v1_5 with SHA-256 and the exponent 65537 aren't supported.""" with pytest.raises(webrtc.NotSupportedError): await webrtc.RTCCertificate.generate(algorithm) @pytest.mark.asyncio -async def test_configured_certificate(create_pc): - """The certificate of a configuration is the one of the connection, whose offer has its fingerprint""" +async def test_configured_certificate(create_pc: CreatePC) -> None: + """The certificate of a configuration is the one of the connection, whose offer has its fingerprint.""" certificate = await webrtc.RTCCertificate.generate('ECDSA') (fingerprint,) = certificate.get_fingerprints() pc = create_pc(webrtc.RTCConfiguration(certificates=[certificate])) @@ -181,8 +196,8 @@ async def test_configured_certificate(create_pc): @pytest.mark.asyncio -async def test_certificates_can_not_change(create_pc): - """A configuration without certificates keeps the ones of the connection, other ones aren't allowed""" +async def test_certificates_can_not_change(create_pc: CreatePC) -> None: + """A configuration without certificates keeps the ones of the connection, other ones aren't allowed.""" ecdsa, other = [await webrtc.RTCCertificate.generate('ECDSA') for _ in range(2)] pc = create_pc(webrtc.RTCConfiguration(certificates=[ecdsa])) pc.set_configuration(webrtc.RTCConfiguration()) @@ -191,8 +206,8 @@ async def test_certificates_can_not_change(create_pc): @pytest.mark.asyncio -async def test_expired_certificate(): - """A connection can't be created with an expired certificate""" +async def test_expired_certificate() -> None: + """A connection can't be created with an expired certificate.""" expired = await webrtc.RTCCertificate.generate('ECDSA', expires=0) # it expires at the millisecond it's generated, which has passed 10 ms later await asyncio.sleep(0.01) @@ -200,8 +215,8 @@ async def test_expired_certificate(): webrtc.RTCPeerConnection(webrtc.RTCConfiguration(certificates=[expired])) -def test_ice_candidate_parsing(): - """The attributes of a candidate are parsed from its candidate string""" +def test_ice_candidate_parsing() -> None: + """The attributes of a candidate are parsed from its candidate string.""" candidate = webrtc.RTCIceCandidate( 'candidate:1 2 TCP 1845501695 192.168.0.1 4444 typ srflx raddr 10.0.0.1 rport 5555 tcptype active', sdp_mid='0', @@ -221,49 +236,60 @@ def test_ice_candidate_parsing(): assert webrtc.RTCIceCandidate.from_json(candidate.to_json()).candidate == candidate.candidate -def test_invalid_ice_candidate(): - """An invalid candidate string isn't validated, nor parsed, but a candidate needs an m-line""" +def test_invalid_ice_candidate() -> None: + """An invalid candidate string isn't validated, nor parsed, but a candidate needs an m-line.""" invalid = webrtc.RTCIceCandidate('a=candidate:1 1 udp 1 1.2.3.4 5 typ host', sdp_m_line_index=0) - assert invalid.foundation is None and invalid.port is None + assert invalid.foundation is None + assert invalid.port is None with pytest.raises(TypeError): webrtc.RTCIceCandidate('candidate:1 1 udp 1 1.2.3.4 5 typ host') -def test_ice_candidate_ufrag_characters(): - """A ufrag may have "/" and "+" """ +def test_ice_candidate_ufrag_characters() -> None: + """A ufrag may have "/" and "+".""" host = webrtc.RTCIceCandidate('candidate:1 1 udp 2121940223 ::1 60645 typ host ufrag /h4t+', sdp_mid='0') assert host.type == webrtc.RTCIceCandidateType.host -def test_peer_reflexive_candidate_hides_its_address(): - """A remote peer-reflexive candidate, which the remote peer didn't signal, exposes no address""" +def test_peer_reflexive_candidate_hides_its_address() -> None: + """A remote peer-reflexive candidate, which the remote peer didn't signal, exposes no address.""" # made from what libwebrtc reports of a selected pair (see RTCIceTransport), as connectivity checks with # a peer-reflexive candidate can't be produced on demand - candidate = webrtc.RTCIceCandidate._peer_reflexive( - { - 'candidate': 'candidate:1 1 udp 1853504767 redacted-ip.invalid 62341 typ prflx generation 0 ufrag a/b+', - 'sdp_mid': '0', - } - ) - assert candidate.candidate == '' + candidate = webrtc.RTCIceCandidate._peer_reflexive({ + 'candidate': 'candidate:1 1 udp 1853504767 redacted-ip.invalid 62341 typ prflx generation 0 ufrag a/b+', + 'sdp_mid': '0', + }) + assert not candidate.candidate assert candidate.type == webrtc.RTCIceCandidateType.prflx assert candidate.address is None assert candidate.port == 62341 -def test_rtc_error(): - """An RTCError is an OperationError with the details of the error""" - error = webrtc.RTCError('sctp-failure', 'failed', sctp_cause_code=12) - assert isinstance(error, webrtc.OperationError) and isinstance(error, webrtc.RTCException) +def test_rtc_error() -> None: + """An RTCError is an OperationError with the details of the error.""" + error = webrtc.RTCError(webrtc.RTCErrorInit('sctp-failure', sctp_cause_code=12), 'failed') + assert isinstance(error, webrtc.OperationError) + assert isinstance(error, webrtc.RTCException) assert error.error_detail == webrtc.RTCErrorDetailType.sctp_failure - assert error.sctp_cause_code == 12 and error.sdp_line_number is None - with pytest.raises(ValueError): - webrtc.RTCError('nonsense') + assert error.sctp_cause_code == 12 + assert error.sdp_line_number is None + with pytest.raises(ValueError, match='not a valid RTCErrorDetailType'): + webrtc.RTCErrorInit('nonsense') + + +def test_rtc_error_options_dict() -> None: + """Options can be a dictionary with camelCase names, as in browsers.""" + error = webrtc.RTCError({'errorDetail': 'sdp-syntax-error', 'sdpLineNumber': 3, 'unknown': 1}, 'bad') + assert error.error_detail == webrtc.RTCErrorDetailType.sdp_syntax_error + assert error.sdp_line_number == 3 + assert str(error) == 'bad' + with pytest.raises(TypeError, match='error_detail'): + webrtc.RTCError({}) @pytest.mark.asyncio -async def test_description_errors(pc): - """An answer without an offer is in the wrong state, and invalid SDP is an RTCError of its syntax""" +async def test_description_errors(pc: webrtc.RTCPeerConnection) -> None: + """An answer without an offer is in the wrong state, and invalid SDP is an RTCError of its syntax.""" with pytest.raises(webrtc.InvalidStateError): await pc.set_remote_description({'type': 'answer', 'sdp': 'invalid'}) with pytest.raises(webrtc.RTCError) as info: @@ -272,8 +298,8 @@ async def test_description_errors(pc): @pytest.mark.asyncio -async def test_created_descriptions(pc): - """Only the descriptions the connection created can be set as local ones, unmodified""" +async def test_created_descriptions(pc: webrtc.RTCPeerConnection) -> None: + """Only the descriptions the connection created can be set as local ones, unmodified.""" pc.add_transceiver(webrtc.MediaType.audio) offer = await pc.create_offer() assert isinstance(offer, webrtc.RTCSessionDescriptionInit) @@ -287,8 +313,10 @@ async def test_created_descriptions(pc): @pytest.mark.asyncio -async def test_provisional_answers_without_sdp(caller, callee): - """A provisional answer, and the final answer after it, are set without SDP""" +async def test_provisional_answers_without_sdp( + caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection +) -> None: + """A provisional answer, and the final answer after it, are set without SDP.""" caller.add_transceiver(webrtc.MediaType.audio) await exchange_offer(caller, callee) @@ -301,8 +329,8 @@ async def test_provisional_answers_without_sdp(caller, callee): assert callee.current_local_description.type == webrtc.RTCSdpType.answer -def test_closed_connection_keeps_its_transceivers(pc): - """The transceivers of a closed connection are stopped, but still there""" +def test_closed_connection_keeps_its_transceivers(pc: webrtc.RTCPeerConnection) -> None: + """The transceivers of a closed connection are stopped, but still there.""" pc.add_transceiver(webrtc.MediaType.audio) pc.close() [transceiver] = pc.get_transceivers() diff --git a/tests/test_data_channel.py b/tests/test_data_channel.py index 4be7527..4f1e106 100644 --- a/tests/test_data_channel.py +++ b/tests/test_data_channel.py @@ -7,6 +7,8 @@ """Data channels: opening, messages, closing, and their limits.""" +from __future__ import annotations + import asyncio import pytest @@ -15,11 +17,13 @@ from tests.helpers import connect, wait_for_event, wait_until -async def open_pair(caller, callee, **options): - """Opens a channel of the caller and returns it with its remote end""" - channel = caller.create_data_channel('chat', **options) +async def open_pair( + caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection, **options: object +) -> tuple[webrtc.RTCDataChannel, webrtc.RTCDataChannel]: + """Opens a channel of the caller and returns it with its remote end.""" + channel = caller.create_data_channel('chat', options) if options.get('negotiated'): - remote = callee.create_data_channel('chat', **options) + remote = callee.create_data_channel('chat', options) opened = [wait_for_event(channel, 'open'), wait_for_event(remote, 'open')] await connect(caller, callee) await asyncio.gather(*opened) @@ -34,13 +38,15 @@ async def open_pair(caller, callee, **options): @pytest.mark.asyncio -@pytest.mark.parametrize('negotiated', [False, True]) -async def test_messages_both_ways(caller, callee, negotiated): - """Both ends of an announced or negotiated channel have its options and id, and messages arrive in order""" - options = {'negotiated': True, 'id': 3} if negotiated else {} +@pytest.mark.parametrize('options', [{}, {'negotiated': True, 'id': 3}], ids=['announced', 'negotiated']) +async def test_messages_both_ways( + caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection, options: dict[str, object] +) -> None: + """Both ends of an announced or negotiated channel have its options and id, and messages arrive in order.""" channel, remote = await open_pair(caller, callee, protocol='proto', **options) - assert remote.label == 'chat' and remote.protocol == 'proto' + assert remote.label == 'chat' + assert remote.protocol == 'proto' assert channel.ready_state == remote.ready_state == webrtc.RTCDataChannelState.open assert channel.id == remote.id @@ -48,7 +54,7 @@ async def test_messages_both_ways(caller, callee, negotiated): got_all = asyncio.get_running_loop().create_future() @remote.on('message') - def on_message(event): + def on_message(event: webrtc.MessageEvent) -> None: received.append(event.data) if len(received) == 3 and not got_all.done(): got_all.set_result(None) @@ -61,8 +67,10 @@ def on_message(event): @pytest.mark.asyncio -async def test_buffered_amount_and_low_event(caller, callee): - """A message sent is buffered until it's handed to SCTP, then bufferedamountlow is fired""" +async def test_buffered_amount_and_low_event( + caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection +) -> None: + """A message sent is buffered until it's handed to SCTP, then bufferedamountlow is fired.""" channel, _ = await open_pair(caller, callee) channel.buffered_amount_low_threshold = 0 @@ -73,12 +81,12 @@ async def test_buffered_amount_and_low_event(caller, callee): @pytest.mark.asyncio -async def test_close_states_and_events(caller, callee): - """Closing a channel makes it closing at once; the remote end fires closing then close, the local end nothing""" +async def test_close_states_and_events(caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection) -> None: + """Closing a channel makes it closing at once; the remote end fires closing then close, the local end nothing.""" channel, remote = await open_pair(caller, callee) remote_events = [] - def on_remote_closing(event): + def on_remote_closing(_event: webrtc.Event) -> None: remote_events.append(('closing', remote.ready_state)) remote.on('closing', on_remote_closing) @@ -98,35 +106,35 @@ def on_remote_closing(event): @pytest.mark.parametrize( - 'init', + ('init', 'error'), [ - {'max_packet_life_time': 1, 'max_retransmits': 1}, - {'negotiated': True}, - {'negotiated': True, 'id': 65535}, + ({'max_packet_life_time': 1, 'max_retransmits': 1}, 'can not both be set'), + ({'negotiated': True}, 'needs an id'), + ({'negotiated': True, 'id': 65535}, 'id must be from 0 to 65534'), ], ids=['both limits', 'negotiated without id', 'id out of range'], ) -def test_invalid_data_channel_init(pc, init): - """A channel has at most one of the limits, and a negotiated one an id in range""" - with pytest.raises(ValueError): - pc.create_data_channel('x', **init) +def test_invalid_data_channel_init(pc: webrtc.RTCPeerConnection, init: dict[str, object], error: str) -> None: + """A channel has at most one of the limits, and a negotiated one an id in range.""" + with pytest.raises(ValueError, match=error): + pc.create_data_channel('x', init) -def test_id_is_ignored_unless_negotiated(pc): - """The id of a channel that isn't negotiated is chosen once SCTP is up""" - assert pc.create_data_channel('x', id=65535).id is None +def test_id_is_ignored_unless_negotiated(pc: webrtc.RTCPeerConnection) -> None: + """The id of a channel that isn't negotiated is chosen once SCTP is up.""" + assert pc.create_data_channel('x', {'id': 65535}).id is None -def test_id_taken(pc): - """Two negotiated channels can't have the same id""" - pc.create_data_channel('taken', negotiated=True, id=1) +def test_id_taken(pc: webrtc.RTCPeerConnection) -> None: + """Two negotiated channels can't have the same id.""" + pc.create_data_channel('taken', {'negotiated': True, 'id': 1}) with pytest.raises(webrtc.OperationError): - pc.create_data_channel('again', negotiated=True, id=1) + pc.create_data_channel('again', webrtc.RTCDataChannelInit(negotiated=True, id=1)) -def test_data_channel_options(pc): - """A new channel has its options and is connecting, and only sends text and bytes""" - channel = pc.create_data_channel('x', priority=webrtc.RTCPriorityType.high, ordered=False) +def test_data_channel_options(pc: webrtc.RTCPeerConnection) -> None: + """A new channel has its options and is connecting, and only sends text and bytes.""" + channel = pc.create_data_channel('x', {'priority': webrtc.RTCPriorityType.high, 'ordered': False}) assert channel.priority == webrtc.RTCPriorityType.high assert channel.ordered is False assert channel.ready_state == webrtc.RTCDataChannelState.connecting @@ -134,37 +142,39 @@ def test_data_channel_options(pc): channel.send(42) -def test_create_data_channel_on_closed_connection(pc): - """A closed connection can't create channels""" +def test_create_data_channel_on_closed_connection(pc: webrtc.RTCPeerConnection) -> None: + """A closed connection can't create channels.""" pc.close() with pytest.raises(webrtc.InvalidStateError): pc.create_data_channel('x') @pytest.mark.asyncio -async def test_max_message_size_before_an_answer(pc): - """The max message size is 65536 until an answer negotiates the max-message-size of the remote peer""" +async def test_max_message_size_before_an_answer(pc: webrtc.RTCPeerConnection) -> None: + """The max message size is 65536 until an answer negotiates the max-message-size of the remote peer.""" pc.create_data_channel('size') await pc.set_local_description() assert pc.sctp.max_message_size == 65536 @pytest.mark.asyncio -async def test_send_larger_than_max_message_size(caller, callee): - """A message larger than the negotiated max message size isn't sent, and the channel stays open""" +async def test_send_larger_than_max_message_size( + caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection +) -> None: + """A message larger than the negotiated max message size isn't sent, and the channel stays open.""" channel, _ = await open_pair(caller, callee) size = caller.sctp.max_message_size assert size > 65536 - with pytest.raises(ValueError): + with pytest.raises(ValueError, match='larger than the maxMessageSize'): channel.send(bytes(int(size) + 1)) assert channel.buffered_amount == 0 assert channel.ready_state == webrtc.RTCDataChannelState.open @pytest.mark.asyncio -async def test_stats_are_current(caller, callee): - """The stats of a channel count a message right after it's received""" +async def test_stats_are_current(caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection) -> None: + """The stats of a channel count a message right after it's received.""" channel, remote = await open_pair(caller, callee) before = (await callee.get_stats()).of_type('data-channel')[0].bytes_received received = wait_for_event(remote, 'message') @@ -176,8 +186,8 @@ async def test_stats_are_current(caller, callee): @pytest.mark.asyncio -async def test_max_channels_once_connected(caller, callee): - """The max number of channels is known once SCTP is connected""" +async def test_max_channels_once_connected(caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection) -> None: + """The max number of channels is known once SCTP is connected.""" caller.create_data_channel('channels') await caller.set_local_description() assert caller.sctp.max_channels is None @@ -188,8 +198,8 @@ async def test_max_channels_once_connected(caller, callee): @pytest.mark.asyncio -async def test_binary_type(caller, callee): - """Binary messages arrive as bytes, or as a Blob once binary_type is 'blob'; a Blob is sent as binary""" +async def test_binary_type(caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection) -> None: + """Binary messages arrive as bytes, or as a Blob once binary_type is 'blob'; a Blob is sent as binary.""" channel, remote = await open_pair(caller, callee) assert remote.binary_type == webrtc.BinaryType.arraybuffer == remote.binaryType @@ -203,23 +213,25 @@ async def test_binary_type(caller, callee): channel.send(webrtc.Blob([b'\x03', 'a', webrtc.Blob([b'\x04'])])) blob = (await second).data assert isinstance(blob, webrtc.Blob) - assert blob.size == 3 and await blob.array_buffer() == b'\x03a\x04' + assert blob.size == 3 + assert await blob.array_buffer() == b'\x03a\x04' text = wait_for_event(remote, 'message') channel.send('text') assert (await text).data == 'text' - with pytest.raises(ValueError): + with pytest.raises(ValueError, match='not a valid BinaryType'): remote.binary_type = 'buffer' assert remote.binary_type == webrtc.BinaryType.blob @pytest.mark.asyncio -async def test_blob(): - """A Blob is immutable bytes with a type, sliced like a sequence""" +async def test_blob() -> None: + """A Blob is immutable bytes with a type, sliced like a sequence.""" blob = webrtc.Blob(['héllo', b' ', bytearray(b'world')], type='Text/Plain') - assert blob.size == len(bytes(blob)) == 12 and blob.type == 'text/plain' + assert blob.size == len(bytes(blob)) == 12 + assert blob.type == 'text/plain' assert await blob.text() == 'héllo world' assert await blob.slice(-5).bytes() == b'world' assert await blob.slice(1, 3).array_buffer() == b'\xc3\xa9' - assert webrtc.Blob(type='é').type == '' + assert not webrtc.Blob(type='é').type diff --git a/tests/test_e2e.py b/tests/test_e2e.py index 2329526..99e80e5 100644 --- a/tests/test_e2e.py +++ b/tests/test_e2e.py @@ -5,6 +5,8 @@ # that can be found in the LICENSE.md file in the root of the project. # +from __future__ import annotations + import asyncio import pytest @@ -15,7 +17,9 @@ TIMEOUT = 20 -async def set_local_and_gather(pc, description): +async def set_local_and_gather( + pc: webrtc.RTCPeerConnection, description: webrtc.RTCSessionDescriptionInit +) -> webrtc.RTCSessionDescription | None: """Non-trickle ICE: returns the local description once all candidates are gathered.""" await pc.set_local_description(description) await wait_for_ice_gathering_complete(pc, TIMEOUT) @@ -23,7 +27,7 @@ async def set_local_and_gather(pc, description): @pytest.mark.asyncio -async def test_peers_connect_and_send_audio(caller, callee): +async def test_peers_connect_and_send_audio(caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection) -> None: """Two peer connections negotiate, connect over ICE/DTLS and stream audio.""" track = webrtc.MediaStreamTrackGenerator('audio') caller.add_track(track) @@ -48,10 +52,10 @@ async def test_peers_connect_and_send_audio(caller, callee): lambda pc=pc: pc.connection_state == webrtc.RTCPeerConnectionState.connected, 'connection', TIMEOUT ) assert pc.signaling_state == webrtc.RTCSignalingState.stable - assert pc.ice_connection_state in ( + assert pc.ice_connection_state in { webrtc.RTCIceConnectionState.connected, webrtc.RTCIceConnectionState.completed, - ) + } receivers = callee.get_receivers() assert len(receivers) == 1 diff --git a/tests/test_enums.py b/tests/test_enums.py index b492486..78338d1 100644 --- a/tests/test_enums.py +++ b/tests/test_enums.py @@ -7,6 +7,8 @@ """Enums: their members are their spec strings, in both directions of the native API.""" +from __future__ import annotations + import enum import pytest @@ -15,15 +17,15 @@ import webrtc.enums -def test_members_are_their_values(): - """Members equal their values, and print as them""" +def test_members_are_their_values() -> None: + """Members equal their values, and print as them.""" assert webrtc.RTCSignalingState.have_local_offer == 'have-local-offer' assert webrtc.RTCPriorityType('very-low') is webrtc.RTCPriorityType.very_low assert str(webrtc.MediaType.audio) == f'{webrtc.MediaType.audio}' == 'audio' -def test_enums_are_exported(): - """Every enum is exported from the package, and defined in one module""" +def test_enums_are_exported() -> None: + """Every enum is exported from the package, and defined in one module.""" enums = [ (name, value) for name, value in vars(webrtc.enums).items() @@ -32,19 +34,20 @@ def test_enums_are_exported(): assert enums for name, value in enums: if not name.startswith('_'): - assert getattr(webrtc, name) is value and name in webrtc.__all__ + assert getattr(webrtc, name) is value + assert name in webrtc.__all__ -def test_native_getters_return_members(pc): - """The native API returns members""" +def test_native_getters_return_members(pc: webrtc.RTCPeerConnection) -> None: + """The native API returns members.""" assert pc.signaling_state is webrtc.RTCSignalingState.stable transceiver = pc.add_transceiver('audio') assert transceiver.kind is webrtc.MediaType.audio assert transceiver.direction is webrtc.TransceiverDirection.sendrecv -def test_native_setters_take_members_and_values(pc): - """The native API takes members and their values""" +def test_native_setters_take_members_and_values(pc: webrtc.RTCPeerConnection) -> None: + """The native API takes members and their values.""" transceiver = pc.add_transceiver(webrtc.MediaType.audio) transceiver.direction = 'recvonly' assert transceiver.direction is webrtc.TransceiverDirection.recvonly @@ -52,8 +55,8 @@ def test_native_setters_take_members_and_values(pc): assert transceiver.direction == 'inactive' -def test_invalid_values_are_type_errors(pc): - """Like for a WebIDL enum, a value the enum doesn't have is a TypeError""" +def test_invalid_values_are_type_errors(pc: webrtc.RTCPeerConnection) -> None: + """Like for a WebIDL enum, a value the enum doesn't have is a TypeError.""" transceiver = pc.add_transceiver(webrtc.MediaType.audio) with pytest.raises(TypeError): transceiver.direction = 'nonsense' @@ -65,7 +68,7 @@ def test_invalid_values_are_type_errors(pc): webrtc.RTCPeerConnection(webrtc.RTCConfiguration(bundle_policy='nonsense')) -def test_data_channel_priority(pc): - """Every priority of a data channel round-trips through libwebrtc""" +def test_data_channel_priority(pc: webrtc.RTCPeerConnection) -> None: + """Every priority of a data channel round-trips through libwebrtc.""" for priority in webrtc.RTCPriorityType: - assert pc.create_data_channel('x', priority=priority.value).priority is priority + assert pc.create_data_channel('x', {'priority': priority.value}).priority is priority diff --git a/tests/test_events.py b/tests/test_events.py index e878fb5..7b43ae2 100644 --- a/tests/test_events.py +++ b/tests/test_events.py @@ -7,6 +7,8 @@ """Events: on/once/off, delivery on the event loop, and states that change along with their events.""" +from __future__ import annotations + import asyncio import pytest @@ -16,18 +18,18 @@ @pytest.mark.asyncio -async def test_on_decorator_once_and_off(pc): - """Handlers registered with the decorator and once are called, a handler removed with off isn't""" +async def test_on_decorator_once_and_off(pc: webrtc.RTCPeerConnection) -> None: + """Handlers registered with the decorator and once are called, a handler removed with off isn't.""" calls = [] @pc.on('negotiationneeded') - def decorated(event): + def decorated(event: webrtc.Event) -> None: calls.append(('decorated', event.type, event.target is pc)) - def once(event): + def once(event: webrtc.Event) -> None: calls.append(('once', event.type)) - def removed(event): + def removed(_event: webrtc.Event) -> None: calls.append(('removed',)) pc.once('negotiationneeded', once) @@ -43,12 +45,12 @@ def removed(event): @pytest.mark.asyncio -async def test_async_handlers_run_as_tasks(pc): - """A coroutine function handler is run as a task on the loop""" +async def test_async_handlers_run_as_tasks(pc: webrtc.RTCPeerConnection) -> None: + """A coroutine function handler is run as a task on the loop.""" done = asyncio.get_running_loop().create_future() @pc.on('negotiationneeded') - async def handler(event): + async def handler(event: webrtc.Event) -> None: await asyncio.sleep(0) if not done.done(): done.set_result(event.type) @@ -57,24 +59,24 @@ async def handler(event): assert await asyncio.wait_for(done, 5) == 'negotiationneeded' -def test_unknown_event(pc): - """Registering a handler of an event the object doesn't have is a ValueError""" - with pytest.raises(ValueError): - pc.on('nosuchevent', lambda event: None) +def test_unknown_event(pc: webrtc.RTCPeerConnection) -> None: + """Registering a handler of an event the object doesn't have is a ValueError.""" + with pytest.raises(ValueError, match="no event 'nosuchevent'"): + pc.on('nosuchevent', lambda _: None) -def test_handlers_need_a_running_loop(pc): - """Handlers are called on the loop they were registered from, so registering needs a running loop""" +def test_handlers_need_a_running_loop(pc: webrtc.RTCPeerConnection) -> None: + """Handlers are called on the loop they were registered from, so registering needs a running loop.""" with pytest.raises(RuntimeError): - pc.on('track', lambda event: None) + pc.on('track', lambda _: None) @pytest.mark.asyncio -async def test_signaling_state_changes_with_its_event(pc): - """Every signalingstatechange event sees its state, and comes before the operation that changed it resolves""" +async def test_signaling_state_changes_with_its_event(pc: webrtc.RTCPeerConnection) -> None: + """Every signalingstatechange event sees its state, and comes before the operation that changed it resolves.""" states = [] - def on_change(event): + def on_change(_event: webrtc.Event) -> None: states.append(pc.signaling_state) pc.on('signalingstatechange', on_change) @@ -87,11 +89,13 @@ def on_change(event): @pytest.mark.asyncio -async def test_closed_connection_emits_nothing(caller, callee): - """Closing a connection changes its states without emitting their events""" +async def test_closed_connection_emits_nothing( + caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection +) -> None: + """Closing a connection changes its states without emitting their events.""" events = [] - def on_event(event): + def on_event(event: webrtc.Event) -> None: events.append(event.type) for name in ('connectionstatechange', 'iceconnectionstatechange', 'signalingstatechange'): @@ -109,11 +113,11 @@ def on_event(event): @pytest.mark.asyncio -async def test_ice_candidates_and_end_of_candidates(pc): - """Candidates are parsed, each transport ends with an empty one, and the final None adds a=end-of-candidates""" +async def test_ice_candidates_and_end_of_candidates(pc: webrtc.RTCPeerConnection) -> None: + """Candidates are parsed, each transport ends with an empty one, and the final None adds a=end-of-candidates.""" candidates = [] - def on_candidate(event): + def on_candidate(event: webrtc.RTCPeerConnectionIceEvent) -> None: candidates.append(event.candidate) pc.on('icecandidate', on_candidate) @@ -123,25 +127,26 @@ def on_candidate(event): await gathered host = [c for c in candidates if c is not None and c.candidate] - assert host and all(c.type is not None for c in host), candidates - assert any(c is not None and c.candidate == '' for c in candidates), candidates + assert host, candidates + assert all(c.type is not None for c in host), candidates + assert any(c is not None and not c.candidate for c in candidates), candidates assert 'a=end-of-candidates' in pc.local_description.sdp, pc.local_description.sdp @pytest.mark.asyncio -async def test_descriptions_change_with_signaling_events(caller, callee): - """An offer received in have-local-offer rolls back first: each signalingstatechange sees its descriptions""" +async def test_descriptions_change_with_signaling_events( + caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection +) -> None: + """An offer received in have-local-offer rolls back first: each signalingstatechange sees its descriptions.""" await caller.set_local_description(await caller.create_offer()) seen = [] - def on_change(event): - seen.append( - { - 'state': caller.signaling_state, - 'local': caller.pending_local_description, - 'remote': caller.pending_remote_description, - } - ) + def on_change(_event: webrtc.Event) -> None: + seen.append({ + 'state': caller.signaling_state, + 'local': caller.pending_local_description, + 'remote': caller.pending_remote_description, + }) caller.on('signalingstatechange', on_change) await caller.set_remote_description(await callee.create_offer()) @@ -154,8 +159,8 @@ def on_change(event): @pytest.mark.asyncio -async def test_restart_ice_before_negotiation_needs_nothing(pc): - """restart_ice before the first negotiation doesn't fire negotiationneeded""" +async def test_restart_ice_before_negotiation_needs_nothing(pc: webrtc.RTCPeerConnection) -> None: + """restart_ice before the first negotiation doesn't fire negotiationneeded.""" events = [] pc.on('negotiationneeded', events.append) pc.restart_ice() @@ -163,11 +168,11 @@ async def test_restart_ice_before_negotiation_needs_nothing(pc): assert events == [] -def test_objects_used_from_another_loop_see_their_events(): - """Once its first loop is closed, an object is updated on the loop of its handlers""" +def test_objects_used_from_another_loop_see_their_events() -> None: + """Once its first loop is closed, an object is updated on the loop of its handlers.""" objects = {} - async def first(): + async def first() -> None: caller, callee = webrtc.RTCPeerConnection(), webrtc.RTCPeerConnection() channel = caller.create_data_channel('loops') opened = wait_for_event(channel, 'open') @@ -175,7 +180,7 @@ async def first(): await opened objects.update(caller=caller, callee=callee, channel=channel) - async def second(): + async def second() -> None: channel = objects['channel'] closed = wait_for_event(channel, 'close') objects['callee'].close() diff --git a/tests/test_ice_transport.py b/tests/test_ice_transport.py index 2460816..c7c177f 100644 --- a/tests/test_ice_transport.py +++ b/tests/test_ice_transport.py @@ -5,8 +5,12 @@ # that can be found in the LICENSE.md file in the root of the project. # -"""ICE transports: the ones of a connection, and ones of their own (the WebRTC ICE extension), which gather, start -and connect without a connection.""" +"""ICE transports: the ones of a connection, and ones of their own. + +The ones of their own (the WebRTC ICE extension) gather, start and connect without a connection. +""" + +from __future__ import annotations import asyncio @@ -17,8 +21,10 @@ @pytest.mark.asyncio -async def test_candidates_parameters_and_role(caller, callee): - """A connection's transport learns its role from the answer, and has the signaled parameters and candidates""" +async def test_candidates_parameters_and_role( + caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection +) -> None: + """A connection's transport learns its role from the answer, and has the signaled parameters and candidates.""" caller.create_data_channel('ice') await caller.set_local_description() ice = caller.sctp.transport.ice_transport @@ -29,25 +35,28 @@ async def test_candidates_parameters_and_role(caller, callee): await wait_until(lambda: ice.role == webrtc.RTCIceRole.controlling, 'the controlling role') remote_ice = callee.sctp.transport.ice_transport local, remote = ice.get_local_parameters(), ice.get_remote_parameters() - assert isinstance(local, webrtc.RTCIceParameters) and local.username_fragment and local.password + assert isinstance(local, webrtc.RTCIceParameters) + assert local.username_fragment + assert local.password assert remote.username_fragment == remote_ice.get_local_parameters().username_fragment - assert ice.get_local_candidates() and all(c.candidate for c in ice.get_local_candidates()) + assert ice.get_local_candidates() + assert all(c.candidate for c in ice.get_local_candidates()) assert {c.candidate for c in ice.get_remote_candidates()} <= { c.candidate for c in remote_ice.get_local_candidates() } @pytest.mark.asyncio -async def test_component(caller, callee): - """RTP and RTCP are multiplexed on the transport of the RTP component""" +async def test_component(caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection) -> None: + """RTP and RTCP are multiplexed on the transport of the RTP component.""" transceiver = caller.add_transceiver(webrtc.MediaType.audio) await connect(caller, callee) assert transceiver.sender.transport.ice_transport.component == webrtc.RTCIceComponent.rtp @pytest.mark.asyncio -async def test_close_keeps_gathering_state(pc): - """Closing a connection closes its transports, it doesn't complete their gathering""" +async def test_close_keeps_gathering_state(pc: webrtc.RTCPeerConnection) -> None: + """Closing a connection closes its transports, it doesn't complete their gathering.""" transceiver = pc.add_transceiver(webrtc.MediaType.audio) await pc.set_local_description() ice = transceiver.sender.transport.ice_transport @@ -58,15 +67,17 @@ async def test_close_keeps_gathering_state(pc): @pytest.mark.asyncio -async def test_two_transports_connect(): - """Two transports gather, start with each other's parameters and connect, one switching to the controlled role""" +async def test_two_transports_connect() -> None: + """Two transports gather, start with each other's parameters and connect, one switching to the controlled role.""" local, remote = webrtc.RTCIceTransport(), webrtc.RTCIceTransport() - assert local.role is None and local.state == webrtc.RTCIceTransportState.new - assert local.get_local_parameters() and local.get_remote_parameters() is None + assert local.role is None + assert local.state == webrtc.RTCIceTransportState.new + assert local.get_local_parameters() + assert local.get_remote_parameters() is None for transport, other in ((local, remote), (remote, local)): - def on_candidate(event, other=other): + def on_candidate(event: webrtc.RTCPeerConnectionIceEvent, other: webrtc.RTCIceTransport = other) -> None: if event.candidate: other.add_remote_candidate(event.candidate) @@ -94,8 +105,8 @@ def on_candidate(event, other=other): remote.stop() -def test_start_validation(): - """start checks the remote parameters and the role, and other remote parameters restart the checks""" +def test_start_validation() -> None: + """Start checks the remote parameters and the role, and other remote parameters restart the checks.""" transport = webrtc.RTCIceTransport() with pytest.raises(webrtc.InvalidSyntaxError): transport.start(webrtc.RTCIceParameters('ab', 'p' * 22)) diff --git a/tests/test_lifetime.py b/tests/test_lifetime.py index 76f3d54..970520d 100644 --- a/tests/test_lifetime.py +++ b/tests/test_lifetime.py @@ -7,15 +7,19 @@ """Ownership and lifetime of the native wrappers: no leaks, no use-after-free, stable identity and state.""" +from __future__ import annotations + import asyncio import gc -import os +import pathlib import subprocess import sys import textwrap import threading import time import weakref +from concurrent.futures import ThreadPoolExecutor +from typing import TYPE_CHECKING import pytest @@ -23,15 +27,18 @@ import wrtc from tests.helpers import QUIET_PERIOD, connect, exchange_offer_answer, wait_for_event +if TYPE_CHECKING: + from collections.abc import Callable, Iterator + -def collect(): +def collect() -> None: gc.collect() gc.collect() -def alive_factories(): - """Factories constructed and not destroyed yet. The last reference to one may be released - on a helper thread (see CreateSessionDescriptionObserver), so wait for the count to settle""" +def alive_factories() -> int: + """Factories constructed and not destroyed yet, once the count settles.""" + # the last reference to one may be released on a helper thread (see CreateSessionDescriptionObserver) collect() count = wrtc._alive_factories() deadline = time.monotonic() + 2 @@ -44,7 +51,7 @@ def alive_factories(): return count -def peer_connection_cycle(): +def peer_connection_cycle() -> None: pc = webrtc.RTCPeerConnection() pc.add_transceiver(webrtc.MediaType.audio) pc.add_transceiver(webrtc.MediaType.video) @@ -55,8 +62,8 @@ def peer_connection_cycle(): @pytest.fixture -def isolated(): - """Nothing created by earlier tests may keep the default factory alive""" +def isolated() -> Iterator[None]: + """Nothing created by earlier tests may keep the default factory alive.""" collect() yield collect() @@ -65,8 +72,8 @@ def isolated(): pytestmark = pytest.mark.usefixtures('isolated') -def test_peer_connection_cycles_do_not_leak_factories(): - """Closed and collected connections release their factory""" +def test_peer_connection_cycles_do_not_leak_factories() -> None: + """Closed and collected connections release their factory.""" baseline = alive_factories() for _ in range(30): @@ -75,8 +82,8 @@ def test_peer_connection_cycles_do_not_leak_factories(): assert alive_factories() == baseline -def test_factories_return_to_baseline_when_everything_is_gone(): - """Everything shares one factory, which is gone with the last object using it""" +def test_factories_return_to_baseline_when_everything_is_gone() -> None: + """Everything shares one factory, which is gone with the last object using it.""" baseline = alive_factories() pc = webrtc.RTCPeerConnection() @@ -95,8 +102,8 @@ def test_factories_return_to_baseline_when_everything_is_gone(): assert alive_factories() == baseline -def test_everything_alive_shares_one_factory(): - """New connections use the factory of the media alive""" +def test_everything_alive_shares_one_factory() -> None: + """New connections use the factory of the media alive.""" stream = webrtc.get_user_media() source_track = webrtc.MediaStreamTrackGenerator('audio') before = alive_factories() @@ -110,8 +117,8 @@ def test_everything_alive_shares_one_factory(): pc.close() -def test_closed_connection_keeps_its_factory_shared(): - """A closed connection and its tracks keep their factory, which new connections reuse""" +def test_closed_connection_keeps_its_factory_shared() -> None: + """A closed connection and its tracks keep their factory, which new connections reuse.""" pc = webrtc.RTCPeerConnection() receiver_track = pc.add_transceiver(webrtc.MediaType.audio).receiver.track pc.close() @@ -125,8 +132,8 @@ def test_closed_connection_keeps_its_factory_shared(): assert receiver_track.ready_state == webrtc.MediaStreamTrackState.ended -def test_dropped_track_wrappers_are_not_notified(): - """Toggling a track notifies its observers, a dropped wrapper must not be one of them""" +def test_dropped_track_wrappers_are_not_notified() -> None: + """Toggling a track notifies its observers, a dropped wrapper must not be one of them.""" streams = [webrtc.get_user_media() for _ in range(50)] tracks = [stream.get_tracks()[0] for stream in streams] del tracks @@ -137,8 +144,8 @@ def test_dropped_track_wrappers_are_not_notified(): stream.get_tracks()[0].enabled = bool(i % 2) -def test_destroyed_track_wrapper_is_not_notified(): - """A track wrapper dies while libwebrtc keeps the track alive in a sender""" +def test_destroyed_track_wrapper_is_not_notified() -> None: + """A track wrapper dies while libwebrtc keeps the track alive in a sender.""" pc = webrtc.RTCPeerConnection() stream = webrtc.get_user_media() pc.add_track(stream.get_tracks()[0]) @@ -152,8 +159,8 @@ def test_destroyed_track_wrapper_is_not_notified(): pc.close() -def test_track_state_survives_gc(): - """The state of a track is kept by its native object, not by its wrapper""" +def test_track_state_survives_gc() -> None: + """The state of a track is kept by its native object, not by its wrapper.""" stream = webrtc.get_user_media() track = stream.get_tracks()[0] track.enabled = False @@ -166,8 +173,8 @@ def test_track_state_survives_gc(): assert track.enabled is False -def test_track_state_survives_gc_of_all_python_references(): - """The state of a remote track survives when no wrapper of it is left""" +def test_track_state_survives_gc_of_all_python_references() -> None: + """The state of a remote track survives when no wrapper of it is left.""" pc = webrtc.RTCPeerConnection() pc.add_transceiver(webrtc.MediaType.audio) pc.get_transceivers()[0].receiver.track.stop() @@ -177,34 +184,33 @@ def test_track_state_survives_gc_of_all_python_references(): pc.close() -def test_transceiver_identity(): - """The same transceiver, sender, receiver and track are wrapped by the same native wrappers""" +def test_transceiver_identity() -> None: + """The same transceiver, sender, receiver and track are wrapped by the same native wrappers.""" pc = webrtc.RTCPeerConnection() transceiver = pc.add_transceiver(webrtc.MediaType.audio) collect() assert pc.get_transceivers()[0] == transceiver - assert pc.get_transceivers()[0]._native_obj is transceiver._native_obj - assert pc.get_senders()[0]._native_obj is transceiver.sender._native_obj - assert pc.get_receivers()[0]._native_obj is transceiver.receiver._native_obj - assert transceiver.receiver.track._native_obj is transceiver.receiver.track._native_obj + assert pc.get_senders()[0] == transceiver.sender + assert pc.get_receivers()[0] == transceiver.receiver + assert transceiver.receiver.track == transceiver.receiver.track pc.close() -def test_sender_identity(audio_stream): - """A sender and its track are the same objects after a collection""" +def test_sender_identity(audio_stream: webrtc.MediaStream) -> None: + """A sender and its track are the same objects after a collection.""" pc = webrtc.RTCPeerConnection() track = audio_stream.get_tracks()[0] sender = pc.add_track(track) collect() assert pc.get_senders()[0] == sender - assert sender.track._native_obj is track._native_obj + assert sender.track == track pc.close() -def test_drop_and_refetch_before_negotiation(): - """Wrappers dropped and fetched again, many times, stay valid""" +def test_drop_and_refetch_before_negotiation() -> None: + """Wrappers dropped and fetched again, many times, stay valid.""" pc = webrtc.RTCPeerConnection() pc.add_transceiver(webrtc.MediaType.audio) pc.add_transceiver(webrtc.MediaType.video) @@ -224,8 +230,10 @@ def test_drop_and_refetch_before_negotiation(): @pytest.mark.asyncio -async def test_drop_and_refetch_transports(caller, callee, audio_stream): - """Transports dropped and fetched again, many times, stay valid and shared""" +async def test_drop_and_refetch_transports( + caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection, audio_stream: webrtc.MediaStream +) -> None: + """Transports dropped and fetched again, many times, stay valid and shared.""" caller.add_track(audio_stream.get_tracks()[0]) await exchange_offer_answer(caller, callee) @@ -236,11 +244,8 @@ async def test_drop_and_refetch_transports(caller, callee, audio_stream): for _ in range(50): for pc in (caller, callee): - for transceiver in pc.get_transceivers(): - for transport in (transceiver.sender.transport, transceiver.receiver.transport): - if transport is not None: - _ = transport.state, transport.ice_transport.state, transport.ice_transport.gathering_state - del pc, transceiver, transport + read_transports(pc) + del pc collect() assert caller.get_senders()[0].transport == caller.get_transceivers()[0].receiver.transport @@ -248,8 +253,10 @@ async def test_drop_and_refetch_transports(caller, callee, audio_stream): @pytest.mark.asyncio -async def test_transports_outlive_closed_connection(caller, callee, audio_stream): - """Transports kept after their connection is closed and collected report closed""" +async def test_transports_outlive_closed_connection( + caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection, audio_stream: webrtc.MediaStream +) -> None: + """Transports kept after their connection is closed and collected report closed.""" caller.add_track(audio_stream.get_tracks()[0]) await exchange_offer_answer(caller, callee) @@ -263,8 +270,26 @@ async def test_transports_outlive_closed_connection(caller, callee, audio_stream assert ice_transport.state == webrtc.RTCIceTransportState.closed -def test_getters_while_connection_is_closed_and_dropped(audio_stream): - """Reading from other threads while the connection is closed and dropped neither crashes nor deadlocks""" +def read_transports(pc: webrtc.RTCPeerConnection) -> None: + for transceiver in pc.get_transceivers(): + for transport in (transceiver.sender.transport, transceiver.receiver.transport): + if transport is not None: + _ = transport.state, transport.ice_transport.state, transport.ice_transport.gathering_state + + +def read_until_stopped(pc: webrtc.RTCPeerConnection, stop: threading.Event) -> None: + while not stop.is_set(): + for transceiver in pc.get_transceivers(): + _ = transceiver.direction, transceiver.receiver.track.ready_state, transceiver.sender.track + for sender in pc.get_senders(): + _ = sender.track, sender.transport + for receiver in pc.get_receivers(): + _ = receiver.track.enabled, receiver.transport + _ = pc.connection_state, pc.signaling_state + + +def test_getters_while_connection_is_closed_and_dropped(audio_stream: webrtc.MediaStream) -> None: + """Reading from other threads while the connection is closed and dropped neither crashes nor deadlocks.""" track = audio_stream.get_tracks()[0] for _ in range(20): @@ -272,48 +297,32 @@ def test_getters_while_connection_is_closed_and_dropped(audio_stream): pc.add_track(track) pc.add_transceiver(webrtc.MediaType.video) stop = threading.Event() - errors = [] - - def read(pc=pc, stop=stop, errors=errors): - try: - while not stop.is_set(): - for transceiver in pc.get_transceivers(): - _ = transceiver.direction, transceiver.receiver.track.ready_state, transceiver.sender.track - for sender in pc.get_senders(): - _ = sender.track, sender.transport - for receiver in pc.get_receivers(): - _ = receiver.track.enabled, receiver.transport - _ = pc.connection_state, pc.signaling_state - except Exception as e: # noqa: BLE001 - errors.append(e) - - readers = [threading.Thread(target=read) for _ in range(4)] - for reader in readers: - reader.start() + executor = ThreadPoolExecutor(max_workers=4) + readers = [executor.submit(read_until_stopped, pc, stop) for _ in range(4)] pc.close() del pc collect() stop.set() + # raises what a reader raised, or TimeoutError if it's stuck for reader in readers: - reader.join(timeout=10) - assert not reader.is_alive(), 'reader thread is stuck' - assert not errors, errors + reader.result(timeout=10) + executor.shutdown() collect() @pytest.mark.asyncio -async def test_handlers_referencing_their_connection_do_not_keep_it_alive(): - """The garbage collector sees the handlers while Python alone owns the connection""" +async def test_handlers_referencing_their_connection_do_not_keep_it_alive() -> None: + """The garbage collector sees the handlers while Python alone owns the connection.""" baseline = alive_factories() delivered = asyncio.get_running_loop().create_future() - def create(): + def create() -> weakref.ref[webrtc.RTCPeerConnection]: pc = webrtc.RTCPeerConnection() @pc.on('negotiationneeded') - def on_negotiation(event): + def on_negotiation(_event: webrtc.Event) -> None: pc.get_transceivers() if not delivered.done(): delivered.set_result(None) @@ -330,24 +339,24 @@ def on_negotiation(event): @pytest.mark.asyncio -async def test_channel_handlers_do_not_dangle_after_close_and_gc(): - """Handlers of a channel collected with echoed messages in flight are never called again, and nothing crashes""" +async def test_channel_handlers_do_not_dangle_after_close_and_gc() -> None: + """Handlers of a channel collected with echoed messages in flight are never called again, and nothing crashes.""" caller, callee = webrtc.RTCPeerConnection(), webrtc.RTCPeerConnection() channel = caller.create_data_channel('lifetime') received = [] - def on_close(event): + def on_close(_event: webrtc.Event) -> None: received.append('close') channel.on('message', lambda event: received.append(event.data)) channel.on('close', on_close) @callee.on('datachannel') - def on_channel(event): + def on_channel(event: webrtc.RTCDataChannelEvent) -> None: remote = event.channel @remote.on('message') - def echo(message): + def echo(message: webrtc.MessageEvent) -> None: remote.send(message.data) opened = wait_for_event(channel, 'open') @@ -365,12 +374,12 @@ def echo(message): @pytest.mark.asyncio -async def test_connections_dropped_while_emitting_events(): - """Events of connections that are destroyed meanwhile are dropped safely: a crash test, for sanitizers""" +async def test_connections_dropped_while_emitting_events() -> None: + """Events of connections that are destroyed meanwhile are dropped safely: a crash test, for sanitizers.""" for _ in range(20): pc = webrtc.RTCPeerConnection() - pc.on('icecandidate', lambda event: None) - pc.on('icegatheringstatechange', lambda event: None) + pc.on('icecandidate', lambda _: None) + pc.on('icegatheringstatechange', lambda _: None) pc.create_data_channel('gather') await pc.set_local_description() del pc @@ -380,10 +389,10 @@ async def test_connections_dropped_while_emitting_events(): collect() -def test_process_exits_after_connecting(): - """Wrappers destroyed at exit unregister while libwebrtc threads may be wrapping objects, without deadlock""" +def test_process_exits_after_connecting() -> None: + """Wrappers destroyed at exit unregister while libwebrtc threads may be wrapping objects, without deadlock.""" script = textwrap.dedent( - ''' + """ import asyncio import webrtc from tests.helpers import connect, wait_for_event @@ -401,18 +410,22 @@ async def main(): print('connected') asyncio.run(main()) - ''' + """ ) # from the root of the project, which has the tests package (pytest may run from elsewhere, like cibuildwheel) - root = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) - result = subprocess.run([sys.executable, '-c', script], capture_output=True, text=True, timeout=60, cwd=root) + root = pathlib.Path(pathlib.Path(pathlib.Path(__file__).resolve()).parent).parent + result = subprocess.run( + [sys.executable, '-c', script], capture_output=True, text=True, timeout=60, cwd=root, check=False + ) assert 'connected' in result.stdout, result.stderr[-2000:] assert result.returncode == 0, result.stderr[-2000:] @pytest.mark.asyncio -async def test_stream_tracks_read_while_they_change(caller, callee, audio_stream): - """Reading a remote stream's tracks while a description changes them doesn't deadlock with its observer""" +async def test_stream_tracks_read_while_they_change( + caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection, audio_stream: webrtc.MediaStream +) -> None: + """Reading a remote stream's tracks while a description changes them doesn't deadlock with its observer.""" audio = audio_stream.get_audio_tracks()[0] sender = caller.add_track(audio, audio_stream) track_event = wait_for_event(callee, 'track') @@ -421,7 +434,7 @@ async def test_stream_tracks_read_while_they_change(caller, callee, audio_stream stop = threading.Event() - def read(): + def read() -> None: while not stop.is_set(): remote.get_tracks() @@ -440,10 +453,10 @@ def read(): assert not reader.is_alive() -def test_collected_on_libwebrtc_thread(): - """The garbage collector running on a libwebrtc thread (to emit an event) releases a connection elsewhere""" +def test_collected_on_libwebrtc_thread() -> None: + """The garbage collector running on a libwebrtc thread (to emit an event) releases a connection elsewhere.""" script = textwrap.dedent( - ''' + """ import asyncio import gc import threading @@ -474,34 +487,37 @@ def collecting(self, name, *args): print('collected') asyncio.run(main()) - ''' + """ + ) + root = pathlib.Path(pathlib.Path(pathlib.Path(__file__).resolve()).parent).parent + result = subprocess.run( + [sys.executable, '-c', script], capture_output=True, text=True, timeout=60, cwd=root, check=False ) - root = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) - result = subprocess.run([sys.executable, '-c', script], capture_output=True, text=True, timeout=60, cwd=root) assert 'collected' in result.stdout, result.stderr[-2000:] @pytest.mark.asyncio -async def test_generators_are_collected(): - """An audio generator is its own track: its handlers mustn't keep it alive""" +async def test_generators_are_collected() -> None: + """An audio generator is its own track: its handlers mustn't keep it alive.""" baseline = alive_factories() - def create(): + def create() -> tuple[weakref.ref[object], ...]: audio = webrtc.MediaStreamTrackGenerator('audio') video = webrtc.VideoTrackGenerator() - audio.on('ended', lambda event: audio.kind) - video.track.on('ended', lambda event: video.track) + audio.on('ended', lambda _: audio.kind) + video.track.on('ended', lambda _: video.track) return weakref.ref(audio), weakref.ref(video), weakref.ref(video.track) refs = [ref for _ in range(10) for ref in create()] + await asyncio.sleep(QUIET_PERIOD) collect() assert [ref for ref in refs if ref() is not None] == [] assert alive_factories() == baseline -def test_generator_track_stays_ended_without_its_wrapper(): - """A generator whose stopped track is collected keeps dropping what's written""" +def test_generator_track_stays_ended_without_its_wrapper() -> None: + """A generator whose stopped track is collected keeps dropping what's written.""" generator = wrtc.TrackGenerator('video') track = generator.track assert generator.live @@ -513,24 +529,24 @@ def test_generator_track_stays_ended_without_its_wrapper(): assert webrtc.MediaStreamTrack._wrap(generator.track).ready_state == webrtc.MediaStreamTrackState.ended -def processor_with_handler_on_its_track(): +def processor_with_handler_on_its_track() -> webrtc.MediaStreamTrackProcessor: track = webrtc.get_user_media(audio=False, video=True).get_tracks()[0] processor = webrtc.MediaStreamTrackProcessor(track) - track.on('ended', lambda event: processor.readable) + track.on('ended', lambda _: processor.readable) track.stop() return processor -def stream_with_handler_on_its_track(): +def stream_with_handler_on_its_track() -> webrtc.MediaStream: stream = webrtc.get_user_media(audio=True, video=False) - stream.get_tracks()[0].on('ended', lambda event: stream.id) + stream.get_tracks()[0].on('ended', lambda _: stream.id) return stream -def processor_of_generator_with_handler(): +def processor_of_generator_with_handler() -> webrtc.MediaStreamTrackProcessor: generator = webrtc.MediaStreamTrackGenerator('video') processor = webrtc.MediaStreamTrackProcessor(generator) - generator.on('ended', lambda event: processor.readable) + generator.on('ended', lambda _: processor.readable) generator.stop() return processor @@ -540,8 +556,8 @@ def processor_of_generator_with_handler(): 'create', [processor_with_handler_on_its_track, stream_with_handler_on_its_track, processor_of_generator_with_handler], ) -async def test_handlers_of_owned_tracks_do_not_keep_owners_alive(create): - """Handlers of a track referencing its processor or stream don't keep them alive""" +async def test_handlers_of_owned_tracks_do_not_keep_owners_alive(create: Callable[[], object]) -> None: + """Handlers of a track referencing its processor or stream don't keep them alive.""" baseline = alive_factories() refs = [weakref.ref(create()) for _ in range(5)] await asyncio.sleep(QUIET_PERIOD) @@ -551,8 +567,8 @@ async def test_handlers_of_owned_tracks_do_not_keep_owners_alive(create): assert alive_factories() == baseline -def test_stream_keeps_the_state_of_its_tracks(): - """The native stream keeps its tracks weakly, the Python one keeps them: a stopped track stays stopped""" +def test_stream_keeps_the_state_of_its_tracks() -> None: + """The native stream keeps its tracks weakly, the Python one keeps them: a stopped track stays stopped.""" stream = webrtc.MediaStream(webrtc.get_user_media(audio=True, video=True).get_tracks()) for track in stream.get_tracks(): track.stop() @@ -570,16 +586,16 @@ def test_stream_keeps_the_state_of_its_tracks(): ) @pytest.mark.asyncio @pytest.mark.parametrize('part', ['sender', 'receiver']) -async def test_handler_of_a_track_referencing_its_sender_or_receiver(part): +async def test_handler_of_a_track_referencing_its_sender_or_receiver(part: str) -> None: baseline = alive_factories() - def create(): + def create() -> weakref.ref[webrtc.RTCRtpSender | webrtc.RTCRtpReceiver]: pc = webrtc.RTCPeerConnection() if part == 'sender': owner = pc.add_track(webrtc.get_user_media(audio=True, video=False).get_tracks()[0]) else: owner = pc.add_transceiver(webrtc.MediaType.audio).receiver - owner.track.on('ended', lambda event: owner.track) + owner.track.on('ended', lambda _: owner.track) pc.close() return weakref.ref(owner) @@ -591,8 +607,8 @@ def create(): assert alive_factories() == baseline -def alive_objects(): - """The native objects alive by type, once releases on helper threads are done""" +def alive_objects() -> dict[str, int]: + """The native objects alive by type, once releases on helper threads are done.""" collect() alive = wrtc._alive() deadline = time.monotonic() + 2 @@ -607,11 +623,11 @@ def alive_objects(): @pytest.mark.asyncio -async def test_a_session_releases_every_native_object(): - """Connections with media, channels, processors and generators, closed and dropped: nothing native is left""" +async def test_a_session_releases_every_native_object() -> None: + """Connections with media, channels, processors and generators, closed and dropped: nothing native is left.""" baseline = alive_objects() - async def session(): + async def session() -> None: caller, callee = webrtc.RTCPeerConnection(), webrtc.RTCPeerConnection() stream = webrtc.get_user_media(audio=True, video=True) for track in stream.get_tracks(): diff --git a/tests/test_media_e2e.py b/tests/test_media_e2e.py index 090f4cb..9178698 100644 --- a/tests/test_media_e2e.py +++ b/tests/test_media_e2e.py @@ -7,6 +7,8 @@ """Media through a connection: generated on one end, read with a processor on the other.""" +from __future__ import annotations + import array import asyncio import math @@ -20,13 +22,13 @@ WIDTH, HEIGHT = 320, 240 -def solid_i420(y, u, v): +def solid_i420(y: int, u: int, v: int) -> bytes: chroma = (WIDTH // 2) * (HEIGHT // 2) return bytes([y] * (WIDTH * HEIGHT) + [u] * chroma + [v] * chroma) -async def write_sine(generator, frequency, stop): - """Writes a sine of 48 kHz mono in 10 ms frames, at the pace of real time""" +async def write_sine(generator: webrtc.MediaStreamTrackGenerator, frequency: float, *, stop: asyncio.Event) -> None: + """Writes a sine of 48 kHz mono in 10 ms frames, at the pace of real time.""" writer = generator.writable.get_writer() loop = asyncio.get_running_loop() start = loop.time() @@ -48,45 +50,54 @@ async def write_sine(generator, frequency, stop): await asyncio.sleep(max(0.0, start + written / 48000 - loop.time())) -def dominant_frequency(samples, rate): - """The frequency of a sine, from its rising zero crossings""" +def dominant_frequency(samples: list[float], rate: int) -> float: + """The frequency of a sine, from its rising zero crossings.""" crossings = [i for i in range(1, len(samples)) if samples[i - 1] < 0 <= samples[i]] return (len(crossings) - 1) * rate / (crossings[-1] - crossings[0]) +async def read_frames(reader: webrtc.ReadableStreamDefaultReader, count: int) -> tuple[list[int], bytearray]: + """Reads frames of the size, returns their timestamps and the last one in RGBA.""" + timestamps = [] + for _ in range(count): + frame = (await asyncio.wait_for(reader.read(), TIMEOUT)).value + assert (frame.coded_width, frame.coded_height) == (WIDTH, HEIGHT) + assert frame.metadata().rtp_timestamp > 0 + timestamps.append(frame.timestamp) + rgba = bytearray(frame.allocation_size({'format': 'RGBA'})) + await frame.copy_to(rgba, {'format': 'RGBA'}) + frame.close() + return timestamps, rgba + + @pytest.mark.asyncio -async def test_video_through_a_connection(caller, callee): - """Frames of one color arrive in that color, at their size, with increasing timestamps""" +async def test_video_through_a_connection(caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection) -> None: + """Frames of one color arrive in that color, at their size, with increasing timestamps.""" generator = webrtc.VideoTrackGenerator() # pure red in BT.601 limited range - async with writing(write_video, generator, solid_i420(81, 90, 240), WIDTH, HEIGHT): - remote = await connect_track(caller, callee, generator.track, TIMEOUT) + async with writing(write_video, generator, solid_i420(81, 90, 240), (WIDTH, HEIGHT)): + remote = await connect_track(caller, callee, generator.track, timeout=TIMEOUT) reader = webrtc.MediaStreamTrackProcessor(remote, max_buffer_size=5).readable.get_reader() - timestamps = [] - for _ in range(10): - frame = (await asyncio.wait_for(reader.read(), TIMEOUT)).value - assert (frame.coded_width, frame.coded_height) == (WIDTH, HEIGHT) - assert frame.metadata().rtp_timestamp > 0 - timestamps.append(frame.timestamp) - rgba = bytearray(frame.allocation_size({'format': 'RGBA'})) - await frame.copy_to(rgba, {'format': 'RGBA'}) - frame.close() + timestamps, rgba = await read_frames(reader, 10) await reader.cancel() generator.track.stop() center = (HEIGHT // 2 * WIDTH + WIDTH // 2) * 4 r, g, b, a = rgba[center : center + 4] # encoding adds some noise - assert r > 230 and g < 25 and b < 25 and a == 255 + assert r > 230 + assert g < 25 + assert b < 25 + assert a == 255 assert timestamps == sorted(set(timestamps)) @pytest.mark.asyncio -async def test_audio_through_a_connection(caller, callee): - """A sine arrives as a sine of the same frequency, decoded at 48 kHz""" +async def test_audio_through_a_connection(caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection) -> None: + """A sine arrives as a sine of the same frequency, decoded at 48 kHz.""" generator = webrtc.MediaStreamTrackGenerator('audio') async with writing(write_sine, generator, 440): - remote = await connect_track(caller, callee, generator, TIMEOUT) + remote = await connect_track(caller, callee, generator, timeout=TIMEOUT) reader = webrtc.MediaStreamTrackProcessor(remote, max_buffer_size=100).readable.get_reader() samples = [] for chunk in range(150): @@ -106,11 +117,13 @@ async def test_audio_through_a_connection(caller, callee): @pytest.mark.asyncio -async def test_remote_track_end_closes_the_processor(caller, callee): - """The processor of a remote track closes when the remote peer stops sending it""" +async def test_remote_track_end_closes_the_processor( + caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection +) -> None: + """The processor of a remote track closes when the remote peer stops sending it.""" generator = webrtc.VideoTrackGenerator() - async with writing(write_video, generator, solid_i420(128, 128, 128), WIDTH, HEIGHT): - remote = await connect_track(caller, callee, generator.track, TIMEOUT) + async with writing(write_video, generator, solid_i420(128, 128, 128), (WIDTH, HEIGHT)): + remote = await connect_track(caller, callee, generator.track, timeout=TIMEOUT) reader = webrtc.MediaStreamTrackProcessor(remote).readable.get_reader() (await asyncio.wait_for(reader.read(), TIMEOUT)).value.close() callee.close() diff --git a/tests/test_media_stream_track_processor.py b/tests/test_media_stream_track_processor.py index 5654361..804f538 100644 --- a/tests/test_media_stream_track_processor.py +++ b/tests/test_media_stream_track_processor.py @@ -7,6 +7,8 @@ """MediaStreamTrackProcessor, VideoTrackGenerator and MediaStreamTrackGenerator on local tracks.""" +from __future__ import annotations + import array import asyncio @@ -18,24 +20,24 @@ TIMEOUT = 10 -def i420(width, height, y=81, u=90, v=240): +def i420(width: int, height: int) -> bytes: chroma = ((width + 1) // 2) * ((height + 1) // 2) - return bytes([y] * (width * height) + [u] * chroma + [v] * chroma) + return bytes([81] * (width * height) + [90] * chroma + [240] * chroma) -def video_frame(timestamp, width=4, height=2): +def video_frame(timestamp: int, width: int = 4, height: int = 2) -> webrtc.VideoFrame: return webrtc.VideoFrame( i420(width, height), format='I420', coded_width=width, coded_height=height, timestamp=timestamp ) -async def read(reader): +async def read(reader: webrtc.ReadableStreamDefaultReader) -> webrtc.ReadableStreamReadResult: return await asyncio.wait_for(reader.read(), TIMEOUT) @pytest.mark.asyncio -async def test_video_frames_of_a_camera(video_stream): - """A processor of a video track reads its frames, and closes when the track stops""" +async def test_video_frames_of_a_camera(video_stream: webrtc.MediaStream) -> None: + """A processor of a video track reads its frames, and closes when the track stops.""" track = video_stream.get_tracks()[0] processor = webrtc.MediaStreamTrackProcessor(track) reader = processor.readable.get_reader() @@ -53,8 +55,8 @@ async def test_video_frames_of_a_camera(video_stream): @pytest.mark.asyncio -async def test_audio_data_of_a_microphone(audio_stream): - """A processor of an audio track reads its samples, 10 ms at a time""" +async def test_audio_data_of_a_microphone(audio_stream: webrtc.MediaStream) -> None: + """A processor of an audio track reads its samples, 10 ms at a time.""" track = audio_stream.get_tracks()[0] reader = webrtc.MediaStreamTrackProcessor(track=track).readable.get_reader() audio = (await read(reader)).value @@ -66,9 +68,8 @@ async def test_audio_data_of_a_microphone(audio_stream): await asyncio.wait_for(reader.closed, TIMEOUT) -@pytest.mark.asyncio -async def test_init_forms(): - """The processor takes a track, an init or a dictionary, and rejects anything else""" +def test_init_forms() -> None: + """The processor takes a track, an init or a dictionary, and rejects anything else.""" generator = webrtc.VideoTrackGenerator() track = generator.track for processor in ( @@ -85,8 +86,8 @@ async def test_init_forms(): @pytest.mark.asyncio -async def test_full_buffer_drops_the_oldest_frames(video_stream): - """Frames nobody reads are dropped once the buffer is full, oldest first, and counted""" +async def test_full_buffer_drops_the_oldest_frames(video_stream: webrtc.MediaStream) -> None: + """Frames nobody reads are dropped once the buffer is full, oldest first, and counted.""" track = video_stream.get_tracks()[0] processor = webrtc.MediaStreamTrackProcessor(track, max_buffer_size=2) reader = processor.readable.get_reader() @@ -101,8 +102,8 @@ async def test_full_buffer_drops_the_oldest_frames(video_stream): @pytest.mark.asyncio -async def test_cancel_stops_reading(video_stream): - """Canceling the stream detaches the processor from the track""" +async def test_cancel_stops_reading(video_stream: webrtc.MediaStream) -> None: + """Canceling the stream detaches the processor from the track.""" processor = webrtc.MediaStreamTrackProcessor(video_stream.get_tracks()[0]) reader = processor.readable.get_reader() (await read(reader)).value.close() @@ -114,8 +115,8 @@ async def test_cancel_stops_reading(video_stream): @pytest.mark.asyncio -async def test_processor_of_an_ended_track(video_stream): - """The stream of an ended track is closed""" +async def test_processor_of_an_ended_track(video_stream: webrtc.MediaStream) -> None: + """The stream of an ended track is closed.""" track = video_stream.get_tracks()[0] track.stop() reader = webrtc.MediaStreamTrackProcessor(track).readable.get_reader() @@ -123,8 +124,8 @@ async def test_processor_of_an_ended_track(video_stream): @pytest.mark.asyncio -async def test_generator_forwards_frames_with_their_timestamps(): - """A frame written to a generator reaches a processor of its track, with its size and timestamp, and is closed""" +async def test_generator_forwards_frames_with_their_timestamps() -> None: + """A frame written to a generator reaches a processor of its track, with its size and timestamp, and is closed.""" generator = webrtc.VideoTrackGenerator() track = generator.track reader = webrtc.MediaStreamTrackProcessor(track, max_buffer_size=10).readable.get_reader() @@ -147,8 +148,8 @@ async def test_generator_forwards_frames_with_their_timestamps(): @pytest.mark.asyncio -async def test_generator_rejects_what_it_cant_send(): - """A video generator takes open VideoFrames only""" +async def test_generator_rejects_what_it_cant_send() -> None: + """A video generator takes open VideoFrames only.""" generator = webrtc.VideoTrackGenerator() writer = generator.writable.get_writer() closed = video_frame(0) @@ -161,8 +162,8 @@ async def test_generator_rejects_what_it_cant_send(): @pytest.mark.asyncio -async def test_closing_the_generator_ends_its_track(): - """Closing the writable ends the track, which closes the processor""" +async def test_closing_the_generator_ends_its_track() -> None: + """Closing the writable ends the track, which closes the processor.""" generator = webrtc.VideoTrackGenerator() track = generator.track ended = wait_for_event(track, 'ended') @@ -174,14 +175,15 @@ async def test_closing_the_generator_ends_its_track(): @pytest.mark.asyncio -async def test_muted_generator_drops_frames(): - """A muted generator mutes its track and drops the frames written""" +async def test_muted_generator_drops_frames() -> None: + """A muted generator mutes its track and drops the frames written.""" generator = webrtc.VideoTrackGenerator() track = generator.track muted = wait_for_event(track, 'mute') generator.muted = True await muted - assert generator.muted and track.muted + assert generator.muted + assert track.muted processor = webrtc.MediaStreamTrackProcessor(track, max_buffer_size=10) reader = processor.readable.get_reader() @@ -198,8 +200,8 @@ async def test_muted_generator_drops_frames(): @pytest.mark.asyncio -async def test_audio_generator_sends_10_ms_frames(): - """An audio generator is a track: it sends what's written in 10 ms frames, converted to 16 bits""" +async def test_audio_generator_sends_10_ms_frames() -> None: + """An audio generator is a track: it sends what's written in 10 ms frames, converted to 16 bits.""" generator = webrtc.MediaStreamTrackGenerator('audio') assert isinstance(generator, webrtc.MediaStreamTrack) assert generator.kind == webrtc.MediaType.audio @@ -226,8 +228,8 @@ async def test_audio_generator_sends_10_ms_frames(): generator.stop() -def test_generator_kinds(): - """A MediaStreamTrackGenerator is created for a kind, as a string, an init or a dictionary""" +def test_generator_kinds() -> None: + """A MediaStreamTrackGenerator is created for a kind, as a string, an init or a dictionary.""" assert webrtc.MediaStreamTrackGenerator('video').kind == webrtc.MediaType.video assert webrtc.MediaStreamTrackGenerator({'kind': 'audio'}).kind == webrtc.MediaType.audio assert ( @@ -239,20 +241,22 @@ def test_generator_kinds(): @pytest.mark.asyncio -async def test_pipe_processor_to_generator(video_stream): - """The frames of a track are piped through a transform to a generator, as in a browser""" +async def test_pipe_processor_to_generator(video_stream: webrtc.MediaStream) -> None: + """The frames of a track are piped through a transform to a generator, as in a browser.""" generator = webrtc.VideoTrackGenerator() reader = webrtc.MediaStreamTrackProcessor(generator.track).readable.get_reader() class Stamp: - def transform(self, frame, controller): + @staticmethod + def transform(frame: webrtc.VideoFrame, controller: webrtc.TransformStreamDefaultController) -> None: controller.enqueue(webrtc.VideoFrame(frame, timestamp=42)) frame.close() source = webrtc.MediaStreamTrackProcessor(video_stream.get_tracks()[0]).readable pipe = asyncio.ensure_future(source.pipe_through(webrtc.TransformStream(Stamp())).pipe_to(generator.writable)) frame = (await read(reader)).value - assert frame.timestamp == 42 and frame.coded_width == 640 + assert frame.timestamp == 42 + assert frame.coded_width == 640 frame.close() video_stream.get_tracks()[0].stop() @@ -261,8 +265,8 @@ def transform(self, frame, controller): @pytest.mark.asyncio -async def test_frames_wait_in_the_native_queue_only(video_stream): - """Frames are taken from the processor for pending reads only: the rest stays in its buffer, where it's dropped""" +async def test_frames_wait_in_the_native_queue_only(video_stream: webrtc.MediaStream) -> None: + """Frames are taken from the processor for pending reads only: the rest stays in its buffer, where it's dropped.""" processor = webrtc.MediaStreamTrackProcessor(video_stream.get_tracks()[0], max_buffer_size=2) reader = processor.readable.get_reader() frames = await asyncio.wait_for(asyncio.gather(*(reader.read() for _ in range(4))), TIMEOUT) diff --git a/tests/test_media_stress.py b/tests/test_media_stress.py index 914ed5f..3bd8d9f 100644 --- a/tests/test_media_stress.py +++ b/tests/test_media_stress.py @@ -7,6 +7,8 @@ """Processors and generators under stress: races that could deadlock, and cycles that could leak.""" +from __future__ import annotations + import asyncio import gc import time @@ -15,27 +17,26 @@ import pytest import webrtc -import wrtc -from tests.helpers import connect_track, rss_bytes, skip_if_sanitized, wait_until, write_video, writing +from tests.helpers import SANITIZED, connect_track, rss_bytes, skip_if_sanitized, wait_until, write_video, writing from webrtc.utils.task_queue import TaskQueue TIMEOUT = 20 -def frame(timestamp=0, width=16, height=16): +def frame(timestamp: int = 0, width: int = 16, height: int = 16) -> webrtc.VideoFrame: return webrtc.VideoFrame( bytes(width * height * 3 // 2), format='I420', coded_width=width, coded_height=height, timestamp=timestamp ) -def collected(refs): +def collected(refs: list[weakref.ref[object]]) -> bool: gc.collect() return all(ref() is None for ref in refs) @pytest.mark.asyncio -async def test_cancel_during_pending_reads(video_stream): - """Canceling while reads wait for frames settles them, over and over""" +async def test_cancel_during_pending_reads(video_stream: webrtc.MediaStream) -> None: + """Canceling while reads wait for frames settles them, over and over.""" track = video_stream.get_tracks()[0] for _ in range(50): reader = webrtc.MediaStreamTrackProcessor(track).readable.get_reader() @@ -48,13 +49,13 @@ async def test_cancel_during_pending_reads(video_stream): @pytest.mark.asyncio -async def test_stop_track_while_reading(video_stream, audio_stream): - """Tracks stopped while tasks read them close their streams""" +async def test_stop_track_while_reading(video_stream: webrtc.MediaStream, audio_stream: webrtc.MediaStream) -> None: + """Tracks stopped while tasks read them close their streams.""" tracks = [*video_stream.get_tracks(), *audio_stream.get_tracks()] read = [] - async def read_all(track): + async def read_all(track: webrtc.MediaStreamTrack) -> int: count = 0 async for media in webrtc.MediaStreamTrackProcessor(track).readable: media.close() @@ -72,8 +73,8 @@ async def read_all(track): @pytest.mark.asyncio -async def test_many_processors_of_one_track(video_stream): - """Every processor of a track gets its frames""" +async def test_many_processors_of_one_track(video_stream: webrtc.MediaStream) -> None: + """Every processor of a track gets its frames.""" track = video_stream.get_tracks()[0] readers = [webrtc.MediaStreamTrackProcessor(track).readable.get_reader() for _ in range(20)] for result in await asyncio.wait_for(asyncio.gather(*(r.read() for r in readers)), TIMEOUT): @@ -83,8 +84,8 @@ async def test_many_processors_of_one_track(video_stream): @pytest.mark.asyncio -async def test_garbage_collected_with_pending_reads(video_stream): - """Processors dropped with reads pending are collected, while their track goes on""" +async def test_garbage_collected_with_pending_reads(video_stream: webrtc.MediaStream) -> None: + """Processors dropped with reads pending are collected, while their track goes on.""" track = video_stream.get_tracks()[0] refs = [] for _ in range(20): @@ -98,11 +99,13 @@ async def test_garbage_collected_with_pending_reads(video_stream): @pytest.mark.asyncio -async def test_close_connection_while_reading(caller, callee): - """Closing the connection closes the processors of its tracks, while frames still arrive""" +async def test_close_connection_while_reading( + caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection +) -> None: + """Closing the connection closes the processors of its tracks, while frames still arrive.""" generator = webrtc.VideoTrackGenerator() - async with writing(write_video, generator, bytes(64 * 64 * 3 // 2), 64, 64, interval=0.01): - remote = await connect_track(caller, callee, generator.track, TIMEOUT) + async with writing(write_video, generator, bytes(64 * 64 * 3 // 2), (64, 64), interval=0.01): + remote = await connect_track(caller, callee, generator.track, timeout=TIMEOUT) readers = [webrtc.MediaStreamTrackProcessor(remote).readable.get_reader() for _ in range(5)] (await asyncio.wait_for(readers[0].read(), TIMEOUT)).value.close() pending = [r.read() for r in readers] @@ -115,11 +118,11 @@ async def test_close_connection_while_reading(caller, callee): generator.track.stop() -def test_loop_closed_while_frames_arrive(): - """A processor whose loop is closed doesn't block the media threads, nor fail once collected""" +def test_loop_closed_while_frames_arrive() -> None: + """A processor whose loop is closed doesn't block the media threads, nor fail once collected.""" stream = webrtc.get_user_media(audio=True, video=True) - async def start(): + async def start() -> list[webrtc.MediaStreamTrackProcessor]: processors = [webrtc.MediaStreamTrackProcessor(t) for t in stream.get_tracks()] for p in processors: p.readable.get_reader().read() @@ -138,10 +141,10 @@ async def start(): @pytest.mark.asyncio -async def test_create_and_destroy_cycles_do_not_leak(): - """Thousands of processors, generators and frames leave nothing behind""" +async def test_create_and_destroy_cycles_do_not_leak() -> None: + """Thousands of processors, generators and frames leave nothing behind.""" - async def cycle(): + async def cycle() -> tuple[weakref.ref[object], weakref.ref[object]]: generator = webrtc.VideoTrackGenerator() processor = webrtc.MediaStreamTrackProcessor(generator.track, max_buffer_size=2) reader = processor.readable.get_reader() @@ -159,13 +162,13 @@ async def cycle(): # callbacks still queued on the loop hold the last objects await wait_until(lambda: collected(refs), 'every cycle to be collected') growth = rss_bytes() - before - assert wrtc._sanitized or growth < 20 * 1024 * 1024, f'{growth / 1e6:.1f} MB more after 1000 cycles' + assert SANITIZED or growth < 20 * 1024 * 1024, f'{growth / 1e6:.1f} MB more after 1000 cycles' @skip_if_sanitized @pytest.mark.asyncio -async def test_unread_frames_do_not_grow_memory(video_stream): - """A processor nobody reads keeps at most its buffer: memory stays flat while frames keep coming""" +async def test_unread_frames_do_not_grow_memory(video_stream: webrtc.MediaStream) -> None: + """A processor nobody reads keeps at most its buffer: memory stays flat while frames keep coming.""" processor = webrtc.MediaStreamTrackProcessor(video_stream.get_tracks()[0], max_buffer_size=3) reader = processor.readable.get_reader() # warm up, then measure for 2 s @@ -178,8 +181,8 @@ async def test_unread_frames_do_not_grow_memory(video_stream): @pytest.mark.asyncio -async def test_reader_that_never_yields_queues_nothing(): - """Frames read as fast as they're written, without the loop running, leave no callbacks piling up for it""" +async def test_reader_that_never_yields_queues_nothing() -> None: + """Frames read as fast as they're written, without the loop running, leave no callbacks piling up for it.""" generator = webrtc.VideoTrackGenerator() reader = webrtc.MediaStreamTrackProcessor(generator.track, max_buffer_size=2).readable.get_reader() writer = generator.writable.get_writer() diff --git a/tests/test_native_calls.py b/tests/test_native_calls.py index 9e0895c..867ef1b 100644 --- a/tests/test_native_calls.py +++ b/tests/test_native_calls.py @@ -7,58 +7,62 @@ """call_native, which awaits native methods reporting to callbacks from libwebrtc threads.""" +from __future__ import annotations + import asyncio import threading +from types import SimpleNamespace +from typing import Callable import pytest from webrtc.utils.native_calls import call_native +OnSuccess = Callable[[object], None] +OnFailure = Callable[[SimpleNamespace], None] -class _Error: - def __init__(self, error): - self._error = error - def toPython(self): - return self._error +def _error(error: Exception) -> SimpleNamespace: + """A stand-in of the native exception passed to on_failure.""" + return SimpleNamespace(toPython=lambda: error) -def _later(callback, *args, delay=0.0): +def _later(callback: Callable[..., None], *args: object, delay: float = 0.0) -> None: threading.Timer(delay, callback, args).start() @pytest.mark.asyncio -async def test_result(): - def method(on_success, on_failure, a, b): - _later(on_success, a + b) +async def test_result() -> None: + def method(on_success: Callable[[int], None], _on_failure: OnFailure, *numbers: int) -> None: + _later(on_success, sum(numbers)) assert await call_native(method, 1, 2) == 3 @pytest.mark.asyncio -async def test_no_result(): - assert await call_native(lambda on_success, on_failure: _later(on_success)) is None +async def test_no_result() -> None: + assert await call_native(lambda on_success, _: _later(on_success)) is None @pytest.mark.asyncio -async def test_failure_is_raised_as_python_error(): - def method(on_success, on_failure): - _later(on_failure, _Error(ValueError('native'))) +async def test_failure_is_raised_as_python_error() -> None: + def method(_on_success: OnSuccess, on_failure: OnFailure) -> None: + _later(on_failure, _error(ValueError('native'))) with pytest.raises(ValueError, match='native'): await call_native(method) @pytest.mark.asyncio -async def test_late_result_after_cancel_is_dropped(): - """A result arriving after the caller was canceled doesn't reach the loop's exception handler""" +async def test_late_result_after_cancel_is_dropped() -> None: + """A result arriving after the caller was canceled doesn't reach the loop's exception handler.""" loop = asyncio.get_running_loop() errors = [] - loop.set_exception_handler(lambda loop, context: errors.append(context)) + loop.set_exception_handler(lambda _, context: errors.append(context)) settled = threading.Event() - def method(on_success, on_failure): - def succeed(): + def method(on_success: OnSuccess, _on_failure: OnFailure) -> None: + def succeed() -> None: on_success('late') settled.set() diff --git a/tests/test_robustness_chaos.py b/tests/test_robustness_chaos.py index 8cd5382..34e322a 100644 --- a/tests/test_robustness_chaos.py +++ b/tests/test_robustness_chaos.py @@ -5,8 +5,12 @@ # that can be found in the LICENSE.md file in the root of the project. # -"""Random sequences of API calls (tests/chaos.py) in processes of their own: no crash, no deadlock. A failure -prints the seed and the steps, replayed with ``python -m tests.chaos --seed --steps ``.""" +"""Random sequences of API calls (tests/chaos.py) in processes of their own: no crash, no deadlock. + +A failure prints the seed and the steps, replayed with ``python -m tests.chaos --seed --steps ``. +""" + +from __future__ import annotations import subprocess import sys @@ -16,10 +20,10 @@ from tests.helpers import ROOT -def run_chaos(seed, steps, timeout): +def run_chaos(seed: int, steps: int, timeout: float) -> None: command = [sys.executable, '-m', 'tests.chaos', '--seed', str(seed), '--steps', str(steps)] try: - result = subprocess.run(command, capture_output=True, text=True, timeout=timeout, cwd=ROOT) + result = subprocess.run(command, capture_output=True, text=True, timeout=timeout, cwd=ROOT, check=False) except subprocess.TimeoutExpired as e: pytest.fail(f'seed {seed} is stuck after:\n{(e.stdout or b"")[-3000:]}') output = result.stdout + result.stderr @@ -28,12 +32,12 @@ def run_chaos(seed, steps, timeout): @pytest.mark.parametrize('seed', range(2)) -def test_chaos(seed): +def test_chaos(seed: int) -> None: run_chaos(seed, steps=150, timeout=120) @pytest.mark.stress @pytest.mark.timeout(900) @pytest.mark.parametrize('seed', range(100, 120)) -def test_chaos_long(seed): +def test_chaos_long(seed: int) -> None: run_chaos(seed, steps=1000, timeout=600) diff --git a/tests/test_robustness_exit.py b/tests/test_robustness_exit.py index 0ffa087..09ef9b5 100644 --- a/tests/test_robustness_exit.py +++ b/tests/test_robustness_exit.py @@ -7,13 +7,15 @@ """The interpreter exits while objects are alive and busy: no hang, no crash.""" +from __future__ import annotations + import os import pytest from tests.helpers import run_isolated -BUSY_AT_EXIT = ''' +BUSY_AT_EXIT = """ import asyncio import threading import webrtc @@ -42,16 +44,16 @@ def spin(): # kept alive until the interpreter finalizes objects = asyncio.run(main()) print('exiting') -''' +""" -@pytest.mark.parametrize('attempt', range(5)) -def test_exit_while_objects_are_busy(attempt): - """Wrappers released by the last collection must not block on threads hung in the GIL""" +@pytest.mark.parametrize('_attempt', range(5)) +def test_exit_while_objects_are_busy(_attempt: int) -> None: + """Wrappers released by the last collection must not block on threads hung in the GIL.""" assert 'exiting' in run_isolated(BUSY_AT_EXIT, timeout=30) -PENDING_AT_EXIT = ''' +PENDING_AT_EXIT = """ import threading import time import webrtc @@ -68,18 +70,18 @@ def spin(): threading.Thread(target=spin, daemon=True).start() time.sleep(0.3) print('exiting') -''' +""" -@pytest.mark.parametrize('attempt', range(5)) -def test_exit_while_operations_are_pending(attempt): - """Callbacks of operations completing at exit are dropped, not run by a libwebrtc thread taking the GIL""" +@pytest.mark.parametrize('_attempt', range(5)) +def test_exit_while_operations_are_pending(_attempt: int) -> None: + """Callbacks of operations completing at exit are dropped, not run by a libwebrtc thread taking the GIL.""" assert 'exiting' in run_isolated(PENDING_AT_EXIT, timeout=30) @pytest.mark.skipif(not hasattr(os, 'fork'), reason='no fork') -def test_forked_child_leaves_the_objects_of_its_parent_alone(): - """The child of a fork doesn't block on its parent's threads; new objects work (raise on macOS)""" +def test_forked_child_leaves_the_objects_of_its_parent_alone() -> None: + """The child of a fork doesn't block on its parent's threads; new objects work (raise on macOS).""" output = run_isolated( """ import asyncio diff --git a/tests/test_robustness_media.py b/tests/test_robustness_media.py index 224617e..7f93845 100644 --- a/tests/test_robustness_media.py +++ b/tests/test_robustness_media.py @@ -7,6 +7,10 @@ """Hostile media input: buffers, sizes and formats that must be rejected rather than read, written or sent.""" +from __future__ import annotations + +from typing import Callable + import pytest import webrtc @@ -17,21 +21,21 @@ I420_SIZE = WIDTH * HEIGHT * 3 // 2 -def i420_frame(): +def i420_frame() -> webrtc.VideoFrame: return webrtc.VideoFrame(bytes(I420_SIZE), format='I420', coded_width=WIDTH, coded_height=HEIGHT, timestamp=0) -def reversed_view(size): - """A view of size bytes whose pointer is its last byte: read or written forward, it's out of its buffer""" +def reversed_view(size: int) -> memoryview: + """A view of size bytes whose pointer is its last byte: read or written forward, it's out of its buffer.""" return memoryview(bytearray(size))[::-1] -def strided_view(size): +def strided_view(size: int) -> memoryview: return memoryview(bytearray(size * 2))[::2] @pytest.mark.parametrize('view', [reversed_view, strided_view]) -def test_frame_from_non_contiguous_buffer_is_rejected(view): +def test_frame_from_non_contiguous_buffer_is_rejected(view: Callable[[int], memoryview]) -> None: with pytest.raises(TypeError, match='contiguous'): webrtc.VideoFrame(view(I420_SIZE), format='I420', coded_width=WIDTH, coded_height=HEIGHT, timestamp=0) @@ -39,7 +43,9 @@ def test_frame_from_non_contiguous_buffer_is_rejected(view): @pytest.mark.asyncio @pytest.mark.parametrize('view', [reversed_view, strided_view]) @pytest.mark.parametrize('options', [None, {'format': 'RGBA'}]) -async def test_frame_copy_to_non_contiguous_destination_is_rejected(view, options): +async def test_frame_copy_to_non_contiguous_destination_is_rejected( + view: Callable[[int], memoryview], options: dict[str, str] | None +) -> None: frame = i420_frame() with pytest.raises(TypeError, match='contiguous'): await frame.copy_to(view(frame.allocation_size(options)), options) @@ -47,7 +53,7 @@ async def test_frame_copy_to_non_contiguous_destination_is_rejected(view, option @pytest.mark.parametrize('view', [reversed_view, strided_view]) -def test_audio_copy_to_non_contiguous_destination_is_rejected(view): +def test_audio_copy_to_non_contiguous_destination_is_rejected(view: Callable[[int], memoryview]) -> None: data = webrtc.AudioData( format='s16', sample_rate=48000, number_of_frames=480, number_of_channels=2, timestamp=0, data=bytes(1920) ) @@ -56,49 +62,49 @@ def test_audio_copy_to_non_contiguous_destination_is_rejected(view): data.close() -def test_native_bounds_checks_do_not_overflow(): - """Offsets and strides near the top of size_t wrap around in naive bounds checks""" +def test_native_bounds_checks_do_not_overflow() -> None: + """Offsets and strides near the top of size_t wrap around in naive bounds checks.""" top = 2**64 - 1 planes = [(0, WIDTH), (WIDTH * HEIGHT, WIDTH // 2), (WIDTH * HEIGHT * 5 // 4, WIDTH // 2)] - with pytest.raises(ValueError): - wrtc.VideoFrameBuffer.fromData('I420', WIDTH, HEIGHT, bytes(I420_SIZE), [(top, WIDTH)] + planes[1:]) - with pytest.raises(ValueError): + with pytest.raises(ValueError, match="layout doesn't fit"): + wrtc.VideoFrameBuffer.fromData('I420', WIDTH, HEIGHT, bytes(I420_SIZE), [(top, WIDTH), *planes[1:]]) + with pytest.raises(ValueError, match="layout doesn't fit"): wrtc.VideoFrameBuffer.fromData('I420', WIDTH, HEIGHT, bytes(I420_SIZE), [(0, 2**63)] * 3) - with pytest.raises(ValueError): + with pytest.raises(ValueError, match='positive size'): wrtc.VideoFrameBuffer.fromData('I420', 2**31 - 1, 2**31 - 1, bytes(I420_SIZE), [(0, 2**31 - 1)] * 3) buffer = wrtc.VideoFrameBuffer.fromData('I420', WIDTH, HEIGHT, bytes(I420_SIZE), planes) destination = bytearray(I420_SIZE) half = (0, 0, WIDTH // 2, HEIGHT // 2, 0, WIDTH // 2) - with pytest.raises(ValueError): + with pytest.raises(ValueError, match='out of the bounds of the frame'): buffer.copyPlanes(destination, [(0, 0, WIDTH, HEIGHT, top - 100, 1), half, half]) - with pytest.raises(ValueError): + with pytest.raises(ValueError, match='out of the bounds of the frame'): buffer.copyPlanes(destination, [(0, top, WIDTH, 2, 0, WIDTH), half, half]) - with pytest.raises(ValueError): - buffer.convertTo(bytearray(16), 'RGBA', 0, 0, WIDTH, HEIGHT, top - 100, WIDTH * 4, '', False) - with pytest.raises(ValueError): - buffer.convertTo(bytearray(WIDTH * HEIGHT * 4), 'RGBA', 2**31 - 1, 0, 2, 1, 0, WIDTH * 4, '', False) + with pytest.raises(ValueError, match='destination is too small'): + buffer.convertTo(bytearray(16), 'RGBA', 0, 0, WIDTH, HEIGHT, top - 100, WIDTH * 4, '', fullRange=False) + with pytest.raises(ValueError, match='rect is out of the bounds'): + buffer.convertTo(bytearray(WIDTH * HEIGHT * 4), 'RGBA', 2**31 - 1, 0, 2, 1, 0, WIDTH * 4, '', fullRange=False) - with pytest.raises(ValueError): + with pytest.raises(ValueError, match='out of the bounds of the samples'): wrtc.copyAudioSamples(bytes(16), 's16', 2**40, 2**40, bytearray(16), 's16', 0, 0, 1) - with pytest.raises(ValueError): + with pytest.raises(ValueError, match='out of the bounds of the samples'): wrtc.copyAudioSamples(bytes(16), 's16', 1, 4, bytearray(16), 's16', 0, top, 2) - with pytest.raises(ValueError): + with pytest.raises(ValueError, match='out of the bounds of the samples'): wrtc.copyAudioSamples(bytes(16), 's16', 0, 4, bytearray(16), 's16-planar', 0, 0, 1) @pytest.mark.parametrize('rate', [float('inf'), float('nan'), 0, -1]) -def test_audio_data_sample_rate_is_positive_and_finite(rate): +def test_audio_data_sample_rate_is_positive_and_finite(rate: float) -> None: with pytest.raises(TypeError): webrtc.AudioData( format='s16', sample_rate=rate, number_of_frames=1, number_of_channels=1, timestamp=0, data=bytes(2) ) -def test_generator_rejects_audio_libwebrtc_cannot_send(): - """Audio beyond libwebrtc's frames or resampler is rejected, it aborted the process""" +def test_generator_rejects_audio_libwebrtc_cannot_send() -> None: + """Audio beyond libwebrtc's frames or resampler is rejected, it aborted the process.""" output = run_isolated( - ''' + """ import asyncio import webrtc from tests.helpers import connect @@ -133,7 +139,7 @@ async def main(): print(rate, channels, await write(rate, channels)) asyncio.run(main()) - ''' + """ ) results = [line.split()[-1] for line in output.splitlines() if line.endswith(('written', 'rejected'))] assert results == ['rejected'] * 7 + ['written'] * 4, output diff --git a/tests/test_robustness_threads.py b/tests/test_robustness_threads.py index e76e949..c2efcd1 100644 --- a/tests/test_robustness_threads.py +++ b/tests/test_robustness_threads.py @@ -7,15 +7,17 @@ """The same objects used from many Python threads at once: no crash, no deadlock, no corruption.""" +from __future__ import annotations + import pytest from tests.helpers import run_isolated -def test_constructors_from_many_threads(): - """Constructors register their Python object with the GIL: pybind11's registry was corrupted""" +def test_constructors_from_many_threads() -> None: + """Constructors register their Python object with the GIL: pybind11's registry was corrupted.""" output = run_isolated( - ''' + """ import gc import threading import time @@ -46,15 +48,15 @@ def construct(): track.stop() assert not errors, errors print('constructed') - ''', + """, timeout=90, ) assert 'constructed' in output @pytest.mark.parametrize('kind', ['video generator', 'audio generator', 'processor']) -def test_released_while_libwebrtc_threads_wait_for_the_gil(kind): - """A track's proxy released with the GIL deadlocked with the signaling thread""" +def test_released_while_libwebrtc_threads_wait_for_the_gil(kind: str) -> None: + """A track's proxy released with the GIL deadlocked with the signaling thread.""" output = run_isolated( f""" import asyncio @@ -100,8 +102,8 @@ async def main(): assert 'released' in output -def test_wrappers_created_and_released_on_many_threads(): - """Releasing a wrapper with the GIL waited for a holder lock held across a BlockingCall""" +def test_wrappers_created_and_released_on_many_threads() -> None: + """Releasing a wrapper with the GIL waited for a holder lock held across a BlockingCall.""" output = run_isolated( """ import asyncio @@ -153,8 +155,8 @@ def release(): assert 'done' in output -def test_wrappers_created_while_a_description_wraps_them(): - """A thread creating a wrapper held the holder's lock waiting for the signaling thread, which waited for it""" +def test_wrappers_created_while_a_description_wraps_them() -> None: + """A thread creating a wrapper held the holder's lock waiting for the signaling thread, which waited for it.""" output = run_isolated( """ import asyncio @@ -201,8 +203,8 @@ def create(): assert 'done' in output -def test_objects_of_connections_read_while_they_connect(): - """Wrapping under a lock of the connection (its SCTP transport, its tracks) waited for the signaling thread""" +def test_objects_of_connections_read_while_they_connect() -> None: + """Wrapping under a lock of the connection (its SCTP transport, its tracks) waited for the signaling thread.""" output = run_isolated( """ import asyncio diff --git a/tests/test_rtp_sender_receiver.py b/tests/test_rtp_sender_receiver.py index c7d9a78..941fba6 100644 --- a/tests/test_rtp_sender_receiver.py +++ b/tests/test_rtp_sender_receiver.py @@ -7,6 +7,8 @@ """Senders and receivers: capabilities, parameters, codecs, tracks, DTMF and synchronization sources.""" +from __future__ import annotations + import asyncio import time @@ -16,8 +18,8 @@ from tests.helpers import connect, exchange_offer_answer, next_task, wait_for_event, wait_until -def test_capabilities(): - """Senders and receivers of audio and video have codecs and header extensions, of data none""" +def test_capabilities() -> None: + """Senders and receivers of audio and video have codecs and header extensions, of data none.""" audio = webrtc.RTCRtpSender.get_capabilities(webrtc.MediaType.audio) assert any(codec.mime_type == 'audio/opus' for codec in audio.codecs) assert audio.header_extensions @@ -25,15 +27,15 @@ def test_capabilities(): assert webrtc.RTCRtpSender.get_capabilities('data') is None -def add_simulcast_sender(pc): +def add_simulcast_sender(pc: webrtc.RTCPeerConnection) -> webrtc.RTCRtpSender: init = webrtc.RtpTransceiverInit( send_encodings=[webrtc.RTCRtpEncodingParameters(rid='hi'), webrtc.RTCRtpEncodingParameters(rid='lo')] ) return pc.add_transceiver(webrtc.MediaType.video, init).sender -def test_default_send_parameters(pc): - """Video encodings scale down by powers of 2 by default, and the parameters of a task share a transaction""" +def test_default_send_parameters(pc: webrtc.RTCPeerConnection) -> None: + """Video encodings scale down by powers of 2 by default, and the parameters of a task share a transaction.""" sender = add_simulcast_sender(pc) parameters = sender.get_parameters() assert [e.scale_resolution_down_by for e in parameters.encodings] == [2.0, 1.0] @@ -41,8 +43,8 @@ def test_default_send_parameters(pc): @pytest.mark.asyncio -async def test_set_send_parameters(pc): - """Changed encodings are applied""" +async def test_set_send_parameters(pc: webrtc.RTCPeerConnection) -> None: + """Changed encodings are applied.""" sender = add_simulcast_sender(pc) parameters = sender.get_parameters() parameters.encodings[0].max_bitrate = 500_000 @@ -52,12 +54,13 @@ async def test_set_send_parameters(pc): await next_task() changed = sender.get_parameters() - assert changed.encodings[0].max_bitrate == 500_000 and not changed.encodings[1].active + assert changed.encodings[0].max_bitrate == 500_000 + assert not changed.encodings[1].active @pytest.mark.asyncio -async def test_send_parameters_out_of_range(pc): - """An encoding can't scale the resolution up""" +async def test_send_parameters_out_of_range(pc: webrtc.RTCPeerConnection) -> None: + """An encoding can't scale the resolution up.""" sender = add_simulcast_sender(pc) parameters = sender.get_parameters() parameters.encodings[0].scale_resolution_down_by = 0.5 @@ -66,8 +69,8 @@ async def test_send_parameters_out_of_range(pc): @pytest.mark.asyncio -async def test_send_parameters_with_other_encodings(pc): - """The number of encodings can't change""" +async def test_send_parameters_with_other_encodings(pc: webrtc.RTCPeerConnection) -> None: + """The number of encodings can't change.""" sender = add_simulcast_sender(pc) parameters = sender.get_parameters() parameters.encodings.pop() @@ -76,8 +79,8 @@ async def test_send_parameters_with_other_encodings(pc): @pytest.mark.asyncio -async def test_parameters_expire_with_their_task(pc): - """Parameters are only accepted in the task that got them""" +async def test_parameters_expire_with_their_task(pc: webrtc.RTCPeerConnection) -> None: + """Parameters are only accepted in the task that got them.""" sender = add_simulcast_sender(pc) parameters = sender.get_parameters() await next_task() @@ -86,32 +89,32 @@ async def test_parameters_expire_with_their_task(pc): @pytest.mark.parametrize( - 'encodings', + ('encodings', 'error'), [ - [{'rid': 'a'}, {'rid': 'a'}], - [{'rid': 'a'}, {}], - [{'rid': 'no-dash'}], - [{'rid': ''}], + ([{'rid': 'a'}, {'rid': 'a'}], 'needs a distinct rid'), + ([{'rid': 'a'}, {}], 'needs a distinct rid'), + ([{'rid': 'no-dash'}], 'not a valid rid'), + ([{'rid': ''}], 'not a valid rid'), ], ids=['duplicate rid', 'missing rid', 'invalid rid', 'empty rid'], ) -def test_invalid_send_encodings(pc, encodings): - """The rids of send encodings are unique, present when there are several, and alphanumeric""" +def test_invalid_send_encodings(pc: webrtc.RTCPeerConnection, encodings: list[dict[str, str]], error: str) -> None: + """The rids of send encodings are unique, present when there are several, and alphanumeric.""" init = webrtc.RtpTransceiverInit(send_encodings=[webrtc.RTCRtpEncodingParameters(**e) for e in encodings]) - with pytest.raises(ValueError): + with pytest.raises(ValueError, match=error): pc.add_transceiver(webrtc.MediaType.video, init) -def test_send_encoding_of_an_unknown_codec(pc): - """The codec of a send encoding must be one the sender supports""" +def test_send_encoding_of_an_unknown_codec(pc: webrtc.RTCPeerConnection) -> None: + """The codec of a send encoding must be one the sender supports.""" unknown = webrtc.RTCRtpEncodingParameters(codec=webrtc.RTCRtpCodec('audio/unknown', 8000)) with pytest.raises(webrtc.OperationError): pc.add_transceiver(webrtc.MediaType.audio, webrtc.RtpTransceiverInit(send_encodings=[unknown])) @pytest.mark.asyncio -async def test_negotiated_codecs(caller, callee): - """The parameters of senders and receivers list their negotiated codecs""" +async def test_negotiated_codecs(caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection) -> None: + """The parameters of senders and receivers list their negotiated codecs.""" transceiver = caller.add_transceiver(webrtc.MediaType.audio) await exchange_offer_answer(caller, callee) @@ -120,8 +123,10 @@ async def test_negotiated_codecs(caller, callee): @pytest.mark.asyncio -async def test_replace_track(pc, audio_stream, video_stream): - """A track is replaced by one of the same kind, or by None""" +async def test_replace_track( + pc: webrtc.RTCPeerConnection, audio_stream: webrtc.MediaStream, video_stream: webrtc.MediaStream +) -> None: + """A track is replaced by one of the same kind, or by None.""" (audio,), (video,) = audio_stream.get_tracks(), video_stream.get_tracks() sender = pc.add_transceiver(webrtc.MediaType.audio).sender await sender.replace_track(audio) @@ -132,8 +137,8 @@ async def test_replace_track(pc, audio_stream, video_stream): assert sender.track is None -def test_codec_preferences_and_header_extensions(pc): - """Codec preferences take supported codecs only, and header extensions to negotiate can be stopped""" +def test_codec_preferences_and_header_extensions(pc: webrtc.RTCPeerConnection) -> None: + """Codec preferences take supported codecs only, and header extensions to negotiate can be stopped.""" transceiver = pc.add_transceiver(webrtc.MediaType.audio) opus = [c for c in webrtc.RTCRtpReceiver.get_capabilities('audio').codecs if c.mime_type == 'audio/opus'] transceiver.set_codec_preferences(opus) @@ -148,8 +153,10 @@ def test_codec_preferences_and_header_extensions(pc): @pytest.mark.asyncio -async def test_sender_codecs_leave_out_unknown_remote_codecs(caller, callee): - """The codecs of a sender are the negotiated ones it knows, and are read-only""" +async def test_sender_codecs_leave_out_unknown_remote_codecs( + caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection +) -> None: + """The codecs of a sender are the negotiated ones it knows, and are read-only.""" sender = caller.add_transceiver(webrtc.MediaType.audio).sender await caller.set_local_description() await callee.set_remote_description(caller.local_description) @@ -169,8 +176,8 @@ async def test_sender_codecs_leave_out_unknown_remote_codecs(caller, callee): @pytest.mark.asyncio -async def test_setting_a_description_expires_sender_parameters(pc): - """Setting a description changes what parameters are valid, so earlier ones expire""" +async def test_setting_a_description_expires_sender_parameters(pc: webrtc.RTCPeerConnection) -> None: + """Setting a description changes what parameters are valid, so earlier ones expire.""" sender = pc.add_transceiver(webrtc.MediaType.audio).sender parameters = sender.get_parameters() await pc.set_local_description() @@ -179,8 +186,8 @@ async def test_setting_a_description_expires_sender_parameters(pc): @pytest.mark.asyncio -async def test_set_parameters_key_frames(caller, callee): - """A key frame is requested per encoding""" +async def test_set_parameters_key_frames(caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection) -> None: + """A key frame is requested per encoding.""" sender = caller.add_transceiver(webrtc.MediaType.video).sender await exchange_offer_answer(caller, callee) with pytest.raises(webrtc.InvalidModificationError): @@ -189,8 +196,8 @@ async def test_set_parameters_key_frames(caller, callee): @pytest.mark.asyncio -async def test_set_parameters_after_rollback(pc): - """A sender rolled back out of its offer has no media channel: setting parameters rejects, not hangs""" +async def test_set_parameters_after_rollback(pc: webrtc.RTCPeerConnection) -> None: + """A sender rolled back out of its offer has no media channel: setting parameters rejects, not hangs.""" sender = pc.add_transceiver(webrtc.MediaType.video).sender await pc.set_local_description() await pc.set_local_description({'type': 'rollback'}) @@ -199,26 +206,32 @@ async def test_set_parameters_after_rollback(pc): @pytest.mark.asyncio -async def test_simulcast_receiver_parameters(caller, callee): - """The receiver of simulcast has the negotiated codecs and header extensions""" +async def test_simulcast_receiver_parameters( + caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection +) -> None: + """The receiver of simulcast has the negotiated codecs and header extensions.""" encodings = [webrtc.RTCRtpEncodingParameters(rid='a'), webrtc.RTCRtpEncodingParameters(rid='b')] caller.add_transceiver(webrtc.MediaType.video, webrtc.RtpTransceiverInit(send_encodings=encodings)) await exchange_offer_answer(caller, callee) parameters = callee.get_transceivers()[0].receiver.get_parameters() - assert parameters.codecs and parameters.header_extensions + assert parameters.codecs + assert parameters.header_extensions @pytest.mark.asyncio -async def test_dtmf(caller, callee, audio_stream): - """An audio sender plays tones, normalized and checked, announcing each one, then an empty one""" +async def test_dtmf( + caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection, audio_stream: webrtc.MediaStream +) -> None: + """An audio sender plays tones, normalized and checked, announcing each one, then an empty one.""" sender = caller.add_track(audio_stream.get_tracks()[0], audio_stream) await connect(caller, callee) dtmf = sender.dtmf - assert dtmf is not None and dtmf.can_insert_dtmf + assert dtmf is not None + assert dtmf.can_insert_dtmf tones = [] dtmf.on('tonechange', lambda event: tones.append(event.tone)) - done = wait_for_event(dtmf, 'tonechange', timeout=5, predicate=lambda event: event.tone == '') + done = wait_for_event(dtmf, 'tonechange', timeout=5, predicate=lambda event: not event.tone) with pytest.raises(webrtc.InvalidCharacterError): dtmf.insert_dtmf('12X') @@ -229,8 +242,10 @@ async def test_dtmf(caller, callee, audio_stream): @pytest.mark.asyncio -async def test_synchronization_sources(caller, callee, video_stream): - """A receiver reports the source of the media it decodes""" +async def test_synchronization_sources( + caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection, video_stream: webrtc.MediaStream +) -> None: + """A receiver reports the source of the media it decodes.""" caller.add_track(video_stream.get_tracks()[0], video_stream) remote_track = wait_for_event(callee, 'track') await connect(caller, callee) @@ -258,8 +273,8 @@ async def test_synchronization_sources(caller, callee, video_stream): {'scale_resolution_down_by': float('nan')}, ], ) -def test_encodings_have_their_webidl_types(pc, encoding): - """An [EnforceRange] unsigned long and restricted doubles: other values are a TypeError, not sent to libwebrtc""" +def test_encodings_have_their_webidl_types(pc: webrtc.RTCPeerConnection, encoding: dict[str, float]) -> None: + """An [EnforceRange] unsigned long and restricted doubles: other values are a TypeError, not sent to libwebrtc.""" with pytest.raises(TypeError): pc.add_transceiver( webrtc.MediaType.video, @@ -267,14 +282,14 @@ def test_encodings_have_their_webidl_types(pc, encoding): ) -def test_encoding_bitrate_beyond_an_int_is_no_limit(pc): +def test_encoding_bitrate_beyond_an_int_is_no_limit(pc: webrtc.RTCPeerConnection) -> None: init = webrtc.RtpTransceiverInit(send_encodings=[webrtc.RTCRtpEncodingParameters(max_bitrate=2**32 - 1)]) sender = pc.add_transceiver(webrtc.MediaType.video, init).sender assert sender.get_parameters().encodings[0].max_bitrate == 2**31 - 1 -def test_transceiver_init_as_a_dictionary(pc): - """As in browsers, with camelCase or snake_case names, the encodings too; unknown members are ignored""" +def test_transceiver_init_as_a_dictionary(pc: webrtc.RTCPeerConnection) -> None: + """As in browsers, with camelCase or snake_case names, the encodings too; unknown members are ignored.""" init = {'direction': 'sendonly', 'sendEncodings': [{'rid': 'a', 'maxBitrate': 100000}, {'rid': 'b'}], 'x': 1} transceiver = pc.add_transceiver(webrtc.MediaType.video, init) assert transceiver.direction == webrtc.TransceiverDirection.sendonly diff --git a/tests/test_session_description.py b/tests/test_session_description.py index a85c3ae..78a5bed 100644 --- a/tests/test_session_description.py +++ b/tests/test_session_description.py @@ -7,13 +7,15 @@ """RTCSessionDescription, created as in the specification: a type is required, the SDP is optional.""" +from __future__ import annotations + import pytest import webrtc -def test_type_is_required(): - """The init of a description needs its type, as its WebIDL dictionary does""" +def test_type_is_required() -> None: + """The init of a description needs its type, as its WebIDL dictionary does.""" with pytest.raises(TypeError): webrtc.RTCSessionDescription() with pytest.raises(TypeError): @@ -25,9 +27,11 @@ def test_type_is_required(): @pytest.mark.parametrize( 'init', [{'type': 'rollback'}, webrtc.RTCSessionDescriptionInit('rollback'), webrtc.RTCSdpType.rollback] ) -def test_sdp_is_empty_by_default(init): - """A description from a dictionary, an init or a type has an empty SDP""" +def test_sdp_is_empty_by_default( + init: dict[str, str] | webrtc.RTCSessionDescriptionInit | webrtc.RTCSdpType, +) -> None: + """A description from a dictionary, an init or a type has an empty SDP.""" description = webrtc.RTCSessionDescription(init) assert description.type == webrtc.RTCSdpType.rollback - assert description.sdp == '' + assert not description.sdp assert description.to_json() == {'type': 'rollback', 'sdp': ''} diff --git a/tests/test_stats.py b/tests/test_stats.py index 3452ed0..1092d77 100644 --- a/tests/test_stats.py +++ b/tests/test_stats.py @@ -7,6 +7,8 @@ """Stats of a connection, its senders and its receivers.""" +from __future__ import annotations + import time import pytest @@ -15,8 +17,10 @@ from tests.helpers import connect, wait_for_event, wait_until, wait_until_unmuted -async def send_audio(caller, callee, stream): - """Connects, sending the audio of the stream, and returns the remote track once media arrives""" +async def send_audio( + caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection, stream: webrtc.MediaStream +) -> webrtc.MediaStreamTrack: + """Connects, sending the audio of the stream, and returns the remote track once media arrives.""" caller.add_track(stream.get_tracks()[0], stream) track_event = wait_for_event(callee, 'track') await connect(caller, callee) @@ -26,8 +30,10 @@ async def send_audio(caller, callee, stream): @pytest.mark.asyncio -async def test_connection_stats(caller, callee, audio_stream): - """A report of the connection has its stats, as dictionaries and attributes, with timestamps in milliseconds""" +async def test_connection_stats( + caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection, audio_stream: webrtc.MediaStream +) -> None: + """A report of the connection has its stats, as dictionaries and attributes, with timestamps in milliseconds.""" await send_audio(caller, callee, audio_stream) report = await caller.get_stats() @@ -39,51 +45,62 @@ async def test_connection_stats(caller, callee, audio_stream): @pytest.mark.asyncio -async def test_sender_stats(caller, callee, audio_stream): - """A report of a sender has the stats of what it sends, and is the one of its track""" +async def test_sender_stats( + caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection, audio_stream: webrtc.MediaStream +) -> None: + """A report of a sender has the stats of what it sends, and is the one of its track.""" await send_audio(caller, callee, audio_stream) sender_report = await caller.get_senders()[0].get_stats() - assert sender_report.of_type('outbound-rtp') and not sender_report.of_type('inbound-rtp') + assert sender_report.of_type('outbound-rtp') + assert not sender_report.of_type('inbound-rtp') assert len(await caller.get_stats(audio_stream.get_tracks()[0])) == len(sender_report) @pytest.mark.asyncio -async def test_receiver_stats(caller, callee, audio_stream): - """A report of a receiver has the stats of what it receives""" +async def test_receiver_stats( + caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection, audio_stream: webrtc.MediaStream +) -> None: + """A report of a receiver has the stats of what it receives.""" await send_audio(caller, callee, audio_stream) receiver = callee.get_receivers()[0] - async def receives(): + async def receives() -> list[webrtc.RTCStats]: return (await receiver.get_stats()).of_type('inbound-rtp') await wait_until(receives, 'inbound-rtp stats') @pytest.mark.asyncio -async def test_stats_of_a_track_the_connection_does_not_send(pc, audio_stream, audio_stream2): - """Only the track of a sender of the connection selects stats""" +async def test_stats_of_a_track_the_connection_does_not_send( + pc: webrtc.RTCPeerConnection, audio_stream: webrtc.MediaStream, audio_stream2: webrtc.MediaStream +) -> None: + """Only the track of a sender of the connection selects stats.""" pc.add_track(audio_stream.get_tracks()[0]) with pytest.raises(webrtc.InvalidAccessError): await pc.get_stats(audio_stream2.get_tracks()[0]) @pytest.mark.asyncio -async def test_closed_connection_has_stats(caller, callee, audio_stream): - """A closed connection still has stats""" +async def test_closed_connection_has_stats( + caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection, audio_stream: webrtc.MediaStream +) -> None: + """A closed connection still has stats.""" await send_audio(caller, callee, audio_stream) caller.close() assert (await caller.get_stats()).of_type('peer-connection') @pytest.mark.asyncio -async def test_remote_audio_is_played_out(caller, callee, audio_stream): - """The audio device pulls playout, so received audio is decoded like in a browser playing it""" +async def test_remote_audio_is_played_out( + caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection, audio_stream: webrtc.MediaStream +) -> None: + """The audio device pulls playout, so received audio is decoded like in a browser playing it.""" await send_audio(caller, callee, audio_stream) receiver = callee.get_receivers()[0] - async def decoded(): + async def decoded() -> bool: inbound = (await receiver.get_stats()).of_type('inbound-rtp') - return inbound and inbound[0].get('totalSamplesReceived', 0) > 0 + return bool(inbound) and inbound[0].get('totalSamplesReceived', 0) > 0 await wait_until(decoded, 'decoded remote audio') diff --git a/tests/test_streams.py b/tests/test_streams.py index a4623fc..270893d 100644 --- a/tests/test_streams.py +++ b/tests/test_streams.py @@ -7,57 +7,63 @@ """The streams media processing uses: readable, writable and transform streams of objects.""" +from __future__ import annotations + import asyncio import gc import weakref +from typing import TYPE_CHECKING, NoReturn import pytest import webrtc from tests.helpers import wait_until +if TYPE_CHECKING: + from collections.abc import Iterable + class Chunks: - """An underlying source of the given chunks, then closed""" + """An underlying source of the given chunks, then closed.""" - def __init__(self, chunks): + def __init__(self, chunks: Iterable[object]) -> None: self.chunks = list(chunks) self.pulls = 0 - self.canceled = None + self.canceled: object = None - def pull(self, controller): + def pull(self, controller: webrtc.ReadableStreamDefaultController) -> None: self.pulls += 1 if self.chunks: controller.enqueue(self.chunks.pop(0)) else: controller.close() - def cancel(self, reason): + def cancel(self, reason: object) -> None: self.canceled = reason class Controlled: - """An underlying source keeping its controller, for the test to enqueue or error""" + """An underlying source keeping its controller, for the test to enqueue or error.""" - def start(self, controller): + def start(self, controller: webrtc.ReadableStreamDefaultController) -> None: self.controller = controller @pytest.mark.asyncio -async def test_read_until_done(): - """A reader reads each chunk, then done""" +async def test_read_until_done() -> None: + """A reader reads each chunk, then done.""" stream = webrtc.ReadableStream(Chunks([1, 2])) reader = stream.get_reader() assert stream.locked - assert await reader.read() == webrtc.ReadableStreamReadResult(1, False) - assert await reader.read() == webrtc.ReadableStreamReadResult(2, False) + assert await reader.read() == webrtc.ReadableStreamReadResult(value=1, done=False) + assert await reader.read() == webrtc.ReadableStreamReadResult(value=2, done=False) assert (await reader.read()).done assert await reader.closed is None @pytest.mark.asyncio -async def test_reads_are_requested_when_called(): - """Reads are pending from the call on, like promises, and settled in order""" +async def test_reads_are_requested_when_called() -> None: + """Reads are pending from the call on, like promises, and settled in order.""" source = Controlled() reader = webrtc.ReadableStream(source, high_water_mark=0).get_reader() reads = [reader.read() for _ in range(3)] @@ -67,8 +73,8 @@ async def test_reads_are_requested_when_called(): @pytest.mark.asyncio -async def test_async_iteration_and_cancel(): - """async for reads every chunk, and breaking out of it cancels the stream""" +async def test_async_iteration_and_cancel() -> None: + """Async for reads every chunk, and breaking out of it cancels the stream.""" source = Chunks(range(10)) stream = webrtc.ReadableStream(source) seen = [] @@ -83,10 +89,10 @@ async def test_async_iteration_and_cancel(): @pytest.mark.asyncio -async def test_cancel_reaches_source(): - """Canceling a stream calls the source and settles pending reads as done""" +async def test_cancel_reaches_source() -> None: + """Canceling a stream calls the source and settles pending reads as done.""" source = Chunks([]) - source.pull = lambda controller: None + source.pull = lambda _: None stream = webrtc.ReadableStream(source, high_water_mark=0) reader = stream.get_reader() read = reader.read() @@ -96,21 +102,21 @@ async def test_cancel_reaches_source(): @pytest.mark.asyncio -async def test_errored_stream_rejects_reads(): - """An error of the source fails pending and later reads""" +async def test_errored_stream_rejects_reads() -> None: + """An error of the source fails pending and later reads.""" source = Controlled() reader = webrtc.ReadableStream(source).get_reader() read = reader.read() source.controller.error(ValueError('broken')) - with pytest.raises(ValueError): + with pytest.raises(ValueError, match='broken'): await read - with pytest.raises(ValueError): + with pytest.raises(ValueError, match='broken'): await reader.read() @pytest.mark.asyncio -async def test_locked_stream(): - """A locked stream has no second reader until the first one is released""" +async def test_locked_stream() -> None: + """A locked stream has no second reader until the first one is released.""" stream = webrtc.ReadableStream(Chunks([1])) reader = stream.get_reader() with pytest.raises(TypeError): @@ -124,12 +130,13 @@ async def test_locked_stream(): @pytest.mark.asyncio -async def test_writer_backpressure_and_order(): - """Writes reach the sink in order, one at a time, and ready follows the queue""" +async def test_writer_backpressure_and_order() -> None: + """Writes reach the sink in order, one at a time, and ready follows the queue.""" written = [] class Sink: - async def write(self, chunk, controller): + @staticmethod + async def write(chunk: int, _controller: webrtc.WritableStreamDefaultController) -> None: await asyncio.sleep(0.01) written.append(chunk) @@ -145,12 +152,14 @@ async def write(self, chunk, controller): @pytest.mark.asyncio -async def test_failed_write_errors_the_stream(): - """A sink that fails a write errors the stream""" +async def test_failed_write_errors_the_stream() -> None: + """A sink that fails a write errors the stream.""" class Sink: - def write(self, chunk, controller): - raise TypeError('not this') + @staticmethod + def write(_chunk: int, _controller: webrtc.WritableStreamDefaultController) -> NoReturn: + msg = 'not this' + raise TypeError(msg) writer = webrtc.WritableStream(Sink()).get_writer() with pytest.raises(TypeError): @@ -162,15 +171,17 @@ def write(self, chunk, controller): @pytest.mark.asyncio -async def test_abort_drops_queued_writes(): - """Aborting fails the writes not done yet and tells the sink""" +async def test_abort_drops_queued_writes() -> None: + """Aborting fails the writes not done yet and tells the sink.""" reasons = [] class Sink: - async def write(self, chunk, controller): + @staticmethod + async def write(_chunk: int, _controller: webrtc.WritableStreamDefaultController) -> None: await asyncio.sleep(0.05) - def abort(self, reason): + @staticmethod + def abort(reason: str) -> None: reasons.append(reason) writer = webrtc.WritableStream(Sink()).get_writer() @@ -183,16 +194,18 @@ def abort(self, reason): @pytest.mark.asyncio -async def test_pipe_through_transform(): - """A readable stream piped through a transform stream into a writable one""" +async def test_pipe_through_transform() -> None: + """A readable stream piped through a transform stream into a writable one.""" written = [] class Double: - def transform(self, chunk, controller): + @staticmethod + def transform(chunk: int, controller: webrtc.TransformStreamDefaultController) -> None: controller.enqueue(chunk * 2) class Sink: - def write(self, chunk, controller): + @staticmethod + def write(chunk: int, _controller: webrtc.WritableStreamDefaultController) -> None: written.append(chunk) readable = webrtc.ReadableStream(Chunks([1, 2, 3])).pipe_through(webrtc.TransformStream(Double())) @@ -201,83 +214,89 @@ def write(self, chunk, controller): @pytest.mark.asyncio -async def test_pipe_to_aborts_on_error(): - """An error of the source aborts the destination""" +async def test_pipe_to_aborts_on_error() -> None: + """An error of the source aborts the destination.""" aborted = [] class Source: - def pull(self, controller): + @staticmethod + def pull(controller: webrtc.ReadableStreamDefaultController) -> None: controller.error(ValueError('source failed')) class Sink: - def abort(self, reason): + @staticmethod + def abort(reason: Exception) -> None: aborted.append(reason) - with pytest.raises(ValueError): + with pytest.raises(ValueError, match='source failed'): await webrtc.ReadableStream(Source()).pipe_to(webrtc.WritableStream(Sink())) assert isinstance(aborted[0], ValueError) -def test_streams_need_a_loop(): - """Readers and writers use futures of the running loop""" +def test_streams_need_a_loop() -> None: + """Readers and writers use futures of the running loop.""" with pytest.raises(RuntimeError): webrtc.ReadableStream(Chunks([])).get_reader() +class Woken: + """An underlying source of 50 numbers, each pulled when what it waits for (known only weakly) is woken.""" + + def __init__(self, waiting: weakref.WeakSet[asyncio.Future[None]]) -> None: + self.waiting = waiting + self.next = 0 + + def pull(self, controller: webrtc.ReadableStreamDefaultController) -> asyncio.Future[None]: + # kept by the source, as a processor keeps its pending read + woken = self.woken = asyncio.get_running_loop().create_future() + self.waiting.add(woken) + woken.add_done_callback(lambda _: self.deliver(controller)) + return woken + + def deliver(self, controller: webrtc.ReadableStreamDefaultController) -> None: + if self.next == 50: + controller.close() + else: + controller.enqueue(self.next) + self.next += 1 + + +def wake(waiting: weakref.WeakSet[asyncio.Future[None]]) -> None: + for woken in list(waiting): + if not woken.done(): + woken.set_result(None) + + @pytest.mark.asyncio -async def test_pipe_goes_on_when_nothing_references_it(): - """A pipe nobody references goes on: a collected one errored its streams with GeneratorExit""" +async def test_pipe_goes_on_when_nothing_references_it() -> None: + """A pipe nobody references goes on: a collected one errored its streams with GeneratorExit.""" # what the source waits for, known only weakly, like a native object waking it - waiting = weakref.WeakSet() - - class Woken: - def __init__(self): - self.next = 0 - - def pull(self, controller): - # kept by the source, as a processor keeps its pending read - woken = self.woken = asyncio.get_running_loop().create_future() - waiting.add(woken) - woken.add_done_callback(lambda _: self.deliver(controller)) - return woken - - def deliver(self, controller): - if self.next == 50: - controller.close() - else: - controller.enqueue(self.next) - self.next += 1 - + waiting: weakref.WeakSet[asyncio.Future[None]] = weakref.WeakSet() written = [] - class Sink: - def write(self, chunk, controller): - written.append(chunk) - - def start(): + def start() -> asyncio.Future[None]: # only the last pipe is referenced - source = webrtc.ReadableStream(Woken(), high_water_mark=0) - return source.pipe_through(webrtc.TransformStream()).pipe_to(webrtc.WritableStream(Sink())) + source = webrtc.ReadableStream(Woken(waiting), high_water_mark=0) + sink = webrtc.WritableStream({'write': lambda chunk, _: written.append(chunk)}) + return source.pipe_through(webrtc.TransformStream()).pipe_to(sink) done = start() deadline = asyncio.get_running_loop().time() + 5 while not done.done() and asyncio.get_running_loop().time() < deadline: gc.collect() - for woken in list(waiting): - if not woken.done(): - woken.set_result(None) + wake(waiting) await asyncio.sleep(0.005) await asyncio.wait_for(done, 1) assert written == list(range(50)) @pytest.mark.asyncio -async def test_sources_sinks_and_transformers_as_dictionaries(): - """As in browsers, methods may be members of a dictionary: they were ignored, a transform changed nothing""" +async def test_sources_sinks_and_transformers_as_dictionaries() -> None: + """As in browsers, methods may be members of a dictionary: they were ignored, a transform changed nothing.""" written = [] source = webrtc.ReadableStream({'pull': lambda controller: controller.enqueue(2)}) transform = webrtc.TransformStream({'transform': lambda chunk, controller: controller.enqueue(chunk * 10)}) - sink = webrtc.WritableStream({'write': lambda chunk, controller: written.append(chunk)}) + sink = webrtc.WritableStream({'write': lambda chunk, _: written.append(chunk)}) pipe = source.pipe_through(transform).pipe_to(sink) await wait_until(lambda: len(written) >= 3, 'chunks written') pipe.cancel() diff --git a/tests/test_task_queue.py b/tests/test_task_queue.py index a6d2c1e..0076b9f 100644 --- a/tests/test_task_queue.py +++ b/tests/test_task_queue.py @@ -7,6 +7,8 @@ """Order of the callbacks of TaskQueue, which delivers events and results of operations.""" +from __future__ import annotations + import asyncio import gc import threading @@ -18,13 +20,13 @@ @pytest.mark.asyncio -async def test_posted_callbacks_run_before_later_timers(): - """Callbacks posted from another thread run before a timer set after they were posted""" +async def test_posted_callbacks_run_before_later_timers() -> None: + """Callbacks posted from another thread run before a timer set after they were posted.""" loop = asyncio.get_running_loop() queue = TaskQueue.of(loop) order = [] - def post_many(): + def post_many() -> None: for i in range(10): queue.post(order.append, i) @@ -33,7 +35,7 @@ def post_many(): thread.join() timer = loop.create_future() - def on_timer(): + def on_timer() -> None: order.append('timer') timer.set_result(None) @@ -43,18 +45,18 @@ def on_timer(): @pytest.mark.asyncio -async def test_microtasks_of_a_callback_run_before_the_next_one(): - """What a callback schedules with call_soon runs before the next posted callback""" +async def test_microtasks_of_a_callback_run_before_the_next_one() -> None: + """What a callback schedules with call_soon runs before the next posted callback.""" loop = asyncio.get_running_loop() queue = TaskQueue.of(loop) order = [] done = loop.create_future() - def first(): + def first() -> None: order.append('first') loop.call_soon(order.append, 'microtask') - def second(): + def second() -> None: order.append('second') done.set_result(None) @@ -65,14 +67,14 @@ def second(): @pytest.mark.asyncio -async def test_resumed_code_runs_before_the_next_callback_only(): - """Code a callback resumes runs before the next callback; after it, callbacks don't wait for timers""" +async def test_resumed_code_runs_before_the_next_callback_only() -> None: + """Code a callback resumes runs before the next callback; after it, callbacks don't wait for timers.""" loop = asyncio.get_running_loop() queue = TaskQueue.of(loop) order = [] resumed = asyncio.Event() - async def awaiting(): + async def awaiting() -> None: await resumed.wait() order.append('resumed') await asyncio.sleep(0) @@ -96,8 +98,8 @@ async def awaiting(): assert order[-1] == 'later' -def test_loops_are_collected_with_what_they_had_queued(): - """A closed loop is collected with what was still queued for it""" +def test_loops_are_collected_with_what_they_had_queued() -> None: + """A closed loop is collected with what was still queued for it.""" class Held: pass @@ -108,7 +110,7 @@ class Held: held = Held() held.loop = loop # never run: the loop closes first - TaskQueue.of(loop).post(lambda held=held: None) + TaskQueue.of(loop).post(lambda _held=held: None) loop.close() refs.append((weakref.ref(loop), weakref.ref(held))) del loop, held diff --git a/tests/test_track_settings.py b/tests/test_track_settings.py index a7c08e1..3e87551 100644 --- a/tests/test_track_settings.py +++ b/tests/test_track_settings.py @@ -7,15 +7,17 @@ """Settings, capabilities, constraints and content hints of tracks.""" +from __future__ import annotations + import pytest import webrtc -from tests.helpers import connect_track, run_isolated, wait_until +from tests.helpers import capture_mode, connect_track, run_isolated, wait_until @pytest.mark.asyncio -async def test_camera_settings_and_capabilities(): - """A camera track has the size and measured frame rate of its frames, and the capabilities of the camera""" +async def test_camera_settings_and_capabilities() -> None: + """A camera track has the size and measured frame rate of its frames, and the capabilities of the camera.""" stream = webrtc.get_user_media(audio=False, video=True, width=320, height=240, frame_rate=30) track = stream.get_tracks()[0] @@ -23,7 +25,8 @@ async def test_camera_settings_and_capabilities(): settings = track.get_settings() assert (settings.width, settings.height, settings.aspect_ratio) == (320, 240, 320 / 240) assert abs(settings.frame_rate - 30) < 5 - assert settings.device_id == 'synthetic-camera' and settings.resize_mode == 'none' + assert settings.device_id == 'synthetic-camera' + assert settings.resize_mode == 'none' capabilities = track.get_capabilities() assert capabilities.width == webrtc.ULongRange(1, 4096) @@ -33,20 +36,21 @@ async def test_camera_settings_and_capabilities(): @pytest.mark.asyncio -async def test_microphone_settings(audio_stream): - """A microphone track has the format of its samples, and no constraints""" +async def test_microphone_settings(audio_stream: webrtc.MediaStream) -> None: + """A microphone track has the format of its samples, and no constraints.""" track = audio_stream.get_tracks()[0] await wait_until(lambda: track.get_settings().sample_rate is not None, 'the audio format') settings = track.get_settings() assert (settings.sample_rate, settings.channel_count, settings.sample_size) == (48000, 1, 16) - assert settings.echo_cancellation is False and settings.device_id == 'synthetic-microphone' + assert settings.echo_cancellation is False + assert settings.device_id == 'synthetic-microphone' assert track.get_capabilities().sample_rate == webrtc.ULongRange(48000, 48000) assert track.get_constraints() == webrtc.MediaTrackConstraints() @pytest.mark.asyncio -async def test_apply_constraints_to_the_camera(video_stream): - """Constraints change the size and frame rate of the camera, and are kept by the track""" +async def test_apply_constraints_to_the_camera(video_stream: webrtc.MediaStream) -> None: + """Constraints change the size and frame rate of the camera, and are kept by the track.""" track = video_stream.get_tracks()[0] await track.apply_constraints({'width': 160, 'height': {'exact': 120}, 'frameRate': {'max': 10}}) await wait_until(lambda: track.get_settings().width == 160, 'the new size') @@ -60,8 +64,8 @@ async def test_apply_constraints_to_the_camera(video_stream): @pytest.mark.asyncio -async def test_overconstrained(video_stream, audio_stream): - """A required constraint the source can't satisfy fails, leaving the track as it was""" +async def test_overconstrained(video_stream: webrtc.MediaStream, audio_stream: webrtc.MediaStream) -> None: + """A required constraint the source can't satisfy fails, leaving the track as it was.""" video = video_stream.get_tracks()[0] await video.apply_constraints({'width': 320}) with pytest.raises(webrtc.OverconstrainedError) as error: @@ -79,8 +83,10 @@ async def test_overconstrained(video_stream, audio_stream): @pytest.mark.asyncio -async def test_remote_track_settings(caller, callee, video_stream): - """A remote track has the settings of what arrives, and no capabilities""" +async def test_remote_track_settings( + caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection, video_stream: webrtc.MediaStream +) -> None: + """A remote track has the settings of what arrives, and no capabilities.""" remote = await connect_track(caller, callee, video_stream.get_tracks()[0]) await wait_until(lambda: remote.get_settings().width is not None, 'frames') settings = remote.get_settings() @@ -91,23 +97,23 @@ async def test_remote_track_settings(caller, callee, video_stream): await remote.apply_constraints({'width': {'exact': 100}}) -@pytest.mark.asyncio -async def test_content_hint(audio_stream, video_stream): - """Hints of the kind of the track are kept, others are ignored""" +def test_content_hint(audio_stream: webrtc.MediaStream, video_stream: webrtc.MediaStream) -> None: + """Hints of the kind of the track are kept, others are ignored.""" video, audio = video_stream.get_tracks()[0], audio_stream.get_tracks()[0] - assert video.content_hint == audio.contentHint == '' + assert not video.content_hint + assert not audio.contentHint video.content_hint = 'text' audio.content_hint = 'music' video.content_hint = 'speech' audio.content_hint = 'detail' assert (video.content_hint, audio.content_hint) == ('text', 'music') video.content_hint = '' - assert video.content_hint == '' + assert not video.content_hint @pytest.mark.asyncio -async def test_constraints_of_an_ended_track(video_stream): - """Constraints of an ended track are accepted, even ones it couldn't satisfy""" +async def test_constraints_of_an_ended_track(video_stream: webrtc.MediaStream) -> None: + """Constraints of an ended track are accepted, even ones it couldn't satisfy.""" track = video_stream.get_tracks()[0] track.stop() await track.apply_constraints({'width': {'exact': 100000}}) @@ -118,8 +124,10 @@ async def test_constraints_of_an_ended_track(video_stream): 'constraints', [{'frame_rate': float('nan')}, {'frame_rate': float('inf')}, {'width': '1'}], ) -async def test_constraints_have_their_webidl_types(video_stream, constraints): - """Unsigned longs and restricted doubles: other values are a TypeError""" +async def test_constraints_have_their_webidl_types( + video_stream: webrtc.MediaStream, constraints: dict[str, object] +) -> None: + """Unsigned longs and restricted doubles: other values are a TypeError.""" with pytest.raises(TypeError): await video_stream.get_tracks()[0].apply_constraints(constraints) with pytest.raises(TypeError): @@ -128,33 +136,35 @@ async def test_constraints_have_their_webidl_types(video_stream, constraints): @pytest.mark.asyncio @pytest.mark.parametrize( - 'constraints, expected', + ('constraints', 'expected'), [ ({'frame_rate': 10**9}, (640, 480, 120)), ({'frame_rate': 0}, (640, 480, 1)), ({'frame_rate': {'ideal': -5}}, (640, 480, 1)), ], ) -async def test_camera_stays_within_its_capabilities(video_stream, constraints, expected): - """Ideal values beyond the capabilities select the nearest ones""" +async def test_camera_stays_within_its_capabilities( + video_stream: webrtc.MediaStream, constraints: dict[str, object], expected: tuple[int, int, float] +) -> None: + """Ideal values beyond the capabilities select the nearest ones.""" track = video_stream.get_tracks()[0] await track.apply_constraints(constraints) - assert track._native_obj._camera() == expected + assert capture_mode(track) == expected track = webrtc.get_user_media(audio=False, video=True, **constraints).get_tracks()[0] - assert track._native_obj._camera() == expected + assert capture_mode(track) == expected track.stop() -def test_get_user_media_rejects_what_the_camera_cannot_do(): +def test_get_user_media_rejects_what_the_camera_cannot_do() -> None: with pytest.raises(webrtc.OverconstrainedError): webrtc.get_user_media(audio=False, video=True, width={'exact': 5000}) with pytest.raises(webrtc.OverconstrainedError): webrtc.get_user_media(audio=False, video=True, frame_rate={'min': 500}) -def test_camera_of_impossible_sizes(): - """A camera of no size, a negative one or a huge one aborted the process""" +def test_camera_of_impossible_sizes() -> None: + """A camera of no size, a negative one or a huge one aborted the process.""" output = run_isolated( """ import asyncio diff --git a/tests/test_tracks.py b/tests/test_tracks.py index e76f3d4..fbd009d 100644 --- a/tests/test_tracks.py +++ b/tests/test_tracks.py @@ -7,6 +7,8 @@ """Remote tracks: their ids and labels, their events and streams, and ending with their transceiver.""" +from __future__ import annotations + import asyncio import pytest @@ -16,8 +18,10 @@ @pytest.mark.asyncio -async def test_remote_tracks_have_their_own_id_and_label(caller, callee, callee2): - """Every connection receiving the same description has other tracks, labeled by their kind""" +async def test_remote_tracks_have_their_own_id_and_label( + caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection, callee2: webrtc.RTCPeerConnection +) -> None: + """Every connection receiving the same description has other tracks, labeled by their kind.""" caller.add_transceiver(webrtc.MediaType.audio) caller.add_transceiver(webrtc.MediaType.video) await caller.set_local_description() @@ -33,8 +37,8 @@ async def test_remote_tracks_have_their_own_id_and_label(caller, callee, callee2 @pytest.mark.asyncio -async def test_stopped_transceiver_ends_the_track_with_its_event(pc): - """The track of a stopped transceiver ends when its ended event is delivered""" +async def test_stopped_transceiver_ends_the_track_with_its_event(pc: webrtc.RTCPeerConnection) -> None: + """The track of a stopped transceiver ends when its ended event is delivered.""" transceiver = pc.add_transceiver(webrtc.MediaType.audio) track = transceiver.receiver.track ended = wait_for_event(track, 'ended') @@ -46,8 +50,10 @@ async def test_stopped_transceiver_ends_the_track_with_its_event(pc): @pytest.mark.asyncio -async def test_rollback_ends_the_track_of_a_removed_transceiver(caller, callee): - """A rollback ends the track of the transceiver it removes, with its event, even if first used after that""" +async def test_rollback_ends_the_track_of_a_removed_transceiver( + caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection +) -> None: + """A rollback ends the track of the transceiver it removes, with its event, even if first used after that.""" caller.add_transceiver(webrtc.MediaType.audio) await exchange_offer(caller, callee) [transceiver] = callee.get_transceivers() @@ -59,21 +65,24 @@ async def test_rollback_ends_the_track_of_a_removed_transceiver(caller, callee): assert track.ready_state == webrtc.MediaStreamTrackState.ended -def test_remove_track_of_another_connection(caller, callee): - """A connection can't remove a sender of another one""" +def test_remove_track_of_another_connection(caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection) -> None: + """A connection can't remove a sender of another one.""" sender = callee.add_transceiver(webrtc.MediaType.audio).sender with pytest.raises(webrtc.InvalidAccessError): caller.remove_track(sender) @pytest.mark.asyncio -async def test_track_event_when_remote_streams_change(caller, callee): - """A track event fires again for the same track when its remote streams change, before the description is set""" +async def test_track_event_when_remote_streams_change( + caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection +) -> None: + """A track event fires again for the same track when its remote streams change, before the description is set.""" events = [] callee.on('track', events.append) transceiver = caller.add_transceiver(webrtc.MediaType.audio) await exchange_offer_answer(caller, callee) - assert len(events) == 1 and events[0].streams == [] + assert len(events) == 1 + assert events[0].streams == [] stream = webrtc.MediaStream() transceiver.sender.set_streams(stream) @@ -84,8 +93,10 @@ async def test_track_event_when_remote_streams_change(caller, callee): @pytest.mark.asyncio -async def test_replace_track_is_chained_after_remove_track(pc, audio_stream, audio_stream2): - """replace_track is an operation of the connection: called before remove_track, it replaces the track after it""" +async def test_replace_track_is_chained_after_remove_track( + pc: webrtc.RTCPeerConnection, audio_stream: webrtc.MediaStream, audio_stream2: webrtc.MediaStream +) -> None: + """replace_track is an operation of the connection: called before remove_track, it replaces the track after it.""" (first,), (second,) = audio_stream.get_tracks(), audio_stream2.get_tracks() sender = pc.add_track(first) replaced = asyncio.ensure_future(sender.replace_track(second)) @@ -96,8 +107,14 @@ async def test_replace_track_is_chained_after_remove_track(pc, audio_stream, aud @pytest.mark.asyncio -async def test_remote_track_mute_and_stream_events(caller, callee, audio_stream, video_stream): - """A remote track that no longer receives media is removed from its stream and muted""" +async def test_remote_track_mute_and_stream_events( + *, + caller: webrtc.RTCPeerConnection, + callee: webrtc.RTCPeerConnection, + audio_stream: webrtc.MediaStream, + video_stream: webrtc.MediaStream, +) -> None: + """A remote track that no longer receives media is removed from its stream and muted.""" (audio,), (video,) = audio_stream.get_tracks(), video_stream.get_tracks() stream = webrtc.MediaStream([audio, video]) caller.add_track(audio, stream) @@ -109,7 +126,7 @@ async def test_remote_track_mute_and_stream_events(caller, callee, audio_stream, remote_stream = events[0].streams[0] assert len(remote_stream.get_tracks()) == 2 removed = wait_for_event(remote_stream, 'removetrack') - remote_video = [e.track for e in events if e.track.kind == webrtc.MediaType.video][0] + remote_video = next(e.track for e in events if e.track.kind == webrtc.MediaType.video) # a track that is muted already doesn't fire mute await wait_until_unmuted(remote_video) muted = wait_for_event(remote_video, 'mute') diff --git a/tests/test_video.py b/tests/test_video.py index f4e002d..0d6fbea 100644 --- a/tests/test_video.py +++ b/tests/test_video.py @@ -7,14 +7,16 @@ """Video: the synthetic camera of get_user_media, video sources and remote video tracks.""" +from __future__ import annotations + import pytest import webrtc -from tests.helpers import connect, wait_for_event +from tests.helpers import capture_mode, connect, wait_for_event -def test_get_user_media_video(): - """A stream of video only has one video track""" +def test_get_user_media_video() -> None: + """A stream of video only has one video track.""" stream = webrtc.get_user_media(audio=False, video=True, width=320, height=240) (track,) = stream.get_tracks() assert track.kind == webrtc.MediaType.video @@ -22,14 +24,14 @@ def test_get_user_media_video(): track.stop() -def test_get_user_media_needs_audio_or_video(): - """A stream of nothing isn't a request""" +def test_get_user_media_needs_audio_or_video() -> None: + """A stream of nothing isn't a request.""" with pytest.raises(TypeError): webrtc.get_user_media(audio=False, video=False) @pytest.mark.parametrize( - 'constraints, error', + ('constraints', 'error'), [ ({'width': {'exact': 0}}, webrtc.OverconstrainedError), ({'frame_rate': {'max': 0}}, webrtc.OverconstrainedError), @@ -37,21 +39,21 @@ def test_get_user_media_needs_audio_or_video(): ], ids=['exact', 'max', 'negative'], ) -def test_get_user_media_constraint_beyond_the_camera(constraints, error): - """A required value the camera can't have is overconstrained, a negative size isn't an unsigned long""" +def test_get_user_media_constraint_beyond_the_camera(constraints: dict[str, object], error: type[Exception]) -> None: + """A required value the camera can't have is overconstrained, a negative size isn't an unsigned long.""" with pytest.raises(error): webrtc.get_user_media(audio=False, video=True, **constraints) -def test_get_user_media_ideal_beyond_the_camera(): - """An ideal value selects the nearest one the camera can have""" +def test_get_user_media_ideal_beyond_the_camera() -> None: + """An ideal value selects the nearest one the camera can have.""" (track,) = webrtc.get_user_media(audio=False, video=True, height={'ideal': 0}).get_tracks() - assert track._native_obj._camera() == (640, 1, 30) + assert capture_mode(track) == (640, 1, 30) track.stop() -def test_get_user_media_constraints(): - """Constraints that select a positive value are accepted""" +def test_get_user_media_constraints() -> None: + """Constraints that select a positive value are accepted.""" stream = webrtc.get_user_media( audio=False, video=True, width={'ideal': 320}, height={'min': 100, 'max': 240}, frame_rate={'exact': 15} ) @@ -59,19 +61,22 @@ def test_get_user_media_constraints(): track.stop() -def test_media_stream_constructor(audio_stream, video_stream): - """A stream is created from tracks or from another stream, which it copies, or empty""" +def test_media_stream_constructor(audio_stream: webrtc.MediaStream, video_stream: webrtc.MediaStream) -> None: + """A stream is created from tracks or from another stream, which it copies, or empty.""" tracks = [*audio_stream.get_tracks(), *video_stream.get_tracks()] stream = webrtc.MediaStream(tracks) assert len(stream.get_tracks()) == 2 copy = webrtc.MediaStream(stream) - assert copy.id != stream.id and len(copy.get_tracks()) == 2 + assert copy.id != stream.id + assert len(copy.get_tracks()) == 2 assert webrtc.MediaStream().get_tracks() == [] @pytest.mark.asyncio -async def test_remote_video_track(caller, callee, video_stream): - """The remote end of a video track has its kind and stream""" +async def test_remote_video_track( + caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection, video_stream: webrtc.MediaStream +) -> None: + """The remote end of a video track has its kind and stream.""" caller.add_track(video_stream.get_tracks()[0], video_stream) track_event = wait_for_event(callee, 'track') await connect(caller, callee) diff --git a/tests/test_video_frame.py b/tests/test_video_frame.py index ac0a992..2ba665c 100644 --- a/tests/test_video_frame.py +++ b/tests/test_video_frame.py @@ -7,6 +7,8 @@ """VideoFrame of WebCodecs: construction, copies, conversions and lifetime.""" +from __future__ import annotations + import gc import math import struct @@ -20,12 +22,12 @@ I420_DATA = bytes(range(1, 13)) -def i420_4x2(data=I420_DATA, **init): +def i420_4x2(data: bytes = I420_DATA, **init: object) -> webrtc.VideoFrame: return webrtc.VideoFrame(data, **{'format': 'I420', 'coded_width': 4, 'coded_height': 2, 'timestamp': 0, **init}) -def test_construct_from_buffer(): - """A frame has the attributes of its init, and defaults for the rest""" +def test_construct_from_buffer() -> None: + """A frame has the attributes of its init, and defaults for the rest.""" frame = i420_4x2(duration=15) assert frame.format == VideoPixelFormat.I420 assert (frame.coded_width, frame.coded_height) == (4, 2) @@ -33,13 +35,14 @@ def test_construct_from_buffer(): assert frame.coded_rect == frame.visible_rect assert (frame.display_width, frame.display_height) == (4, 2) assert (frame.timestamp, frame.duration) == (0, 15) - assert frame.color_space == webrtc.VideoColorSpace('bt709', 'bt709', 'bt709', False) - assert frame.codedWidth == frame.coded_width and frame.allocationSize() == 12 + assert frame.color_space == webrtc.VideoColorSpace('bt709', 'bt709', 'bt709', full_range=False) + assert frame.codedWidth == frame.coded_width + assert frame.allocationSize() == 12 frame.close() -def test_init_as_dataclass_or_dictionary(): - """The init is a dataclass, a dictionary with camelCase names, or keyword arguments""" +def test_init_as_dataclass_or_dictionary() -> None: + """The init is a dataclass, a dictionary with camelCase names, or keyword arguments.""" init = webrtc.VideoFrameBufferInit(format=VideoPixelFormat.I420, coded_width=4, coded_height=2, timestamp=7) for frame in ( webrtc.VideoFrame(I420_DATA, init), @@ -61,14 +64,14 @@ def test_init_as_dataclass_or_dictionary(): {'format': 'I420', 'coded_width': 4, 'coded_height': 2}, ], ) -def test_invalid_init(init): - """An invalid init, or a rect that isn't aligned to the chroma planes, is a TypeError""" +def test_invalid_init(init: dict[str, object]) -> None: + """An invalid init, or a rect that isn't aligned to the chroma planes, is a TypeError.""" with pytest.raises(TypeError): webrtc.VideoFrame(I420_DATA, **init) -def test_buffer_too_small(): - """The buffer must hold the frame at its layout""" +def test_buffer_too_small() -> None: + """The buffer must hold the frame at its layout.""" with pytest.raises(TypeError): i420_4x2(I420_DATA[:11]) with pytest.raises(TypeError): @@ -77,8 +80,8 @@ def test_buffer_too_small(): @pytest.mark.asyncio -async def test_buffer_is_copied(): - """Changing the buffer later doesn't change the frame""" +async def test_buffer_is_copied() -> None: + """Changing the buffer later doesn't change the frame.""" data = bytearray(I420_DATA) frame = i420_4x2(data) data[0] = 99 @@ -89,8 +92,8 @@ async def test_buffer_is_copied(): @pytest.mark.asyncio -async def test_copy_to_layouts(): - """copyTo writes the planes at the layout asked for, and returns it""" +async def test_copy_to_layouts() -> None: + """CopyTo writes the planes at the layout asked for, and returns it.""" frame = i420_4x2() out = bytearray(12) assert await frame.copy_to(out) == [PlaneLayout(0, 4), PlaneLayout(8, 2), PlaneLayout(10, 2)] @@ -105,8 +108,8 @@ async def test_copy_to_layouts(): @pytest.mark.asyncio -async def test_copy_to_rect(): - """A rect copies part of the frame, aligned to the chroma planes""" +async def test_copy_to_rect() -> None: + """A rect copies part of the frame, aligned to the chroma planes.""" frame = i420_4x2() options = {'rect': {'x': 2, 'y': 0, 'width': 2, 'height': 2}} out = bytearray(frame.allocation_size(options)) @@ -118,8 +121,8 @@ async def test_copy_to_rect(): @pytest.mark.asyncio -async def test_copy_to_errors(): - """A small buffer and a layout of the wrong number of planes are TypeErrors, other formats aren't supported""" +async def test_copy_to_errors() -> None: + """A small buffer and a layout of the wrong number of planes are TypeErrors, other formats aren't supported.""" frame = i420_4x2() with pytest.raises(TypeError): await frame.copy_to(bytearray(11)) @@ -135,8 +138,8 @@ async def test_copy_to_errors(): @pytest.mark.asyncio @pytest.mark.parametrize('format', ['RGBA', 'RGBX', 'BGRA', 'BGRX']) -async def test_convert_i420_to_rgb(format): - """copyTo converts YUV to the RGB formats, with the matrix and range of the frame""" +async def test_convert_i420_to_rgb(format: str) -> None: + """CopyTo converts YUV to the RGB formats, with the matrix and range of the frame.""" # pure red in BT.601 limited range: Y 81, U 90, V 240 data = bytes([81] * 16 + [90] * 4 + [240] * 4) frame = webrtc.VideoFrame( @@ -151,15 +154,19 @@ async def test_convert_i420_to_rgb(format): assert len(out) == 64 assert await frame.copy_to(out, {'format': format}) == [PlaneLayout(0, 16)] r, g, b, a = out[:4] if format.startswith('RGB') else (out[2], out[1], out[0], out[3]) - assert r > 245 and g < 10 and b < 10 and a == 255 + assert r > 245 + assert g < 10 + assert b < 10 + assert a == 255 frame.close() @pytest.mark.asyncio -async def test_rgb_formats_swap_and_alpha(): - """RGBA converts to BGRA by swapping R and B, keeping alpha, and to RGBX without it""" +async def test_rgb_formats_swap_and_alpha() -> None: + """RGBA converts to BGRA by swapping R and B, keeping alpha, and to RGBX without it.""" frame = webrtc.VideoFrame(bytes([1, 2, 3, 4] * 4), format='RGBA', coded_width=2, coded_height=2, timestamp=0) - assert frame.color_space.matrix == 'rgb' and frame.color_space.full_range + assert frame.color_space.matrix == 'rgb' + assert frame.color_space.full_range out = bytearray(16) await frame.copy_to(out, {'format': 'BGRA'}) assert list(out[:4]) == [3, 2, 1, 4] @@ -170,7 +177,7 @@ async def test_rgb_formats_swap_and_alpha(): @pytest.mark.asyncio @pytest.mark.parametrize( - 'format, size', + ('format', 'size'), [ ('I420P10', 24), ('I420A', 20), @@ -180,8 +187,8 @@ async def test_rgb_formats_swap_and_alpha(): ('NV12', 12), ], ) -async def test_other_formats_round_trip(format, size): - """Every planar format is kept as it is, and converts to RGBA""" +async def test_other_formats_round_trip(format: str, size: int) -> None: + """Every planar format is kept as it is, and converts to RGBA.""" data = bytes(i % 200 for i in range(size)) frame = webrtc.VideoFrame(data, format=format, coded_width=4, coded_height=2, timestamp=0) assert frame.allocation_size() == size @@ -193,8 +200,8 @@ async def test_other_formats_round_trip(format, size): frame.close() -def test_high_bit_depth_samples_are_little_endian_16_bit(): - """P10 formats have 2 bytes a sample""" +def test_high_bit_depth_samples_are_little_endian_16_bit() -> None: + """P10 formats have 2 bytes a sample.""" y = struct.pack('<8H', *[1023] * 8) uv = struct.pack('<4H', *[512] * 4) frame = webrtc.VideoFrame(y + uv, format='I420P10', coded_width=4, coded_height=2, timestamp=0) @@ -202,8 +209,8 @@ def test_high_bit_depth_samples_are_little_endian_16_bit(): frame.close() -def test_frame_from_frame(): - """A frame from another one shares its pixels, with a visible rect, display size, timestamp or alpha of its own""" +def test_frame_from_frame() -> None: + """A frame from another one shares its pixels, with a visible rect, display size, timestamp or alpha of its own.""" frame = i420_4x2(timestamp=1234, display_width=8, display_height=2) crop = webrtc.VideoFrame(frame, visible_rect={'x': 2, 'y': 0, 'width': 2, 'height': 2}) assert (crop.coded_width, crop.visible_rect.x, crop.visible_rect.width) == (4, 2, 2) @@ -220,8 +227,8 @@ def test_frame_from_frame(): f.close() -def test_rotation_and_flip(): - """Rotations are rounded to a multiple of 90, and combine with the flip of the frame they're added to""" +def test_rotation_and_flip() -> None: + """Rotations are rounded to a multiple of 90, and combine with the flip of the frame they're added to.""" frame = webrtc.VideoFrame(bytes(32), format='RGBX', coded_width=4, coded_height=2, timestamp=0, rotation=-315) assert frame.rotation == 90 assert (frame.display_width, frame.display_height) == (2, 4) @@ -234,23 +241,25 @@ def test_rotation_and_flip(): @pytest.mark.parametrize('rotation', [math.inf, -math.inf, math.nan]) -def test_rotation_must_be_finite(rotation): - """A rotation is a WebIDL double: non-finite values raise TypeError, they raised OverflowError (found by fuzzing)""" +def test_rotation_must_be_finite(rotation: float) -> None: + """A rotation is a WebIDL double: non-finite values are a TypeError, not an OverflowError (found by fuzzing).""" with pytest.raises(TypeError): webrtc.VideoFrame(bytes(32), format='RGBX', coded_width=4, coded_height=2, timestamp=0, rotation=rotation) - with webrtc.VideoFrame(bytes(32), format='RGBX', coded_width=4, coded_height=2, timestamp=0) as frame: - with pytest.raises(TypeError): - webrtc.VideoFrame(frame, rotation=rotation) + frame = webrtc.VideoFrame(bytes(32), format='RGBX', coded_width=4, coded_height=2, timestamp=0) + with frame, pytest.raises(TypeError): + webrtc.VideoFrame(frame, rotation=rotation) @pytest.mark.asyncio -async def test_close_and_clone(): - """A closed frame has no pixels, a clone is closed separately""" +async def test_close_and_clone() -> None: + """A closed frame has no pixels, a clone is closed separately.""" frame = i420_4x2() clone = frame.clone() frame.close() frame.close() - assert frame.format is None and frame.coded_width == 0 and frame.visible_rect is None + assert frame.format is None + assert frame.coded_width == 0 + assert frame.visible_rect is None assert frame.timestamp == 0 with pytest.raises(webrtc.InvalidStateError): frame.allocation_size() @@ -266,16 +275,20 @@ async def test_close_and_clone(): assert clone.format is None -def test_unclosed_frame_warns(): - """A frame garbage collected without being closed warns""" +def drop_unclosed_frame() -> None: + i420_4x2() + gc.collect() + + +def test_unclosed_frame_warns() -> None: + """A frame garbage collected without being closed warns.""" with pytest.warns(ResourceWarning): - i420_4x2() - gc.collect() + drop_unclosed_frame() @pytest.mark.asyncio -async def test_visible_rect_of_a_buffer_is_the_frame(): - """A frame created from a buffer keeps its visible rect only, which becomes the whole frame""" +async def test_visible_rect_of_a_buffer_is_the_frame() -> None: + """A frame created from a buffer keeps its visible rect only, which becomes the whole frame.""" frame = i420_4x2(visible_rect={'x': 2, 'y': 0, 'width': 2, 'height': 2}) assert (frame.coded_width, frame.coded_height) == (2, 2) assert frame.visible_rect == webrtc.DOMRectReadOnly(0, 0, 2, 2) diff --git a/tests/wpt/__init__.py b/tests/wpt/__init__.py index e69de29..f13f15d 100644 --- a/tests/wpt/__init__.py +++ b/tests/wpt/__init__.py @@ -0,0 +1,6 @@ +# +# Copyright 2026 Ilya (Marshal) . All rights reserved. +# +# Use of this source code is governed by a BSD-style license +# that can be found in the LICENSE.md file in the root of the project. +# diff --git a/tests/wpt/__main__.py b/tests/wpt/__main__.py index 7174aa1..9b11ca8 100644 --- a/tests/wpt/__main__.py +++ b/tests/wpt/__main__.py @@ -15,7 +15,9 @@ from __future__ import annotations import argparse +import logging import os +import sys from collections import Counter from concurrent.futures import ThreadPoolExecutor @@ -23,32 +25,35 @@ from tests.wpt.expectations import Expectations from tests.wpt.loader import discover +logger = logging.getLogger(__name__) -def _print_result(case: str, result: dict): + +def _print_result(case: str, result: runner.CaseResult) -> None: harness = result['harness'] - print(f'{case}: harness {harness["status"]}' + (f' ({harness["message"]})' if harness['message'] else '')) + message = f' ({harness["message"]})' if harness['message'] else '' + logger.info('%s: harness %s%s', case, harness['status'], message) for test in result['tests']: - print(f' {test["status"]:<8} {test["name"]}') + logger.info(' %-8s %s', test['status'], test['name']) if test['status'] != 'PASS' and test['message']: - print(f' {test["message"]}') + logger.info(' %s', test['message']) if result['unsupported']: - print(' unsupported: ' + ', '.join(result['unsupported'])) + logger.info(' unsupported: %s', ', '.join(result['unsupported'])) -def run(args): +def run(args: argparse.Namespace) -> None: for case in args.cases: _print_result(case, runner.run(case)) -def update(args): +def update(args: argparse.Namespace) -> None: expectations = Expectations.load() cases = [c for c in args.cases or discover() if not expectations.skip_reason(c)] - tests = Counter() - harness = Counter() - unsupported = Counter() + tests: Counter[str] = Counter() + harness: Counter[str] = Counter() + unsupported: Counter[str] = Counter() - def run_repeatedly(case): + def run_repeatedly(case: str) -> list[runner.CaseResult]: return [runner.run(case) for _ in range(args.repeat)] with ThreadPoolExecutor(args.jobs) as pool: @@ -59,19 +64,19 @@ def run_repeatedly(case): tests.update(t['status'] for t in first['tests']) unsupported.update(first['unsupported']) statuses = '/'.join(sorted({r['harness']['status'] for r in results})) - print(f'[{done}/{len(cases)}] {statuses:<8} {case}', flush=True) + logger.info('[%d/%d] %-8s %s', done, len(cases), statuses, case) expectations.save() - print(f'\nfiles: {dict(harness)}') - print(f'tests: {dict(tests)}') + logger.info('\nfiles: %s', dict(harness)) + logger.info('tests: %s', dict(tests)) if unsupported: - print('unsupported members used by tests:') + logger.info('unsupported members used by tests:') for name, count in unsupported.most_common(): - print(f' {count:>4} {name}') + logger.info(' %4d %s', count, name) -def main(): +def main() -> None: parser = argparse.ArgumentParser(prog='python -m tests.wpt') commands = parser.add_subparsers(required=True) @@ -85,6 +90,7 @@ def main(): update_parser.add_argument('--repeat', type=int, default=1, help='runs of each case, to find flaky tests') update_parser.set_defaults(func=update) + logging.basicConfig(format='%(message)s', level=logging.INFO, stream=sys.stdout) args = parser.parse_args() args.func(args) diff --git a/tests/wpt/bridge.py b/tests/wpt/bridge.py index 450f53a..b5abaee 100644 --- a/tests/wpt/bridge.py +++ b/tests/wpt/bridge.py @@ -5,22 +5,36 @@ # that can be found in the LICENSE.md file in the root of the project. # -"""Python side of the WPT shim: does what shim.js asks on `webrtc` objects, returning {'ok': value} or -{'error': {...}} so exceptions reach JS with their type.""" +"""Python side of the WPT shim: does what shim.js asks on `webrtc` objects. + +It returns {'ok': value} or {'error': {...}}, so exceptions reach JS with their type. +""" + +from __future__ import annotations import asyncio import dataclasses import enum import sys import time +from typing import TYPE_CHECKING, Callable, Union import pythonmonkey as pm import webrtc import webrtc.enums +if TYPE_CHECKING: + from collections.abc import Coroutine + +Result = dict[str, object] +Buffer = Union[bytes, bytearray, memoryview] + # The loop of the test, set by the runner: code called from JS may not see it as the running loop -LOOP = None +LOOP: asyncio.AbstractEventLoop | None = None + +# what the shim turns into JS errors of their type, anything else reaches JS as a plain Error +BRIDGED_ERRORS = (TypeError, ValueError, OverflowError, webrtc.PythonWebRTCExceptionBase) # Python objects without a native object, exposed to JS as interfaces _PLAIN_INTERFACES = ( @@ -36,92 +50,105 @@ ) -def _camel_case(name): +def _camel_case(name: str) -> str: first, *rest = name.split('_') return first + ''.join(part.title() for part in rest) -def to_js(value): - if isinstance(value, enum.Enum): - # the values of the enums are the WebIDL ones - return value.value - if isinstance(value, webrtc.DOMRectReadOnly): - return {'__rect': [value.x, value.y, value.width, value.height]} - if isinstance(value, webrtc.Blob): - return {'__blob': bytearray(bytes(value)), 'type': value.type} - if isinstance(value, webrtc.RTCSessionDescriptionInit): - # a dictionary in WebIDL, but a WebRTCObject (it holds a native one) here, not a dataclass - return value.to_json() - if isinstance(value, webrtc.WebRTCObject): - # Wrappers are created per access, but the native object is shared, so its id identifies the WebRTC object - return {'__type': type(value).__name__, '__id': id(value._native_obj), '__obj': value} - if isinstance(value, _PLAIN_INTERFACES): - return {'__type': type(value).__name__, '__id': id(value), '__obj': value} - if isinstance(value, BaseException): - return {'__error': _error(value)['error']} - if isinstance(value, webrtc.RTCStatsReport): - return {'__statsReport': [[stats_id, dict(stats)] for stats_id, stats in value.items()]} - if isinstance(value, bytes): - # PythonMonkey shares a bytearray as a Uint8Array, the shim copies it - return {'__bytes': bytearray(value)} - if isinstance(value, webrtc.Event): - init = {_camel_case(k): to_js(v) for k, v in vars(value).items() if k not in ('type', 'target')} - return {'__event': type(value).__name__, 'type': value.type, 'init': init} +def _event_to_js(event: webrtc.Event) -> dict[str, object]: + init = {_camel_case(k): to_js(v) for k, v in vars(event).items() if k not in {'type', 'target'}} + return {'__event': type(event).__name__, 'type': event.type, 'init': init} + + +def _dictionary_to_js(value: object) -> dict[str, object]: + # a dictionary in WebIDL, where missing members are left out + members = ((f.name, getattr(value, f.name)) for f in dataclasses.fields(value)) + return {_camel_case(name): to_js(member) for name, member in members if member is not None} + + +# in order: the first type that matches converts the value +_CONVERTERS: list[tuple[type | tuple[type, ...], Callable[..., object]]] = [ + # the values of the enums are the WebIDL ones + (enum.Enum, lambda value: value.value), + (webrtc.DOMRectReadOnly, lambda value: {'__rect': [value.x, value.y, value.width, value.height]}), + (webrtc.Blob, lambda value: {'__blob': bytearray(bytes(value)), 'type': value.type}), + # a dictionary in WebIDL, but a WebRTCObject (it holds a native one) here, not a dataclass + (webrtc.RTCSessionDescriptionInit, lambda value: value.to_json()), + # Wrappers are created per access, but they hash as the native object they share, which identifies it + (webrtc.WebRTCObject, lambda value: {'__type': type(value).__name__, '__id': hash(value), '__obj': value}), + (_PLAIN_INTERFACES, lambda value: {'__type': type(value).__name__, '__id': id(value), '__obj': value}), + (BaseException, lambda value: {'__error': _error(value)['error']}), + ( + webrtc.RTCStatsReport, + lambda value: {'__statsReport': [[stats_id, dict(stats)] for stats_id, stats in value.items()]}, + ), + # PythonMonkey shares a bytearray as a Uint8Array, the shim copies it + (bytes, lambda value: {'__bytes': bytearray(value)}), + (webrtc.Event, _event_to_js), + (dict, lambda value: {k: to_js(v) for k, v in value.items()}), + ((list, tuple), lambda value: [to_js(v) for v in value]), +] + + +def to_js(value: object) -> object: + """A Python value as the shim reads it.""" + for types, convert in _CONVERTERS: + if isinstance(value, types): + return convert(value) if dataclasses.is_dataclass(value) and not isinstance(value, type): - # a dictionary in WebIDL, where missing members are left out - return { - _camel_case(f.name): to_js(getattr(value, f.name)) - for f in dataclasses.fields(value) - if getattr(value, f.name) is not None - } - if isinstance(value, dict): - return {k: to_js(v) for k, v in value.items()} - if isinstance(value, (list, tuple)): - return [to_js(v) for v in value] + return _dictionary_to_js(value) return value -def _to_enum(value): +def _to_enum(value: dict[str, object]) -> object: # named after its webrtc counterpart, whose values are the WebIDL ones - enum_cls = getattr(webrtc.enums, value['__enum']) + enum_cls = getattr(webrtc.enums, str(value['__enum'])) try: return enum_cls(value['value']) except ValueError: if value.get('strict', True): - raise TypeError(f"'{value['value']}' is not a valid value for enumeration {value['__enum']}") from None + msg = f"'{value['value']}' is not a valid value for enumeration {value['__enum']}" + raise TypeError(msg) from None # A DOMString in WebIDL rather than an enum, so it's up to the library to reject it return value['value'] -def from_js(value): - if value is pm.null: - return None - if isinstance(value, float) and value.is_integer(): - return int(value) +def _object_from_js(value: list[object] | memoryview | dict[str, object]) -> object: if isinstance(value, list): return [from_js(v) for v in value] if isinstance(value, memoryview): # an ArrayBuffer or a view of one return bytes(value) - if isinstance(value, dict): - if '__enum' in value: - return _to_enum(value) - if '__model' in value: - # a WebIDL dictionary the library has a keyword model for - return getattr(webrtc, value['__model'])(**from_js(value['kwargs'])) - return {k: from_js(v) for k, v in value.items()} + return _dict_from_js(value) + + +def _dict_from_js(value: dict[str, object]) -> object: + if '__enum' in value: + return _to_enum(value) + if '__model' in value: + # a WebIDL dictionary the library has a keyword model for + return getattr(webrtc, str(value['__model']))(**from_js(value['kwargs'])) + return {k: from_js(v) for k, v in value.items()} + + +def from_js(value: object) -> object: + """A value from the shim as Python takes it.""" + if value is pm.null: + return None + if isinstance(value, float) and value.is_integer(): + return int(value) + if isinstance(value, (list, memoryview, dict)): + return _object_from_js(value) return value -def _error(exc): +def _error(exc: BaseException) -> Result: # the most specific class known to the shim - for cls in type(exc).__mro__: - if cls.__module__ in ('webrtc.exceptions', 'wrtc', 'builtins'): - kind = cls.__name__ - break - else: - kind = 'Error' - error = {'kind': kind, 'message': str(exc)} + kind = next( + (cls.__name__ for cls in type(exc).__mro__ if cls.__module__ in {'webrtc.exceptions', 'wrtc', 'builtins'}), + 'Error', + ) + error: dict[str, object] = {'kind': kind, 'message': str(exc)} if isinstance(exc, webrtc.OverconstrainedError): error['constraint'] = exc.constraint if isinstance(exc, webrtc.RTCError): @@ -136,72 +163,74 @@ def _error(exc): return {'error': error} -def _guard(func): +def _guard(func: Callable[[], object]) -> Result: try: return {'ok': to_js(func())} - except Exception as e: + except BRIDGED_ERRORS as e: + return _error(e) + + +async def _guard_async(awaitable: Callable[[], Coroutine[object, object, object]]) -> Result: + try: + return {'ok': to_js(await awaitable())} + except BRIDGED_ERRORS as e: return _error(e) -def get_attr(obj, name): +def _start(coroutine: Coroutine[object, object, Result]) -> asyncio.Future[Result]: + # Started eagerly so the synchronous steps of the method run when it's called, not on the next iteration of + # the loop. Before Python 3.12 they run one iteration later, so results that depend on event order may differ. + if sys.version_info >= (3, 12): + return asyncio.Task(coroutine, loop=LOOP, eager_start=True) + return asyncio.ensure_future(coroutine, loop=LOOP) + + +def get_attr(obj: object, name: str) -> Result: return _guard(lambda: getattr(obj, name)) -def set_attr(obj, name, value): +def set_attr(obj: object, name: str, value: object) -> Result: return _guard(lambda: setattr(obj, name, from_js(value))) -def _call(obj, name, args, kwargs): +def _call(obj: object, name: str, arguments: dict[str, object]) -> object: + """Calls a method with {'args': [...], 'kwargs': {...}} from JS.""" + args, kwargs = arguments['args'], arguments.get('kwargs') return getattr(obj, name)(*from_js(list(args)), **from_js(dict(kwargs or {}))) -def call_method(obj, name, args, kwargs=None): - return _guard(lambda: _call(obj, name, args, kwargs)) +def call_method(obj: object, name: str, arguments: dict[str, object]) -> Result: + return _guard(lambda: _call(obj, name, arguments)) -async def _call_async_method(obj, name, args, kwargs): - try: - return {'ok': to_js(await _call(obj, name, args, kwargs))} - except Exception as e: - return _error(e) +def call_async_method(obj: object, name: str, arguments: dict[str, object]) -> asyncio.Future[Result]: + return _start(_guard_async(lambda: _call(obj, name, arguments))) -def call_async_method(obj, name, args, kwargs=None): - # Started eagerly so the synchronous steps of the method run when it's called, not on the next iteration of - # the loop. Before Python 3.12 they run one iteration later, so results that depend on event order may differ. - coroutine = _call_async_method(obj, name, args, kwargs) - if sys.version_info >= (3, 12): - return asyncio.Task(coroutine, loop=LOOP, eager_start=True) - return asyncio.ensure_future(coroutine, loop=LOOP) +def await_attr(obj: object, name: str) -> asyncio.Future[Result]: + """An attribute that is a future (like the closed promise of a reader), awaited.""" + async def attr() -> object: + return await getattr(obj, name) -async def _await_attr(obj, name): - try: - return {'ok': to_js(await getattr(obj, name))} - except Exception as e: - return _error(e) - + return asyncio.ensure_future(_guard_async(attr), loop=LOOP) -def await_attr(obj, name): - """An attribute that is a future (like the closed promise of a reader), awaited""" - return asyncio.ensure_future(_await_attr(obj, name), loop=LOOP) +def video_frame_copy_to(frame: webrtc.VideoFrame, destination: Buffer, options: object) -> asyncio.Future[Result]: + """Copies a frame into the bytes of a JS buffer, which the shim writes back.""" -def video_frame_copy_to(frame, destination, options): - """Copies a frame into the bytes of a JS buffer, which the shim writes back""" - - def copy(): + async def copy() -> dict[str, object]: data = bytearray(destination) - layout = frame._copy_to(data, from_js(options)) + layout = await frame.copy_to(data, from_js(options)) return {'layout': layout, 'data': bytes(data)} - return _guard(copy) + return _start(_guard_async(copy)) -def audio_data_copy_to(audio, destination, options): - """Copies samples into the bytes of a JS buffer, which the shim writes back""" +def audio_data_copy_to(audio: webrtc.AudioData, destination: Buffer, options: object) -> Result: + """Copies samples into the bytes of a JS buffer, which the shim writes back.""" - def copy(): + def copy() -> bytes: data = bytearray(destination) audio.copy_to(data, from_js(options)) return bytes(data) @@ -209,34 +238,34 @@ def copy(): return _guard(copy) -def construct(name, kwargs): +def construct(name: str, kwargs: dict[str, object]) -> Result: return _guard(lambda: getattr(webrtc, name)(**from_js(dict(kwargs)))) -def get_user_media(kwargs): +def get_user_media(kwargs: dict[str, object]) -> Result: return _guard(lambda: webrtc.get_user_media(**from_js(dict(kwargs)))) -def call_static(class_name, name, args): +def call_static(class_name: str, name: str, args: list[object]) -> Result: return _guard(lambda: getattr(getattr(webrtc, class_name), name)(*from_js(list(args)))) -def call_async_static(class_name, name, args): - return call_async_method(getattr(webrtc, class_name), name, args) +def call_async_static(class_name: str, name: str, args: list[object]) -> asyncio.Future[Result]: + return call_async_method(getattr(webrtc, class_name), name, {'args': args}) -def now(): - """Milliseconds since the epoch, with the precision of the clock of libwebrtc stats""" +def now() -> float: + """Milliseconds since the epoch, with the precision of the clock of libwebrtc stats.""" return time.time() * 1000 -def subscribe(obj, name, callback): - """Delivers the events of a type to a JS callback, which dispatches them to the JS listeners""" +def subscribe(obj: webrtc.EventTarget, name: str, callback: Callable[[object], object]) -> Result: + """Delivers the events of a type to a JS callback, which dispatches them to the JS listeners.""" - def deliver(event): + def deliver(event: webrtc.Event) -> None: callback(to_js(event)) - def add_listener(): + def add_listener() -> None: # nothing to return to JS: on() returns the handler obj.on(name, deliver) diff --git a/tests/wpt/child.py b/tests/wpt/child.py new file mode 100644 index 0000000..b681be5 --- /dev/null +++ b/tests/wpt/child.py @@ -0,0 +1,113 @@ +# +# Copyright 2026 Ilya (Marshal) . All rights reserved. +# +# Use of this source code is governed by a BSD-style license +# that can be found in the LICENSE.md file in the root of the project. +# + +"""Runs one web-platform-tests case in this process and prints its result, for tests.wpt.runner.""" + +from __future__ import annotations + +import asyncio +import json +import logging +import sys + +import pythonmonkey as pm + +from tests.wpt import bridge +from tests.wpt.loader import build_scripts, load, split_case +from tests.wpt.runner import ( + HARNESS_STATUSES, + LONG_TIMEOUT, + RESULT_PREFIX, + TEST_STATUSES, + TIMEOUT, + CaseResult, + harness_result, +) + +logger = logging.getLogger(__name__) + + +def _text(value: object) -> str | None: + # JS null and undefined arrive as PythonMonkey objects + return value if isinstance(value, str) else None + + +def _log_unhandled_rejections(loop: asyncio.AbstractEventLoop, context: dict[str, object]) -> None: + """Logs unhandled rejections instead of PythonMonkey's handler, which stops its timers. + + That would leave the rest of the file hanging. It also reports rejections that get handled later (as testharness + does), so they are only logged. Any other exception goes to the default handler. + """ + if isinstance(context.get('exception'), pm.SpiderMonkeyError): + logger.warning('unhandled rejection: %s', context['exception']) + else: + loop.default_exception_handler(context) + + +async def run_in_process(case: str) -> CaseResult: + path, variant = split_case(case) + test_file = load(path) + + loop = asyncio.get_running_loop() + completed = loop.create_future() + unsupported: set[str] = set() + + def complete(result: dict) -> None: + if not completed.done(): + completed.set_result(result) + + bridge.LOOP = loop + loop.set_exception_handler(_log_unhandled_rejections) + + pm.eval('(env) => { globalThis.__wpt = env; }')({ + 'bridge': bridge.EXPORTS, + 'unsupported': unsupported.add, + 'complete': complete, + }) + pm.eval('(search, pathname) => { globalThis.location = {search, pathname, href: pathname + search}; }')( + variant, '/' + case.partition('?')[0] + ) + + try: + # one after another, without giving control to the loop in between + for script in build_scripts(test_file): + pm.eval(script) + except pm.SpiderMonkeyError as e: + return harness_result('ERROR', str(e)) + + timeout = LONG_TIMEOUT if test_file.long_timeout else TIMEOUT + try: + result = await asyncio.wait_for(asyncio.shield(completed), timeout) + except asyncio.TimeoutError: + # marks unfinished tests as timed out and completes the harness + pm.eval('timeout')() + try: + result = await asyncio.wait_for(completed, 5) + except asyncio.TimeoutError: + return harness_result('TIMEOUT', 'the harness did not complete after timing out') + + return { + 'harness': { + 'status': HARNESS_STATUSES[int(result['harness']['status'])], + 'message': _text(result['harness']['message']), + }, + 'tests': [ + {'name': t['name'], 'status': TEST_STATUSES[int(t['status'])], 'message': _text(t['message'])} + for t in result['tests'] + ], + 'unsupported': sorted(unsupported), + } + + +def main() -> None: + result = asyncio.run(run_in_process(sys.argv[1])) + sys.stdout.write(RESULT_PREFIX + json.dumps(result) + '\n') + sys.stdout.flush() + + +if __name__ == '__main__': + main() diff --git a/tests/wpt/expectations.py b/tests/wpt/expectations.py index 9dde78c..a7315d5 100644 --- a/tests/wpt/expectations.py +++ b/tests/wpt/expectations.py @@ -23,17 +23,21 @@ import json from dataclasses import dataclass, field from pathlib import Path +from typing import TYPE_CHECKING + +if TYPE_CHECKING: + from tests.wpt.runner import CaseResult PATH = Path(__file__).with_name('expectations.json') HARNESS_KEY = '[harness]' -def _allowed(expected) -> list[str]: +def _allowed(expected: str | list[str]) -> list[str]: return expected if isinstance(expected, list) else [expected] -def _describe(expected) -> str: +def _describe(expected: str | list[str]) -> str: return ' or '.join(_allowed(expected)) @@ -49,14 +53,14 @@ def load(cls) -> Expectations: data = json.loads(PATH.read_text()) return cls(skip=data.get('skip', {}), results=data.get('results', {})) - def save(self): + def save(self) -> None: data = {'skip': dict(sorted(self.skip.items())), 'results': dict(sorted(self.results.items()))} PATH.write_text(json.dumps(data, indent=2) + '\n') def skip_reason(self, case: str) -> str | None: return self.skip.get(case.partition('?')[0]) - def record(self, case: str, results: list[dict]): + def record(self, case: str, results: list[CaseResult]) -> None: """Records the statuses seen over one or more runs of a case. Statuses that differ between runs are recorded as a list. An existing list is kept while it still covers @@ -84,7 +88,7 @@ def record(self, case: str, results: list[dict]): else: self.results.pop(case, None) - def mismatches(self, case: str, result: dict) -> list[str]: + def mismatches(self, case: str, result: CaseResult) -> list[str]: expected = dict(self.results.get(case, {})) expected_harness = expected.pop(HARNESS_KEY, 'OK') @@ -102,6 +106,8 @@ def mismatches(self, case: str, result: dict) -> list[str]: if test['status'] not in _allowed(want): problems.append(f'{test["name"]}: expected {_describe(want)}, got {test["status"]}: {test["message"]}') - for name in sorted(expected.keys() - seen): - problems.append(f'{name}: expected {_describe(expected[name])}, but the test did not run') + problems.extend( + f'{name}: expected {_describe(expected[name])}, but the test did not run' + for name in sorted(expected.keys() - seen) + ) return problems diff --git a/tests/wpt/loader.py b/tests/wpt/loader.py index 2470d91..79c2d23 100644 --- a/tests/wpt/loader.py +++ b/tests/wpt/loader.py @@ -71,14 +71,16 @@ class TestFile: class _HtmlCollector(HTMLParser): - def __init__(self, test_file: TestFile): + def __init__(self, test_file: TestFile) -> None: super().__init__() self.test_file = test_file - self._inline = None + self._inline: list[str] | None = None self._in_title = False - def handle_starttag(self, tag, attrs): - attrs = dict(attrs) + def handle_starttag(self, tag: str, attrs: list[tuple[str, str | None]]) -> None: + self._start(tag, dict(attrs)) + + def _start(self, tag: str, attrs: dict[str, str | None]) -> None: if tag == 'script': if attrs.get('src'): self.test_file.scripts.append(('src', attrs['src'])) @@ -91,13 +93,13 @@ def handle_starttag(self, tag, attrs): elif tag == 'title': self._in_title = True - def handle_data(self, data): + def handle_data(self, data: str) -> None: if self._inline is not None: self._inline.append(data) elif self._in_title: self.test_file.title += data - def handle_endtag(self, tag): + def handle_endtag(self, tag: str) -> None: if tag == 'script' and self._inline is not None: self.test_file.scripts.append(('inline', ''.join(self._inline))) self._inline = None @@ -107,14 +109,14 @@ def handle_endtag(self, tag): def _load_html(path: Path) -> TestFile: test_file = TestFile(path) - _HtmlCollector(test_file).feed(path.read_text()) + _HtmlCollector(test_file).feed(path.read_text(encoding='utf-8')) return test_file def _load_js(path: Path) -> TestFile: """Loads a .window.js or .any.js test, which WPT would wrap into a generated HTML page.""" test_file = TestFile(path, scripts=[('src', HARNESS)]) - for line in path.read_text().splitlines(): + for line in path.read_text(encoding='utf-8').splitlines(): match = _META.match(line.strip()) if not match: continue @@ -180,6 +182,9 @@ def build_scripts(test_file: TestFile) -> list[str]: They are evaluated separately, like the scripts of a page (so a "use strict" directive applies to its own script), but all at once: in a shell, the harness considers the page loaded after the first microtask. + + Returns: + The source of each script, in the order of the page. """ parts = [POLYFILLS.read_text(), SHIM.read_text()] test_src = '/' + test_file.path.relative_to(WPT_ROOT).as_posix() diff --git a/tests/wpt/runner.py b/tests/wpt/runner.py index 7985f5d..20a743d 100644 --- a/tests/wpt/runner.py +++ b/tests/wpt/runner.py @@ -8,7 +8,7 @@ """Runs web-platform-tests cases against python-webrtc in PythonMonkey. PythonMonkey has one global object per process and WPT helpers declare top-level constants, so each case runs -in a child process: `python -m tests.wpt.runner ` prints a single ``WPT_RESULT `` line. +in a child process: `python -m tests.wpt.child ` prints a single ``WPT_RESULT `` line. A result looks like: { @@ -20,13 +20,13 @@ from __future__ import annotations -import asyncio import json import subprocess import sys from pathlib import Path +from typing import TypedDict -from tests.wpt.loader import build_scripts, load, split_case +from tests.wpt.loader import load, split_case REPO_ROOT = Path(__file__).resolve().parents[2] @@ -40,109 +40,44 @@ RESULT_PREFIX = 'WPT_RESULT ' -def _harness_result(status: str, message: str) -> dict: - return {'harness': {'status': status, 'message': message}, 'tests': [], 'unsupported': []} - - -def _text(value): - # JS null and undefined arrive as PythonMonkey objects - return value if isinstance(value, str) else None - - -def _log_unhandled_rejections(loop, context): - """PythonMonkey's handler of unhandled rejections stops its timers, which leaves the rest of the file hanging. - It also reports rejections that get handled later (as testharness does), so they are only logged. Any other - exception goes to the default handler.""" - import pythonmonkey as pm - - if isinstance(context.get('exception'), pm.SpiderMonkeyError): - print('unhandled rejection:', context['exception'], file=sys.stderr) - else: - loop.default_exception_handler(context) - - -async def run_in_process(case: str) -> dict: - import pythonmonkey as pm - - from tests.wpt import bridge +class HarnessResult(TypedDict): + status: str + message: str | None - path, variant = split_case(case) - test_file = load(path) - loop = asyncio.get_running_loop() - completed = loop.create_future() - unsupported = set() +class TestResult(TypedDict): + name: str + status: str + message: str | None - def complete(result): - if not completed.done(): - completed.set_result(result) - bridge.LOOP = loop - loop.set_exception_handler(_log_unhandled_rejections) +class CaseResult(TypedDict): + harness: HarnessResult + tests: list[TestResult] + unsupported: list[str] - pm.eval('(env) => { globalThis.__wpt = env; }')( - {'bridge': bridge.EXPORTS, 'unsupported': unsupported.add, 'complete': complete} - ) - pm.eval('(search, pathname) => { globalThis.location = {search, pathname, href: pathname + search}; }')( - variant, '/' + case.partition('?')[0] - ) - try: - # one after another, without giving control to the loop in between - for script in build_scripts(test_file): - pm.eval(script) - except pm.SpiderMonkeyError as e: - return _harness_result('ERROR', str(e)) - - timeout = LONG_TIMEOUT if test_file.long_timeout else TIMEOUT - try: - result = await asyncio.wait_for(asyncio.shield(completed), timeout) - except asyncio.TimeoutError: - # marks unfinished tests as timed out and completes the harness - pm.eval('timeout')() - try: - result = await asyncio.wait_for(completed, 5) - except asyncio.TimeoutError: - return _harness_result('TIMEOUT', 'the harness did not complete after timing out') - - return { - 'harness': { - 'status': HARNESS_STATUSES[int(result['harness']['status'])], - 'message': _text(result['harness']['message']), - }, - 'tests': [ - {'name': t['name'], 'status': TEST_STATUSES[int(t['status'])], 'message': _text(t['message'])} - for t in result['tests'] - ], - 'unsupported': sorted(unsupported), - } +def harness_result(status: str, message: str) -> CaseResult: + return {'harness': {'status': status, 'message': message}, 'tests': [], 'unsupported': []} -def run(case: str) -> dict: +def run(case: str) -> CaseResult: """Runs a case in a child process. A crash or a hang of the child becomes the harness status.""" path, _ = split_case(case) limit = (LONG_TIMEOUT if load(path).long_timeout else TIMEOUT) + 30 try: proc = subprocess.run( - [sys.executable, '-m', 'tests.wpt.runner', case], + [sys.executable, '-m', 'tests.wpt.child', case], cwd=REPO_ROOT, capture_output=True, text=True, timeout=limit, + check=False, ) except subprocess.TimeoutExpired: - return _harness_result('TIMEOUT', f'the runner did not finish in {limit} seconds') + return harness_result('TIMEOUT', f'the runner did not finish in {limit} seconds') for line in proc.stdout.splitlines(): if line.startswith(RESULT_PREFIX): return json.loads(line[len(RESULT_PREFIX) :]) - return _harness_result('CRASH', f'exit code {proc.returncode}\n{proc.stderr[-2000:]}') - - -def main(): - result = asyncio.run(run_in_process(sys.argv[1])) - print(RESULT_PREFIX + json.dumps(result), flush=True) - - -if __name__ == '__main__': - main() + return harness_result('CRASH', f'exit code {proc.returncode}\n{proc.stderr[-2000:]}') diff --git a/tests/wpt/shim.js b/tests/wpt/shim.js index ba44f1c..ceb0375 100644 --- a/tests/wpt/shim.js +++ b/tests/wpt/shim.js @@ -136,10 +136,10 @@ const setAttr = (self, name, value) => unwrap(bridge.set_attr(pyObjects.get(self), name, value)); const callMethod = (self, name, ...args) => callMethodWithKeywords(self, name, args, {}); const callMethodWithKeywords = (self, name, args, kwargs) => - unwrap(bridge.call_method(pyObjects.get(self), name, args.map(toPy), kwargs)); + unwrap(bridge.call_method(pyObjects.get(self), name, {args: args.map(toPy), kwargs})); const callAsyncMethod = async (self, name, ...args) => callAsyncMethodWithKeywords(self, name, args, {}); const callAsyncMethodWithKeywords = async (self, name, args, kwargs) => - unwrap(await bridge.call_async_method(pyObjects.get(self), name, args.map(toPy), kwargs)); + unwrap(await bridge.call_async_method(pyObjects.get(self), name, {args: args.map(toPy), kwargs})); const callStatic = (className, name, ...args) => unwrap(bridge.call_static(className, name, args.map(toPy))); // an attribute that is a promise, like the closed one of a reader const awaitAttr = async (self, name) => unwrap(await bridge.await_attr(pyObjects.get(self), name)); @@ -800,7 +800,7 @@ requireArguments(arguments, 1, 'RTCPeerConnection.createDataChannel'); const kwargs = convertDictionary(requireDictionary(init, 'RTCDataChannelInit'), 'RTCDataChannelInit', DATA_CHANNEL_INIT); - return callMethodWithKeywords(this, 'create_data_channel', [toUSVString(label)], kwargs); + return callMethod(this, 'create_data_channel', toUSVString(label), kwargs); } async addIceCandidate(candidate) { @@ -884,7 +884,7 @@ throw new TypeError('RTCError: missing required member errorDetail'); } // the library validates the error detail, an RTCErrorDetailType - construct('RTCError', {error_detail: pyEnum('RTCErrorDetailType', init.errorDetail), message: String(message)}); + construct('RTCErrorInit', {error_detail: pyEnum('RTCErrorDetailType', init.errorDetail)}); } super(message, 'OperationError'); const detail = { @@ -1007,7 +1007,7 @@ async copyTo(destination, options) { const bytes = bytesOf(destination, 'destination'); - const {layout, data} = unwrap(bridge.video_frame_copy_to(pyObjects.get(this), bytes, copyToOptions(options))); + const {layout, data} = unwrap(await bridge.video_frame_copy_to(pyObjects.get(this), bytes, copyToOptions(options))); bytes.set(new Uint8Array(data)); return layout; } diff --git a/tests/wpt/test_wpt.py b/tests/wpt/test_wpt.py index 24b8169..b7864e1 100644 --- a/tests/wpt/test_wpt.py +++ b/tests/wpt/test_wpt.py @@ -11,20 +11,27 @@ recorded as well: run `python -m tests.wpt update ` and commit the change. """ +from __future__ import annotations + +from typing import TYPE_CHECKING + import pytest +if TYPE_CHECKING: + from _pytest.mark import ParameterSet + pytest.importorskip('pythonmonkey') -from tests.wpt import runner # noqa: E402 -from tests.wpt.expectations import Expectations # noqa: E402 -from tests.wpt.loader import WPT_ROOT, discover # noqa: E402 +from tests.wpt import runner +from tests.wpt.expectations import Expectations +from tests.wpt.loader import WPT_ROOT, discover pytestmark = pytest.mark.skipif(not WPT_ROOT.is_dir(), reason='no wpt checkout, see tests/wpt/README.md') expectations = Expectations.load() -def _cases(): +def _cases() -> list[str | ParameterSet]: if not WPT_ROOT.is_dir(): return [] return [ @@ -36,7 +43,7 @@ def _cases(): @pytest.mark.parametrize('case', _cases()) -def test_wpt(case): +def test_wpt(case: str) -> None: problems = expectations.mismatches(case, runner.run(case)) if problems: pytest.fail('\n'.join(problems), pytrace=False) diff --git a/uv.lock b/uv.lock index d3553d3..df0f224 100644 --- a/uv.lock +++ b/uv.lock @@ -1547,7 +1547,7 @@ dev = [ { name = "pytest", specifier = ">=8" }, { name = "pytest-asyncio", specifier = ">=0.24" }, { name = "pytest-timeout", specifier = ">=2.3" }, - { name = "ruff", specifier = ">=0.13" }, + { name = "ruff", specifier = ">=0.16.9" }, { name = "scikit-build-core", specifier = ">=0.11" }, ] test = [