diff --git a/.github/scripts/fuzz.sh b/.github/scripts/fuzz.sh index 9d87d9c..d36b498 100755 --- a/.github/scripts/fuzz.sh +++ b/.github/scripts/fuzz.sh @@ -25,7 +25,7 @@ if [ ! -x "$BUILD/venv/bin/python" ]; then # built from source with this Clang, so its sanitizer runtime is the one the extension is instrumented for CLANG_BIN="$(command -v clang)" LIBFUZZER_LIB="$(clang -print-runtime-dir)/libclang_rt.fuzzer_no_main.a" \ "$BUILD/venv/bin/pip" install -q --no-binary atheris atheris \ - cmake ninja "pybind11>=3.0" pytest + cmake ninja "pybind11>=3.0" "typing_extensions>=4.10" pytest fi # shellcheck disable=SC1091 source "$BUILD/venv/bin/activate" diff --git a/.github/scripts/sanitizers-macos.sh b/.github/scripts/sanitizers-macos.sh index 4b4ab97..09e3a20 100755 --- a/.github/scripts/sanitizers-macos.sh +++ b/.github/scripts/sanitizers-macos.sh @@ -24,7 +24,7 @@ fi if [ ! -x "$BUILD/venv/bin/python" ]; then uv venv -q "$BUILD/venv" --python "$PYTHON" - uv pip install -q --python "$BUILD/venv/bin/python" cmake ninja "pybind11>=3.0" pytest pytest-asyncio pytest-timeout + uv pip install -q --python "$BUILD/venv/bin/python" cmake ninja "pybind11>=3.0" "typing_extensions>=4.10" pytest pytest-asyncio pytest-timeout fi export PATH="$BUILD/venv/bin:$PATH" diff --git a/.github/scripts/sanitizers.sh b/.github/scripts/sanitizers.sh index ce03da2..fb35890 100755 --- a/.github/scripts/sanitizers.sh +++ b/.github/scripts/sanitizers.sh @@ -20,7 +20,7 @@ dnf install -y -q clang lld compiler-rt llvm "$PYTHON" -m venv "$BUILD/venv" # shellcheck disable=SC1091 source "$BUILD/venv/bin/activate" -python -m pip install -q cmake ninja "pybind11>=3.0" pytest pytest-asyncio pytest-timeout +python -m pip install -q cmake ninja "pybind11>=3.0" "typing_extensions>=4.10" pytest pytest-asyncio pytest-timeout CC=clang CXX=clang++ cmake -S "$SRC" -B "$BUILD" -G Ninja \ -DCMAKE_BUILD_TYPE=RelWithDebInfo \ diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index efd4736..5c46d1a 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -27,6 +27,9 @@ jobs: - uses: astral-sh/setup-uv@v7 - run: uvx ruff check - run: uvx ruff format --check + # the imports of the checked code, without building the extension: its stub is in stubs/ + - run: uv sync --group dev --group wpt --no-install-project + - run: uvx pyrefly check - run: make format-check # clang-tidy reads the libwebrtc headers, from the cache the wheels use - uses: actions/cache@v6 diff --git a/Makefile b/Makefile index f005aa5..cd19aca 100644 --- a/Makefile +++ b/Makefile @@ -1,4 +1,4 @@ -.PHONY: dev test asan tsan fuzz lint format format-check tidy stub wheels doc clean +.PHONY: dev test asan tsan fuzz lint typecheck format format-check tidy stub wheels doc clean # pinned to the clang-tidy of .github/scripts/tidy.sh CLANG_FORMAT := uvx clang-format==22.1.8 @@ -30,6 +30,9 @@ lint: format-check uvx ruff check uvx ruff format --check +typecheck: + uvx pyrefly check + format: uvx ruff check --fix uvx ruff format @@ -45,6 +48,8 @@ tidy: stub: uv run --no-sync pybind11-stubgen wrtc -o build/stubs cp build/stubs/wrtc.pyi stubs/wrtc/__init__.pyi + # collections.abc.Buffer is 3.12+ + perl -pi -e 's/collections\.abc\.Buffer/typing_extensions.Buffer/g; s/^import typing$$/import typing\nimport typing_extensions/' stubs/wrtc/__init__.pyi # wheels for the current platform, exactly as CI builds them wheels: diff --git a/benchmarks/__main__.py b/benchmarks/__main__.py index 89d38a8..292ebd8 100644 --- a/benchmarks/__main__.py +++ b/benchmarks/__main__.py @@ -44,7 +44,7 @@ def _commit() -> str: try: 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 '') + return commit + (' with uncommitted changes' if dirty != '' else '') except (OSError, subprocess.CalledProcessError): return 'unknown' diff --git a/benchmarks/measure.py b/benchmarks/measure.py index 3d9f5c8..3833585 100644 --- a/benchmarks/measure.py +++ b/benchmarks/measure.py @@ -33,7 +33,7 @@ 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: + if len(values) == 0: return float('nan') ordered = sorted(values) return ordered[min(len(ordered) - 1, int(fraction * len(ordered)))] @@ -46,7 +46,7 @@ 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: asyncio.Task | None = None + self._task: asyncio.Task[None] | None = None async def _probe(self) -> None: loop = asyncio.get_running_loop() @@ -62,7 +62,7 @@ def __enter__(self) -> Self: def __exit__( self, exc_type: type[BaseException] | None, exc: BaseException | None, traceback: TracebackType | None ) -> None: - if self._task: + if self._task is not None: self._task.cancel() @property @@ -103,27 +103,29 @@ def stop(self) -> Usage: @property def cpu_percent(self) -> float: """Of one core: the process uses several threads (encoders, decoders, network).""" - return self.cpu / self.wall * 100 if self.wall else float('nan') + return self.cpu / self.wall * 100 if self.wall != 0 else float('nan') 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: + if len(samples) == 0: return float('nan') xs = [t for t, _ in samples] ys = [b / 1e6 for _, b in samples] mean_x, mean_y = statistics.fmean(xs), statistics.fmean(ys) denominator = sum((x - mean_x) ** 2 for x in xs) - if not denominator: + if denominator == 0: return float('nan') return sum((x - mean_x) * (y - mean_y) for x, y in zip(xs, ys)) / denominator * 60 def machine() -> str: """The CPU, the OS and Python of the machine.""" - cpu = platform.processor() or platform.machine() + cpu = platform.processor() + if cpu == '': + cpu = platform.machine() sysctl = shutil.which('sysctl') if sys.platform == 'darwin' else None - if sysctl: + if sysctl is not None: with contextlib.suppress(OSError, subprocess.CalledProcessError): cpu = subprocess.check_output([sysctl, '-n', 'machdep.cpu.brand_string'], text=True).strip() return ( diff --git a/benchmarks/media.py b/benchmarks/media.py index 0a186ba..30ce734 100644 --- a/benchmarks/media.py +++ b/benchmarks/media.py @@ -23,7 +23,7 @@ from tests.helpers import connect, rss_bytes, wait_for_event if TYPE_CHECKING: - from collections.abc import AsyncIterator, Awaitable, Iterable, Sequence + from collections.abc import AsyncGenerator, Awaitable, Iterable, Sequence # the frame number, drawn as bits in blocks of luma at the top left of each frame BITS = 16 @@ -85,7 +85,7 @@ def received_sizes(self) -> str: 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 = [] + frames: list[bytearray] = [] chroma = (width // 2) * (height // 2) row = bytes((x * 255 // width) for x in range(width)) base = bytearray(row * height + bytes([128]) * (2 * chroma)) @@ -101,14 +101,14 @@ 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 = LUMA_ONE if number >> bit & 1 else LUMA_ZERO + value = LUMA_ONE if number >> bit & 1 != 0 else LUMA_ZERO block = bytes([value]) * BLOCK for y in range(BLOCK): start = y * width + bit * BLOCK frame[start : start + BLOCK] = block -def _read_number(luma: bytes, stride: int) -> int: +def _read_number(luma: bytes | bytearray, stride: int) -> int: number = 0 for bit in range(BITS): # the center of the block, away from the blur of compression at its edges @@ -118,8 +118,24 @@ def _read_number(luma: bytes, stride: int) -> int: return number +def _video(media: object) -> webrtc.VideoFrame: + """The media a processor of a video track reads, as a video frame.""" + if not isinstance(media, webrtc.VideoFrame): + msg = f'expected a video frame, not {media!r}' + raise TypeError(msg) + return media + + +def _audio(media: object) -> webrtc.AudioData: + """The media a processor of an audio track reads, as audio data.""" + if not isinstance(media, webrtc.AudioData): + msg = f'expected audio data, not {media!r}' + raise TypeError(msg) + return media + + @contextlib.asynccontextmanager -async def _connection() -> AsyncIterator[tuple[webrtc.RTCPeerConnection, webrtc.RTCPeerConnection]]: +async def _connection() -> AsyncGenerator[tuple[webrtc.RTCPeerConnection, webrtc.RTCPeerConnection], None]: caller, callee = webrtc.RTCPeerConnection(), webrtc.RTCPeerConnection() try: yield caller, callee @@ -138,12 +154,16 @@ async def _remote_track( sender = caller.add_track(track) track_event = wait_for_event(callee, 'track', 30) await connect(caller, callee, 30) - if max_bitrate: + if max_bitrate is not None and max_bitrate != 0: parameters = sender.get_parameters() for encoding in parameters.encodings: encoding.max_bitrate = max_bitrate await sender.set_parameters(parameters) - return (await track_event).track + event = await track_event + if not isinstance(event, webrtc.RTCTrackEvent): + msg = f'expected a track event, not {event!r}' + raise TypeError(msg) + return event.track @dataclass @@ -207,11 +227,15 @@ async def run(self) -> VideoResult: processor = webrtc.MediaStreamTrackProcessor( webrtc.MediaStreamTrackProcessorInit(remote, max_buffer_size=self.max_buffer_size) ) - sampling = asyncio.ensure_future(run.sample_rss(self.rss_every)) if self.rss_every else None + sampling = ( + asyncio.ensure_future(run.sample_rss(self.rss_every)) + if self.rss_every is not None and self.rss_every != 0 + else None + ) lag = await run.phases.measure( run.result, processor, write=lambda: run.write(generator), read=lambda: run.read(processor) ) - if sampling: + if sampling is not None: await sampling run.result.lag_p95_ms, run.result.lag_max_ms = lag.p95_ms, lag.max_ms generator.track.stop() @@ -258,15 +282,16 @@ async def write(self, generator: webrtc.VideoTrackGenerator) -> None: async def read(self, processor: webrtc.MediaStreamTrackProcessor) -> None: header = webrtc.VideoFrameCopyToOptions(rect=webrtc.DOMRectInit(x=0, y=0, width=BITS * BLOCK, height=BLOCK)) loop = asyncio.get_running_loop() - async for frame in processor.readable: + async for media in processor.readable: now = loop.time() + frame = _video(media) 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: + if self.options.consumer_delay != 0: await asyncio.sleep(self.options.consumer_delay) if self.phases.done.is_set(): break @@ -342,7 +367,8 @@ async def write() -> None: await asyncio.sleep(max(0.0, start + written / 100 - loop.time())) async def read() -> None: - async for audio in processor.readable: + async for media in processor.readable: + audio = _audio(media) if phases.measuring.is_set(): result.received += 1 result.received_frames += audio.number_of_frames @@ -372,10 +398,12 @@ def megapixels_per_second(self) -> float: async def copy_costs( - sizes: Iterable[tuple[int, int]], formats: Sequence[str] = ('I420', 'RGBA', 'BGRA'), budget: float = 1.0 + sizes: Iterable[tuple[int, int]], + formats: Sequence[webrtc.VideoPixelFormatValue] = ('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 = [] + results: list[CopyResult] = [] for width, height in sizes: chroma = (width // 2) * (height // 2) frame = webrtc.VideoFrame( diff --git a/cmake/libcxx/update.py b/cmake/libcxx/update.py index 749e786..8197cbe 100755 --- a/cmake/libcxx/update.py +++ b/cmake/libcxx/update.py @@ -144,7 +144,7 @@ def git_fetch(url: str, ref: str, repo: str) -> str: 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: + if match is None: msg = f'{LLVM_DIRS[name][0]} is not in the WebRTC DEPS' raise SystemExit(msg) url, revision = match.groups() @@ -158,25 +158,26 @@ def resolve_llvm(deps: str, name: str) -> str: 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']) + return commit['sha'] msg = f'no llvm-project commit has the {name} tree of {url}@{revision}' raise SystemExit(msg) def runtime_sources(tag: str) -> list[str]: """The libc++ and libc++abi sources of Chromium's Linux build, as llvm-project paths.""" - sources = [] + sources: list[str] = [] for gn, llvm in (('libc%2B%2B', 'libcxx'), ('libc%2B%2Babi', 'libcxxabi')): for line in chromium(f'buildtools/third_party/{gn}/BUILD.gn', tag).decode().splitlines(): match = re.search(r'"//third_party/libc\+\+(?:abi)?/src/src/([^"]+\.cpp)"', line) - if match and not line.lstrip().startswith('#') and 'win32' not in match[1]: + if match is not None and not line.lstrip().startswith('#') and 'win32' not in match[1]: sources.append(f'{llvm}/src/{match[1]}') return sorted(set(sources)) 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() + found: set[str] = set() + include: bytes # one group, so findall gives its bytes for include in INCLUDE.findall(data): name = include.decode() found.add(posixpath.normpath(posixpath.join(posixpath.dirname(path), name))) @@ -193,13 +194,13 @@ def pin_runtime(commits: dict[str, str], tag: str) -> None: pinned: dict[str, bytes] = {} queue = list(sources) with ThreadPoolExecutor(8) as pool: - while queue: + while len(queue) > 0: pinned.update(zip(queue, pool.map(blob, [tree[p] for p in queue]))) - found = set().union(*(includes(path, pinned[path]) for path in queue)) + found = set[str]().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: + if len(stray) > 0: 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())] @@ -222,7 +223,7 @@ def pin_headers(commit: str, headers: Path) -> None: remote = {e['path']: e['sha'] for e in tree} 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: + if len(mismatches) > 0: msg = f'the prebuilt headers differ from llvm-project@{commit}: {mismatches[:10]}' raise SystemExit(msg) @@ -240,7 +241,7 @@ def sha256(entry: TreeEntry) -> str: def pin_config(tag: str) -> None: """Pins Chromium's build-generated libc++ config.""" - config = [] + config: list[str] = [] for name in CHROMIUM_CONFIG: data = chromium(f'buildtools/third_party/libc%2B%2B/{name}', tag) config.append(f'{hashlib.sha256(data).hexdigest()} {name}\n') diff --git a/examples/echo.py b/examples/echo.py index 183a108..6623c65 100755 --- a/examples/echo.py +++ b/examples/echo.py @@ -40,14 +40,22 @@ async def grayscale(frame: webrtc.VideoFrame, controller: webrtc.TransformStream frame.close() +def video_frame(media: object) -> webrtc.VideoFrame: + """The media a processor of a video track reads, as a video frame.""" + if not isinstance(media, webrtc.VideoFrame): + msg = f'expected a video frame, not {media!r}' + raise TypeError(msg) + return media + + 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(webrtc.MediaStreamTrackProcessorInit(track)).readable.get_reader() loop = asyncio.get_running_loop() end = loop.time() + SECONDS - frames = 0 + frames, rgba = 0, bytearray() while loop.time() < end: - frame = (await reader.read()).value + frame = video_frame((await reader.read()).value) options = webrtc.VideoFrameCopyToOptions(format='RGBA') rgba = bytearray(frame.allocation_size(options)) await frame.copy_to(rgba, options) @@ -64,7 +72,7 @@ def trickle(caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection) async def on_candidate( event: webrtc.RTCPeerConnectionIceEvent, other: webrtc.RTCPeerConnection = other ) -> None: - if event.candidate: + if event.candidate is not None: await other.add_ice_candidate(event.candidate) pc.on('icecandidate', on_candidate) diff --git a/examples/janus_streaming.py b/examples/janus_streaming.py index f35a074..c9efff3 100755 --- a/examples/janus_streaming.py +++ b/examples/janus_streaming.py @@ -26,21 +26,69 @@ import sys import threading import uuid -from typing import TYPE_CHECKING, Any +from typing import TYPE_CHECKING, TypedDict import httpx import sounddevice +from typing_extensions import NotRequired, Required import webrtc if TYPE_CHECKING: - from collections.abc import Iterator + from collections.abc import Generator from types import TracebackType from typing_extensions import Self JANUS = 'https://janus.conf.meetecho.com/janus' -Json = dict[str, Any] + + +class Stream(TypedDict): + """A stream of the streaming plugin.""" + + id: int + description: NotRequired[str] + + +class Jsep(TypedDict): + """A session description of the Janus API.""" + + type: str + sdp: str + + +class PluginReply(TypedDict, total=False): + """The data of a plugin reply: the streams of a list request.""" + + list: list[Stream] + + +class PluginData(TypedDict, total=False): + """The reply of a plugin.""" + + data: PluginReply + + +class Created(TypedDict): + """The id of a created session or handle.""" + + id: int + + +class Error(TypedDict): + """An error of the Janus API.""" + + reason: str + + +class Reply(TypedDict, total=False): + """A message of the Janus API: the fields used here.""" + + janus: Required[str] + data: Created + error: Error + plugindata: PluginData + jsep: Jsep class Janus: @@ -48,12 +96,12 @@ class Janus: def __init__(self) -> None: self.client = httpx.AsyncClient(timeout=60) + self.session = '' + self.handle = '' 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"]}' + self.session = await self._create('', janus='create') + self.handle = await self._create(self.session, janus='attach', plugin='janus.plugin.streaming') return self async def __aexit__( @@ -62,21 +110,31 @@ async def __aexit__( await self._post(self.session, janus='destroy') await self.client.aclose() - async def request(self, body: Json, **extra: Json) -> Json | None: + async def request(self, body: dict[str, object], **extra: object) -> PluginReply | 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) -> Json: + async def event(self) -> Reply: """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() + event: Reply = (await self.client.get(JANUS + self.session)).json() + return event - async def _post(self, path: str, **message: object) -> Json: - reply: Json = (await self.client.post(JANUS + path, json={'transaction': uuid.uuid4().hex, **message})).json() + async def _post(self, path: str, **message: object) -> Reply: + reply: Reply = (await self.client.post(JANUS + path, json={'transaction': uuid.uuid4().hex, **message})).json() if reply['janus'] == 'error': - raise RuntimeError(reply['error']['reason']) + error = reply.get('error') + raise RuntimeError(error['reason'] if error is not None else reply) return reply + async def _create(self, path: str, **message: object) -> str: + """Creates a session or a handle under the path, returns its path.""" + reply = await self._post(path, **message) + if 'data' not in reply: + msg = f'no id in {reply}' + raise RuntimeError(msg) + return f'{path}/{reply["data"]["id"]}' + class Speakers: """Plays 16-bit audio from a buffer that keeps half a second at most.""" @@ -86,7 +144,7 @@ def __init__(self, rate: int, channels: int) -> None: self.stream = sounddevice.RawOutputStream(rate, channels=channels, dtype='int16', callback=self._on_need) self.stream.start() - def play(self, samples: bytes) -> None: + def play(self, samples: bytes | bytearray) -> None: """Queues samples, dropping the oldest ones over the limit.""" with self.lock: self.buffer += samples @@ -99,7 +157,7 @@ def _on_need(self, out: memoryview, *_: object) -> None: out[:] = chunk.ljust(len(out), b'\0') -def draw(rgbx: bytes, width: int, height: int) -> None: +def draw(rgbx: bytes | bytearray, 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) @@ -117,7 +175,7 @@ def color(x: int, y: int) -> str: @contextlib.contextmanager -def fullscreen() -> Iterator[None]: +def fullscreen() -> Generator[None, None, None]: """Switches to the alternate screen, without the cursor.""" sys.stdout.write('\033[?1049h\033[?25l') try: @@ -132,33 +190,44 @@ async def watch(track: webrtc.MediaStreamTrack) -> None: async for frame in webrtc.MediaStreamTrackProcessor( webrtc.MediaStreamTrackProcessorInit(track, max_buffer_size=1) ).readable: + if not isinstance(frame, webrtc.VideoFrame): + msg = f'expected a video frame, not {frame!r}' + raise TypeError(msg) with frame: options = webrtc.VideoFrameCopyToOptions(format='RGBX') rgbx = bytearray(frame.allocation_size(options)) await frame.copy_to(rgbx, options) size = frame.visible_rect + if size is None: + msg = 'the frame has no visible rect' + raise RuntimeError(msg) draw(rgbx, int(size.width), int(size.height)) async def listen(track: webrtc.MediaStreamTrack) -> None: """Plays the audio track until it ends.""" - speakers = None + speakers: Speakers | None = None try: async for data in webrtc.MediaStreamTrackProcessor( webrtc.MediaStreamTrackProcessorInit(track, max_buffer_size=50) ).readable: + if not isinstance(data, webrtc.AudioData): + msg = f'expected audio data, not {data!r}' + raise TypeError(msg) with data: options = webrtc.AudioDataCopyToOptions(plane_index=0, format='s16') samples = bytearray(data.allocation_size(options)) data.copy_to(samples, options) - speakers = speakers or Speakers(int(data.sample_rate), data.number_of_channels) + rate, channels = int(data.sample_rate), data.number_of_channels + if speakers is None: + speakers = Speakers(rate, channels) speakers.play(samples) finally: if speakers is not None: speakers.stream.close() -async def answer(pc: webrtc.RTCPeerConnection, offer: dict[str, str]) -> dict[str, str]: +async def answer(pc: webrtc.RTCPeerConnection, offer: Jsep) -> Jsep: """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()) @@ -166,6 +235,9 @@ async def answer(pc: webrtc.RTCPeerConnection, offer: dict[str, str]) -> dict[st await pc.set_local_description(await pc.create_answer()) with contextlib.suppress(asyncio.TimeoutError): await asyncio.wait_for(gathered.wait(), 5) + if pc.local_description is None: + msg = 'no local description' + raise RuntimeError(msg) return {'type': 'answer', 'sdp': pc.local_description.sdp} @@ -177,7 +249,11 @@ async def keep_alive(janus: Janus) -> None: async def pick_stream(janus: Janus) -> int: """Lists the streams of the server, returns the first one.""" - streams = (await janus.request({'request': 'list'}))['list'] + reply = await janus.request({'request': 'list'}) + if reply is None or 'list' not in reply: + msg = 'the server did not list its streams' + raise RuntimeError(msg) + streams = reply['list'] for stream in streams: print(f'{stream["id"]}: {stream.get("description")}') return streams[0]['id'] @@ -197,7 +273,7 @@ def on_track(event: webrtc.RTCTrackEvent) -> None: if stream_id is None: stream_id = await pick_stream(janus) await janus.request({'request': 'watch', 'id': stream_id}) - event: Json = {} + event = await janus.event() while 'jsep' not in event: event = await janus.event() await janus.request({'request': 'start'}, jsep=await answer(pc, event['jsep'])) diff --git a/examples/openai_live.py b/examples/openai_live.py index 87a0523..21c2a10 100755 --- a/examples/openai_live.py +++ b/examples/openai_live.py @@ -41,10 +41,11 @@ import threading import time from array import array -from typing import Any, ClassVar +from typing import ClassVar, TypedDict import httpx import sounddevice +from typing_extensions import NotRequired import webrtc @@ -57,6 +58,60 @@ SILENT_MIC_WARNING = 10 # seconds of silence before the microphone is suspected +class Session(TypedDict, total=False): + """The session of a reply of the API.""" + + id: str + + +class Transport(TypedDict): + """The transport of a reply of the API: the answer.""" + + sdp: str + + +class SessionReply(TypedDict): + """The reply of the API to a new session.""" + + session: NotRequired[Session] + transport: Transport + + +class Usage(TypedDict, total=False): + """The usage of a session.""" + + seconds: float + + +class ContextWindow(TypedDict, total=False): + """How full the context of the model is.""" + + usage_ratio: object + + +class ApiError(TypedDict, total=False): + """An error of the API.""" + + message: str + + +class ErrorReply(TypedDict, total=False): + """The reply of the API to a request that failed.""" + + error: ApiError + + +class Message(TypedDict, total=False): + """An event of the API on the event channel: the fields used here.""" + + type: str + delta: str + reason: str + usage: Usage + error: ApiError + context_window: ContextWindow + + class Console: """Human-readable output: timestamped status lines and live transcripts that share the terminal.""" @@ -72,7 +127,7 @@ class Console: def __init__(self, *, verbose: bool) -> None: self.verbose = verbose - self.color = sys.stdout.isatty() and not os.environ.get('NO_COLOR') + self.color = sys.stdout.isatty() and os.environ.get('NO_COLOR', '') == '' self.speaker: str | None = None # who the open transcript line belongs to def paint(self, text: str, color: str) -> str: @@ -80,7 +135,7 @@ def paint(self, text: str, color: str) -> str: return f'\033[{self.COLORS[color]}m{text}\033[0m' if self.color else text def _end_transcript(self) -> None: - if self.speaker: + if self.speaker is not None: print(flush=True) self.speaker = None @@ -88,7 +143,7 @@ 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) + print(self.paint(line, color) if color is not None else line, flush=True) def info(self, text: str) -> None: """Prints a status line.""" @@ -121,10 +176,10 @@ def transcript(self, speaker: str, delta: str) -> None: print(delta, end='', flush=True) -def peak(samples: bytes) -> float: +def peak(samples: bytes | bytearray) -> float: """The peak level of 16-bit samples, from 0 to 1.""" values = array('h', samples) - return max(*values, -min(values)) / 32768 if values else 0 + return max(*values, -min(values)) / 32768 if len(values) > 0 else 0 class Microphone: @@ -167,7 +222,7 @@ def __init__(self, device: str | int | None) -> None: self._lock = threading.Lock() self._limit = 0 - def play(self, samples: bytes, sample_rate: int, channels: int) -> None: + def play(self, samples: bytes | bytearray, 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 @@ -221,39 +276,49 @@ async def run(self, hang_up: asyncio.Event) -> None: if not args.barge_in: console.info('Echo guard on: the mic is muted while the assistant speaks (--barge-in turns it off)') - self.pc = webrtc.RTCPeerConnection() - self.pc.on('connectionstatechange', self._on_connection_state) - self.pc.on('track', self._on_track) + self.pc = pc = webrtc.RTCPeerConnection() + pc.on('connectionstatechange', self._on_connection_state) + pc.on('track', self._on_track) generator = webrtc.MediaStreamTrackGenerator('audio') - self.pc.add_track(generator) + pc.add_track(generator) # created before the offer, so the offer negotiates it - self.events = self.pc.create_data_channel('oai-events') + self.events = pc.create_data_channel('oai-events') self.events.on('open', lambda _event: console.ok('Event channel open')) self.events.on('message', self._on_message) - offer = await self.pc.create_offer() - await self.pc.set_local_description(offer) - await self._gathered() + offer = await pc.create_offer() + await pc.set_local_description(offer) + sdp = await self._gathered(pc) console.info(f'Creating a {args.model} session...') - answer = await self._create_session(self.pc.local_description.sdp) - await self.pc.set_remote_description(webrtc.RTCSessionDescriptionInit('answer', answer)) + answer = await self._create_session(sdp) + await pc.set_remote_description(webrtc.RTCSessionDescriptionInit('answer', answer)) self.microphone.start() - self.tasks.append(asyncio.ensure_future(self._send_microphone(generator.writable.get_writer()))) + writer = generator.writable.get_writer() + self.tasks.append(asyncio.ensure_future(self._send_microphone(writer, self.microphone))) await hang_up.wait() await self.close() - async def _gathered(self) -> None: - """Waits for the local ICE candidates, which go in the offer since there is no trickling.""" + async def _gathered(self, pc: webrtc.RTCPeerConnection) -> str: + """Waits for the local ICE candidates, which go in the offer since there is no trickling, returns the offer.""" done = asyncio.Event() - self.pc.on('icegatheringstatechange', lambda _event: self.pc.ice_gathering_state == 'complete' and done.set()) - if self.pc.ice_gathering_state != 'complete': + + def on_gathering_state(_event: webrtc.Event) -> None: + if pc.ice_gathering_state == 'complete': + done.set() + + pc.on('icegatheringstatechange', on_gathering_state) + if 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') + if pc.local_description is None: + msg = 'No local description' + raise CallError(msg) + return pc.local_description.sdp async def _create_session(self, sdp: str) -> str: session: dict[str, object] = {'model': self.args.model} @@ -273,11 +338,14 @@ async def _create_session(self, sdp: str) -> str: raise CallError(msg) from None if response.is_error: raise CallError(_describe_http_error(response)) - reply: dict[str, Any] = response.json() + reply: SessionReply = response.json() self.console.ok(f'Session created: {reply.get("session", {}).get("id", "?")}') return reply['transport']['sdp'] def _on_connection_state(self, _event: webrtc.Event) -> None: + if self.pc is None: + msg = 'No connection' + raise CallError(msg) state = self.pc.connection_state messages = { 'connecting': ('info', 'Connecting audio...'), @@ -290,16 +358,19 @@ def _on_connection_state(self, _event: webrtc.Event) -> None: getattr(self.console, level)(text) def _on_track(self, event: webrtc.RTCTrackEvent) -> None: + if self.speakers is None: + msg = 'No speakers' + raise CallError(msg) self.console.debug(f"Receiving the assistant's {event.track.kind} track") - self.tasks.append(asyncio.ensure_future(self._play(event.track))) + self.tasks.append(asyncio.ensure_future(self._play(event.track, self.speakers))) - async def _send_microphone(self, writer: webrtc.WritableStreamDefaultWriter) -> None: + async def _send_microphone(self, writer: webrtc.WritableStreamDefaultWriter, microphone: Microphone) -> 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() while True: - chunk = await self.microphone.queue.get() + chunk = await microphone.queue.get() if not heard: if peak(chunk) > VOICE_LEVEL: heard = True @@ -324,12 +395,15 @@ async def _send_microphone(self, writer: webrtc.WritableStreamDefaultWriter) -> await writer.write(data) timestamp += 10_000 - async def _play(self, track: webrtc.MediaStreamTrack) -> None: + async def _play(self, track: webrtc.MediaStreamTrack, speakers: Speakers) -> None: """Plays the assistant's audio.""" heard = False async for data in webrtc.MediaStreamTrackProcessor( webrtc.MediaStreamTrackProcessorInit(track, max_buffer_size=50) ).readable: + if not isinstance(data, webrtc.AudioData): + msg = f'expected audio data, not {data!r}' + raise TypeError(msg) with data: options = webrtc.AudioDataCopyToOptions(plane_index=0, format='s16') samples = bytearray(data.allocation_size(options)) @@ -340,13 +414,12 @@ async def _play(self, track: webrtc.MediaStreamTrack) -> None: if not heard: heard = True self.console.ok(f"Receiving the assistant's voice ({rate} Hz, {channels} ch)") - self.speakers.play(samples, rate, channels) + speakers.play(samples, rate, channels) def _on_message(self, event: webrtc.MessageEvent) -> None: console = self.console - try: - message: dict[str, Any] = json.loads(event.data) - except ValueError: + message = parse_message(event.data) + if message is None: console.debug(f'Not JSON: {event.data!r}') return kind = message.get('type') @@ -358,12 +431,15 @@ def _on_message(self, event: webrtc.MessageEvent) -> None: console.ok('Session started, say something! (Ctrl+C to hang up)') elif kind == 'session.closed': usage = message.get('usage', {}) - reason = message.get('reason') - console.info(f'Session closed{f" ({reason})" if reason else ""}, {usage.get("seconds", "?")} s billed') + reason = message.get('reason', '') + console.info( + f'Session closed{f" ({reason})" if reason != "" else ""}, {usage.get("seconds", "?")} s billed' + ) self.session_closed.set() elif kind == 'error': error = message.get('error', {}) - console.error(f'API error: {error.get("message") or error}') + text = error.get('message') + console.error(f'API error: {text if text is not None and text != "" else error}') elif kind == 'session.usage.updated': usage = message.get('usage', {}) ratio = message.get('context_window', {}).get('usage_ratio') @@ -394,15 +470,28 @@ async def close(self) -> None: console.ok('Bye!') +def parse_message(data: str | bytes | webrtc.Blob) -> Message | None: + """An event of the API from the event channel, None if it isn't JSON.""" + if isinstance(data, webrtc.Blob): + msg = 'the binary type of the event channel is bytes, not blobs' + raise TypeError(msg) + try: + message: Message = json.loads(data) + except ValueError: + return None + return message + + class CallError(Exception): """The call couldn't be set up.""" def _describe_http_error(response: httpx.Response) -> str: try: - message = response.json().get('error', {}).get('message') + reply: ErrorReply = response.json() except ValueError: - message = None + reply = {} + message = reply.get('error', {}).get('message', '') hints = { 401: 'the API key is invalid', 403: 'the key has no access to this model', @@ -410,7 +499,7 @@ def _describe_http_error(response: httpx.Response) -> str: 429: 'rate limit or quota exceeded', } summary = hints.get(response.status_code, response.reason_phrase) - return f'OpenAI API returned {response.status_code}, {summary}' + (f': {message}' if message else '') + return f'OpenAI API returned {response.status_code}, {summary}' + (f': {message}' if message != '' else '') def parse_args() -> argparse.Namespace: @@ -430,7 +519,7 @@ def parse_args() -> argparse.Namespace: parser.add_argument('-v', '--verbose', action='store_true', help='log every event from the API') args = parser.parse_args() for name in ('input_device', 'output_device'): - value = getattr(args, name) + value: str | None = getattr(args, name) if value is not None and value.isdigit(): setattr(args, name, int(value)) return args diff --git a/examples/recorder.py b/examples/recorder.py index 02c7fc5..4f98f86 100755 --- a/examples/recorder.py +++ b/examples/recorder.py @@ -33,13 +33,17 @@ async def record(track: webrtc.MediaStreamTrack, file: BinaryIO) -> None: async for media in webrtc.MediaStreamTrackProcessor( webrtc.MediaStreamTrackProcessorInit(track, max_buffer_size=30) ).readable: - if track.kind == 'audio': + # audio data for an audio track, video frames for a video one + if isinstance(media, webrtc.AudioData): options = webrtc.AudioDataCopyToOptions(plane_index=0) data = bytearray(media.allocation_size(options)) media.copy_to(data, options) - else: + elif isinstance(media, webrtc.VideoFrame): data = bytearray(media.allocation_size()) await media.copy_to(data) + else: + msg = f'unexpected media: {media!r}' + raise TypeError(msg) media.close() file.write(data) frames += 1 @@ -53,7 +57,7 @@ def trickle(caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection) async def on_candidate( event: webrtc.RTCPeerConnectionIceEvent, other: webrtc.RTCPeerConnection = other ) -> None: - if event.candidate: + if event.candidate is not None: await other.add_ice_candidate(event.candidate) pc.on('icecandidate', on_candidate) diff --git a/examples/telegram_group_calls.py b/examples/telegram_group_calls.py index 83624f1..255f52b 100755 --- a/examples/telegram_group_calls.py +++ b/examples/telegram_group_calls.py @@ -129,7 +129,7 @@ async def send_audio_data(generator: webrtc.MediaStreamTrackGenerator, file: Bin start = loop.time() chunks = 0 - while data := file.read(480 * 4): # 480 frames of 2 channels of 16 bits + while (data := file.read(480 * 4)) != b'': # 480 frames of 2 channels of 16 bits frames = len(data) // 4 await writer.write( webrtc.AudioData( @@ -168,7 +168,8 @@ async def on_update(update: UpdateGroupCallWrapper | UpdateGroupCallParticipants return if isinstance(update.call, GroupCallWrapper): answered.set() - answer = build_answer(json.loads(update.call.params.data)) + params: CallParams = json.loads(update.call.params.data) + answer = build_answer(params) await pc.set_remote_description( webrtc.RTCSessionDescription(webrtc.RTCSessionDescriptionInit(webrtc.RTCSdpType.answer, answer)) ) diff --git a/pyproject.toml b/pyproject.toml index d8ef7df..8f56269 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -11,6 +11,7 @@ license = "BSD-3-Clause" license-files = ["LICENSE.md", "THIRD_PARTY_LICENSES.md"] authors = [{ name = "Ilya (Marshal)", email = "ilya@marshal.dev" }] requires-python = ">=3.9" +dependencies = ["typing_extensions>=4.10"] classifiers = [ "Development Status :: 1 - Planning", "Natural Language :: English", @@ -130,6 +131,7 @@ ignore = [ "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 + "PLC1901", # pyrefly's implicit-bool wants explicit comparisons, like != '' ] [tool.ruff.lint.per-file-ignores] @@ -175,3 +177,41 @@ max-public-methods = 20 [tool.ruff.lint.isort] required-imports = ["from __future__ import annotations"] known-first-party = ["webrtc", "wrtc", "tests", "benchmarks"] + +[tool.pyrefly] +required-version = ">=1.3.2" +preset = "all" +project-includes = [ + "python-webrtc/python/webrtc/**/*.py", "tests/**/*.py", "examples/**/*.py", "benchmarks/**/*.py", "cmake/**/*.py", + "stubs/**/*.pyi", +] +# tests/fuzz runs with its directory on PYTHONPATH (.github/scripts/fuzz.sh) +search-path = ["python-webrtc/python", "stubs", ".", "tests/fuzz"] +python-version = "3.9" +python-platform = "all" +python-interpreter-path = ".venv/bin/python" +strict-callable-subtyping = true +infer-with-first-use = false +permissive-ignores = false +type-ignore-unknown-tag-behavior = "no-effect" +# optional modules of the fuzzers and examples, not installed with the project +ignore-missing-imports = ["atheris", "httpx", "pytgcalls", "pytgcalls.*", "pyrogram", "pyrogram.*", "sounddevice"] + +# like WebIDL, the library converts what callers pass (bool(), int(), str()), whatever the annotations say +[[tool.pyrefly.sub-config]] +matches = "python-webrtc/python/webrtc/**" +errors = { unnecessary-type-conversion = false } + +# callers of browser-style APIs drop results, like the transceiver of add_track(); the library keeps the check +[[tool.pyrefly.sub-config]] +matches = "tests/**" +errors = { unused-call-result = false } +[[tool.pyrefly.sub-config]] +matches = "examples/**" +errors = { unused-call-result = false } +[[tool.pyrefly.sub-config]] +matches = "benchmarks/**" +errors = { unused-call-result = false } +[[tool.pyrefly.sub-config]] +matches = "cmake/**" +errors = { unused-call-result = false } diff --git a/python-webrtc/python/webrtc/__init__.py b/python-webrtc/python/webrtc/__init__.py index 8d6ea0b..9309667 100644 --- a/python-webrtc/python/webrtc/__init__.py +++ b/python-webrtc/python/webrtc/__init__.py @@ -46,6 +46,24 @@ VideoMatrixCoefficients, AlphaOption, AudioSampleFormat, + RTCSdpTypeValue, + TransceiverDirectionValue, + MediaTypeValue, + RTCPriorityTypeValue, + RTCDegradationPreferenceValue, + RTCIceTransportPolicyValue, + RTCBundlePolicyValue, + RTCRtcpMuxPolicyValue, + RTCRtpHeaderEncryptionPolicyValue, + RTCIceServerTransportProtocolValue, + RTCErrorDetailTypeValue, + BinaryTypeValue, + VideoPixelFormatValue, + VideoColorPrimariesValue, + VideoTransferCharacteristicsValue, + VideoMatrixCoefficientsValue, + AlphaOptionValue, + AudioSampleFormatValue, ) from .base import WebRTCObject from .exceptions import ( @@ -177,11 +195,14 @@ __all__ = [ 'Algorithm', 'AlphaOption', + 'AlphaOptionValue', 'AudioData', 'AudioDataCopyToOptions', 'AudioDataInit', 'AudioSampleFormat', + 'AudioSampleFormatValue', 'BinaryType', + 'BinaryTypeValue', 'Blob', 'ConstrainBooleanOrDOMStringParameters', 'ConstrainBooleanParameters', @@ -216,6 +237,7 @@ 'MediaTrackConstraints', 'MediaTrackSettings', 'MediaType', + 'MediaTypeValue', 'MessageEvent', 'NetworkError', 'NotSupportedError', @@ -225,6 +247,7 @@ 'PythonWebRTCException', 'PythonWebRTCExceptionBase', 'RTCBundlePolicy', + 'RTCBundlePolicyValue', 'RTCCertificate', 'RTCConfiguration', 'RTCDTMFSender', @@ -234,10 +257,12 @@ 'RTCDataChannelInit', 'RTCDataChannelState', 'RTCDegradationPreference', + 'RTCDegradationPreferenceValue', 'RTCDtlsFingerprint', 'RTCDtlsTransport', 'RTCError', 'RTCErrorDetailType', + 'RTCErrorDetailTypeValue', 'RTCErrorEvent', 'RTCErrorInit', 'RTCException', @@ -253,9 +278,11 @@ 'RTCIceRole', 'RTCIceServer', 'RTCIceServerTransportProtocol', + 'RTCIceServerTransportProtocolValue', 'RTCIceTcpCandidateType', 'RTCIceTransport', 'RTCIceTransportPolicy', + 'RTCIceTransportPolicyValue', 'RTCIceTransportState', 'RTCLocalSessionDescriptionInit', 'RTCOAuthCredential', @@ -264,7 +291,9 @@ 'RTCPeerConnectionIceEvent', 'RTCPeerConnectionState', 'RTCPriorityType', + 'RTCPriorityTypeValue', 'RTCRtcpMuxPolicy', + 'RTCRtcpMuxPolicyValue', 'RTCRtcpParameters', 'RTCRtpCapabilities', 'RTCRtpCodec', @@ -272,6 +301,7 @@ 'RTCRtpContributingSource', 'RTCRtpEncodingParameters', 'RTCRtpHeaderEncryptionPolicy', + 'RTCRtpHeaderEncryptionPolicyValue', 'RTCRtpHeaderExtensionCapability', 'RTCRtpHeaderExtensionParameters', 'RTCRtpReceiveParameters', @@ -283,6 +313,7 @@ 'RTCRtpTransceiverInit', 'RTCSctpTransport', 'RTCSdpType', + 'RTCSdpTypeValue', 'RTCSessionDescription', 'RTCSessionDescriptionInit', 'RTCSignalingState', @@ -297,10 +328,12 @@ 'SctpTransportState', 'SdpParseException', 'TransceiverDirection', + 'TransceiverDirectionValue', 'TransformStream', 'TransformStreamDefaultController', 'ULongRange', 'VideoColorPrimaries', + 'VideoColorPrimariesValue', 'VideoColorSpace', 'VideoColorSpaceInit', 'VideoFrame', @@ -309,9 +342,12 @@ 'VideoFrameInit', 'VideoFrameMetadata', 'VideoMatrixCoefficients', + 'VideoMatrixCoefficientsValue', 'VideoPixelFormat', + 'VideoPixelFormatValue', 'VideoTrackGenerator', 'VideoTransferCharacteristics', + 'VideoTransferCharacteristicsValue', 'WebRTCObject', 'WritableStream', 'WritableStreamDefaultController', diff --git a/python-webrtc/python/webrtc/base.py b/python-webrtc/python/webrtc/base.py index c3d3b72..48eca8e 100644 --- a/python-webrtc/python/webrtc/base.py +++ b/python-webrtc/python/webrtc/base.py @@ -9,7 +9,7 @@ from __future__ import annotations -from typing import TYPE_CHECKING, Any, Callable, ClassVar, Generic, TypeVar +from typing import TYPE_CHECKING, Callable, Generic, TypeVar from webrtc.utils.events import EventTarget @@ -29,13 +29,19 @@ class WebRTCObject(Generic[_NativeT]): """ #: The native class, created with no arguments when no native object is given - _class: ClassVar[Callable[[], Any] | None] = None + _class: Callable[..., _NativeT] | None = None + __obj: _NativeT 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() + if native_obj is None: + if self._class is None: + msg = f'{type(self).__name__} has no native class' + raise TypeError(msg) + native_obj = self._class() + self.__obj = native_obj @property def _native_obj(self) -> _NativeT: @@ -64,7 +70,8 @@ def __repr__(self) -> str: def __eq__(self, other: object) -> bool: if isinstance(other, WebRTCObject): - return self._native_obj is other._native_obj + other_obj: object = other._native_obj + return self._native_obj is other_obj return NotImplemented def __hash__(self) -> int: diff --git a/python-webrtc/python/webrtc/enums.py b/python-webrtc/python/webrtc/enums.py index 26bcb0e..a32eac3 100644 --- a/python-webrtc/python/webrtc/enums.py +++ b/python-webrtc/python/webrtc/enums.py @@ -10,11 +10,13 @@ from __future__ import annotations import enum +from typing import Literal class _StrEnum(str, enum.Enum): def __str__(self) -> str: - return self.value + value: str = self.value + return value class RTCPeerConnectionState(_StrEnum): @@ -70,6 +72,10 @@ class RTCSdpType(_StrEnum): rollback = 'rollback' +#: The values of :obj:`RTCSdpType`, which parameters taking it take too +RTCSdpTypeValue = Literal['offer', 'pranswer', 'answer', 'rollback'] + + class MediaStreamTrackState(_StrEnum): """The state of a track.""" @@ -96,6 +102,10 @@ class TransceiverDirection(_StrEnum): stopped = 'stopped' +#: The values of :obj:`TransceiverDirection`, which parameters taking it take too +TransceiverDirectionValue = Literal['sendrecv', 'sendonly', 'recvonly', 'inactive', 'stopped'] + + class MediaType(_StrEnum): """The kind of a track: audio or video. Data and unsupported media sections exist in libwebrtc only.""" @@ -105,6 +115,10 @@ class MediaType(_StrEnum): unsupported = 'unsupported' +#: The values of :obj:`MediaType`, which parameters taking it take too +MediaTypeValue = Literal['audio', 'video', 'data', 'unsupported'] + + class RTCIceComponent(_StrEnum): """The component of an ICE transport or candidate.""" @@ -177,6 +191,10 @@ class RTCPriorityType(_StrEnum): high = 'high' +#: The values of :obj:`RTCPriorityType`, which parameters taking it take too +RTCPriorityTypeValue = Literal['very-low', 'low', 'medium', 'high'] + + class RTCDegradationPreference(_StrEnum): """What a video sender degrades first when it can't keep up.""" @@ -186,6 +204,15 @@ class RTCDegradationPreference(_StrEnum): maintain_framerate_and_resolution = 'maintain-framerate-and-resolution' +#: The values of :obj:`RTCDegradationPreference`, which parameters taking it take too +RTCDegradationPreferenceValue = Literal[ + 'maintain-framerate', + 'maintain-resolution', + 'balanced', + 'maintain-framerate-and-resolution', +] + + class RTCIceTransportPolicy(_StrEnum): """Which ICE candidates may be used.""" @@ -193,6 +220,10 @@ class RTCIceTransportPolicy(_StrEnum): relay = 'relay' +#: The values of :obj:`RTCIceTransportPolicy`, which parameters taking it take too +RTCIceTransportPolicyValue = Literal['all', 'relay'] + + class RTCBundlePolicy(_StrEnum): """How media is bundled when the remote peer doesn't support bundling.""" @@ -201,12 +232,20 @@ class RTCBundlePolicy(_StrEnum): max_bundle = 'max-bundle' +#: The values of :obj:`RTCBundlePolicy`, which parameters taking it take too +RTCBundlePolicyValue = Literal['balanced', 'max-compat', 'max-bundle'] + + class RTCRtcpMuxPolicy(_StrEnum): """Whether RTCP is multiplexed with RTP, which is required.""" require = 'require' +#: The values of :obj:`RTCRtcpMuxPolicy`, which parameters taking it take too +RTCRtcpMuxPolicyValue = Literal['require'] + + class RTCRtpHeaderEncryptionPolicy(_StrEnum): """Whether RTP header extensions are encrypted with cryptex (RFC 9335).""" @@ -214,6 +253,10 @@ class RTCRtpHeaderEncryptionPolicy(_StrEnum): require = 'require' +#: The values of :obj:`RTCRtpHeaderEncryptionPolicy`, which parameters taking it take too +RTCRtpHeaderEncryptionPolicyValue = Literal['negotiate', 'require'] + + class RTCIceCandidateType(_StrEnum): """The type of an ICE candidate.""" @@ -246,6 +289,10 @@ class RTCIceServerTransportProtocol(_StrEnum): tls = 'tls' +#: The values of :obj:`RTCIceServerTransportProtocol`, which parameters taking it take too +RTCIceServerTransportProtocolValue = Literal['udp', 'tcp', 'tls'] + + class RTCErrorDetailType(_StrEnum): """The WebRTC-specific cause of an :obj:`webrtc.RTCError`.""" @@ -258,6 +305,18 @@ class RTCErrorDetailType(_StrEnum): hardware_encoder_error = 'hardware-encoder-error' +#: The values of :obj:`RTCErrorDetailType`, which parameters taking it take too +RTCErrorDetailTypeValue = Literal[ + 'data-channel-failure', + 'dtls-failure', + 'fingerprint-failure', + 'sctp-failure', + 'sdp-syntax-error', + 'hardware-encoder-not-available', + 'hardware-encoder-error', +] + + class BinaryType(_StrEnum): """What the binary messages of a :obj:`webrtc.RTCDataChannel` are delivered as.""" @@ -267,6 +326,10 @@ class BinaryType(_StrEnum): blob = 'blob' +#: The values of :obj:`BinaryType`, which parameters taking it take too +BinaryTypeValue = Literal['arraybuffer', 'blob'] + + class VideoPixelFormat(_StrEnum): """The layout of the pixels of a :obj:`webrtc.VideoFrame`. @@ -299,6 +362,34 @@ class VideoPixelFormat(_StrEnum): BGRX = 'BGRX' +#: The values of :obj:`VideoPixelFormat`, which parameters taking it take too +VideoPixelFormatValue = Literal[ + 'I420', + 'I420P10', + 'I420P12', + 'I420A', + 'I420AP10', + 'I420AP12', + 'I422', + 'I422P10', + 'I422P12', + 'I422A', + 'I422AP10', + 'I422AP12', + 'I444', + 'I444P10', + 'I444P12', + 'I444A', + 'I444AP10', + 'I444AP12', + 'NV12', + 'RGBA', + 'RGBX', + 'BGRA', + 'BGRX', +] + + class VideoColorPrimaries(_StrEnum): """The color primaries of a :obj:`webrtc.VideoColorSpace`.""" @@ -309,6 +400,10 @@ class VideoColorPrimaries(_StrEnum): smpte432 = 'smpte432' +#: The values of :obj:`VideoColorPrimaries`, which parameters taking it take too +VideoColorPrimariesValue = Literal['bt709', 'bt470bg', 'smpte170m', 'bt2020', 'smpte432'] + + class VideoTransferCharacteristics(_StrEnum): """The transfer characteristics of a :obj:`webrtc.VideoColorSpace`.""" @@ -320,6 +415,10 @@ class VideoTransferCharacteristics(_StrEnum): hlg = 'hlg' +#: The values of :obj:`VideoTransferCharacteristics`, which parameters taking it take too +VideoTransferCharacteristicsValue = Literal['bt709', 'smpte170m', 'iec61966-2-1', 'linear', 'pq', 'hlg'] + + class VideoMatrixCoefficients(_StrEnum): """The matrix coefficients of a :obj:`webrtc.VideoColorSpace`.""" @@ -330,6 +429,10 @@ class VideoMatrixCoefficients(_StrEnum): bt2020_ncl = 'bt2020-ncl' +#: The values of :obj:`VideoMatrixCoefficients`, which parameters taking it take too +VideoMatrixCoefficientsValue = Literal['rgb', 'bt709', 'bt470bg', 'smpte170m', 'bt2020-ncl'] + + class AlphaOption(_StrEnum): """Whether a :obj:`webrtc.VideoFrame` created from another one keeps its alpha channel.""" @@ -337,6 +440,10 @@ class AlphaOption(_StrEnum): discard = 'discard' +#: The values of :obj:`AlphaOption`, which parameters taking it take too +AlphaOptionValue = Literal['keep', 'discard'] + + class AudioSampleFormat(_StrEnum): """The type of the samples of an :obj:`webrtc.AudioData`, interleaved or in a plane per channel.""" @@ -348,3 +455,7 @@ class AudioSampleFormat(_StrEnum): s16_planar = 's16-planar' s32_planar = 's32-planar' f32_planar = 'f32-planar' + + +#: The values of :obj:`AudioSampleFormat`, which parameters taking it take too +AudioSampleFormatValue = Literal['u8', 's16', 's32', 'f32', 'u8-planar', 's16-planar', 's32-planar', 'f32-planar'] diff --git a/python-webrtc/python/webrtc/exceptions.py b/python-webrtc/python/webrtc/exceptions.py index 8772a1c..c3bb58a 100644 --- a/python-webrtc/python/webrtc/exceptions.py +++ b/python-webrtc/python/webrtc/exceptions.py @@ -10,12 +10,15 @@ from __future__ import annotations from dataclasses import dataclass -from typing import ClassVar +from typing import TYPE_CHECKING, ClassVar from webrtc import RTCErrorDetailType, wrtc from webrtc.models.dictionary import Dictionary from webrtc.utils.names import Alias, alias +if TYPE_CHECKING: + from webrtc.enums import RTCErrorDetailTypeValue + PythonWebRTCExceptionBase = wrtc.PythonWebRTCExceptionBase PythonWebRTCException = wrtc.PythonWebRTCException SdpParseException = wrtc.SdpParseException @@ -70,7 +73,7 @@ class OverconstrainedError(RTCException): """ def __init__(self, constraint: str, message: str = '') -> None: - super().__init__(message or f"The constraint {constraint} can't be satisfied") + super().__init__(message if message != '' else f"The constraint {constraint} can't be satisfied") self.constraint = constraint @@ -90,7 +93,7 @@ class RTCErrorInit(Dictionary): ValueError: If ``error_detail`` isn't a member of :obj:`RTCErrorDetailType`. """ - error_detail: RTCErrorDetailType + error_detail: RTCErrorDetailType | RTCErrorDetailTypeValue sdp_line_number: int | None = None sctp_cause_code: int | None = None received_alert: int | None = None @@ -101,7 +104,7 @@ def __post_init__(self) -> None: self.error_detail = RTCErrorDetailType(self.error_detail) #: Alias for :attr:`error_detail` - errorDetail: ClassVar[Alias[RTCErrorDetailType]] = alias('error_detail') + errorDetail: ClassVar[Alias[RTCErrorDetailType | RTCErrorDetailTypeValue]] = alias('error_detail') #: Alias for :attr:`sdp_line_number` sdpLineNumber: ClassVar[Alias[int | None]] = alias('sdp_line_number') #: Alias for :attr:`sctp_cause_code` @@ -125,7 +128,7 @@ class RTCError(OperationError): def __init__(self, init: RTCErrorInit, message: str = '') -> None: super().__init__(message) self.message = message - self.error_detail = init.error_detail + self.error_detail = RTCErrorDetailType(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 @@ -133,17 +136,17 @@ def __init__(self, init: RTCErrorInit, message: str = '') -> None: self.http_request_status_code = init.http_request_status_code #: Alias for :attr:`error_detail` - errorDetail = alias('error_detail') + errorDetail: ClassVar[Alias[RTCErrorDetailType]] = alias('error_detail') #: Alias for :attr:`sdp_line_number` - sdpLineNumber = alias('sdp_line_number') + sdpLineNumber: ClassVar[Alias[int | None]] = alias('sdp_line_number') #: Alias for :attr:`sctp_cause_code` - sctpCauseCode = alias('sctp_cause_code') + sctpCauseCode: ClassVar[Alias[int | None]] = alias('sctp_cause_code') #: Alias for :attr:`received_alert` - receivedAlert = alias('received_alert') + receivedAlert: ClassVar[Alias[int | None]] = alias('received_alert') #: Alias for :attr:`sent_alert` - sentAlert = alias('sent_alert') + sentAlert: ClassVar[Alias[int | None]] = alias('sent_alert') #: Alias for :attr:`http_request_status_code` - httpRequestStatusCode = alias('http_request_status_code') + httpRequestStatusCode: ClassVar[Alias[int | None]] = alias('http_request_status_code') _BY_RTC_ERROR_TYPE = { @@ -172,9 +175,18 @@ def _from_native( """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: init = RTCErrorInit( - detail or RTCErrorDetailType.data_channel_failure, + detail if detail is not None else RTCErrorDetailType.data_channel_failure, 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) + + +def _event_error(native: wrtc.RTCCallbackException) -> RTCError: + """The error of an ``error`` event: libwebrtc fails channels and transports with errors that carry a detail.""" + error = native.toPython() + if not isinstance(error, RTCError): + msg = f'an error event carries an RTCError, not {type(error).__name__}' + raise TypeError(msg) + return error diff --git a/python-webrtc/python/webrtc/functions/get_user_media.py b/python-webrtc/python/webrtc/functions/get_user_media.py index df19090..5822daf 100644 --- a/python-webrtc/python/webrtc/functions/get_user_media.py +++ b/python-webrtc/python/webrtc/functions/get_user_media.py @@ -64,10 +64,12 @@ def get_user_media( raise OverconstrainedError(failed, f"The constraint {failed} can't be satisfied") # the camera's defaults, within the constraints and the camera's capabilities capabilities = _CAMERA_CAPABILITIES - width = _selected(width, 640, capabilities.width) - height = _selected(height, 480, capabilities.height) - frame_rate = _selected(frame_rate, 30.0, capabilities.frame_rate) - stream = MediaStream._wrap(wrtc.getUserMedia(bool(audio), bool(video), width, height, float(frame_rate))) + selected_width = _selected(width, 640, capabilities.width) + selected_height = _selected(height, 480, capabilities.height) + selected_frame_rate = _selected(frame_rate, 30.0, capabilities.frame_rate) + stream = MediaStream._wrap( + wrtc.getUserMedia(bool(audio), bool(video), selected_width, selected_height, float(selected_frame_rate)) + ) for track in stream.get_video_tracks(): track._native_obj._constraints = constraints return stream diff --git a/python-webrtc/python/webrtc/interfaces/media_stream.py b/python-webrtc/python/webrtc/interfaces/media_stream.py index 8321cb9..3d37c5b 100644 --- a/python-webrtc/python/webrtc/interfaces/media_stream.py +++ b/python-webrtc/python/webrtc/interfaces/media_stream.py @@ -9,7 +9,9 @@ from __future__ import annotations -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, cast + +from typing_extensions import override from webrtc import MediaStreamTrack, MediaStreamTrackEvent, MediaType, WebRTCObject, wrtc from webrtc.utils.events import EventTarget @@ -37,30 +39,37 @@ class MediaStream(WebRTCObject[wrtc.MediaStream], EventTarget): _class = wrtc.MediaStream _events = ('addtrack', 'removetrack') + #: The native tracks, kept here: the native stream keeps them weakly + _tracks: list[wrtc.MediaStreamTrack] 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 []])) + super().__init__(wrtc.MediaStream.create([track._native_obj for track in tracks] if tracks is not None else [])) self._keep_tracks() @classmethod + @override def _wrap(cls, item: wrtc.MediaStream) -> Self: stream = super()._wrap(item) stream._keep_tracks() return stream - def _keep_tracks(self) -> list[wrtc.MediaStreamTrack]: - """The native tracks, kept here: the native stream keeps them weakly.""" + def _keep_tracks(self) -> None: self._tracks = self._native_obj.getTracks() + + def _kept_tracks(self) -> list[wrtc.MediaStreamTrack]: + self._keep_tracks() return self._tracks + @override def _on_event(self, name: str, *_args: object) -> None: if name in {'addtrack', 'removetrack'}: self._keep_tracks() + @override def _create_event(self, name: str, *args: object) -> webrtc.Event | None: - (track,) = args + (track,) = cast('tuple[wrtc.MediaStreamTrack]', args) return MediaStreamTrackEvent(name, MediaStreamTrack._wrap(track), target=self) @property @@ -79,7 +88,7 @@ def get_audio_tracks(self) -> list[webrtc.MediaStreamTrack]: 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]) + return MediaStreamTrack._wrap_many([t for t in self._kept_tracks() if t.kind == MediaType.audio]) def get_video_tracks(self) -> list[webrtc.MediaStreamTrack]: """Returns the video tracks of the stream, in no defined order. @@ -87,7 +96,7 @@ def get_video_tracks(self) -> list[webrtc.MediaStreamTrack]: 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]) + return MediaStreamTrack._wrap_many([t for t in self._kept_tracks() if t.kind == MediaType.video]) def get_tracks(self) -> list[webrtc.MediaStreamTrack]: """Returns all the tracks of the stream, in no defined order. @@ -95,7 +104,7 @@ def get_tracks(self) -> list[webrtc.MediaStreamTrack]: Returns: :obj:`list` of :obj:`webrtc.MediaStreamTrack`: The tracks. """ - return MediaStreamTrack._wrap_many(self._keep_tracks()) + return MediaStreamTrack._wrap_many(self._kept_tracks()) 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. diff --git a/python-webrtc/python/webrtc/interfaces/media_stream_track.py b/python-webrtc/python/webrtc/interfaces/media_stream_track.py index 666f1c8..c8b3353 100644 --- a/python-webrtc/python/webrtc/interfaces/media_stream_track.py +++ b/python-webrtc/python/webrtc/interfaces/media_stream_track.py @@ -11,7 +11,9 @@ import asyncio import math -from typing import TYPE_CHECKING, Any, Union +from typing import TYPE_CHECKING, Union, cast + +from typing_extensions import override from webrtc import ( ConstrainBooleanOrDOMStringParameters, @@ -210,10 +212,11 @@ class MediaStreamTrack(WebRTCObject[wrtc.MediaStreamTrack], EventTarget): _class = wrtc.MediaStreamTrack _events = ('mute', 'unmute', 'ended') + @override def _on_event(self, name: str, *args: object) -> None: # muted changes along with the events if name in {'mute', 'unmute'}: - (muted,) = args + (muted,) = cast('tuple[bool]', args) self._native_obj._surfaceMuted(muted) elif name == 'ended': self._native_obj._surfaceEnded() @@ -277,16 +280,18 @@ def get_settings(self) -> MediaTrackSettings: Returns: :obj:`webrtc.MediaTrackSettings`: The settings. """ - native: dict[str, Any] = self._native_obj._settings() + native = self._native_obj._settings() settings = MediaTrackSettings() - if 'width' in native: - settings.width, settings.height = native['width'], native['height'] - settings.aspect_ratio = native['width'] / native['height'] if native['height'] else None + # the size, and the format of the audio, come together + width, height = native.get('width'), native.get('height') + if width is not None and height is not None: + settings.width, settings.height = width, height + settings.aspect_ratio = width / height if height != 0 else None settings.frame_rate = native.get('frame_rate') if 'sample_rate' in native: - settings.sample_rate = native['sample_rate'] - settings.sample_size = native['sample_size'] - settings.channel_count = native['channel_count'] + settings.sample_rate = native.get('sample_rate') + settings.sample_size = native.get('sample_size') + settings.channel_count = native.get('channel_count') device = native.get('device') if device == 'camera': settings.resize_mode = 'none' @@ -343,7 +348,7 @@ def apply_constraints(self, constraints: MediaTrackConstraints | None = None) -> return future def _apply_constraints(self, constraints: MediaTrackConstraints) -> None: - advanced = list(constraints.advanced or ()) + advanced: list[MediaTrackConstraintSet] = list(constraints.advanced) if constraints.advanced is not None else [] for constraint_set in [constraints, *advanced]: _check_numbers(constraint_set) if self.ready_state == 'ended': @@ -364,7 +369,7 @@ def _apply_constraints(self, constraints: MediaTrackConstraints) -> None: height = _selected(constraint_set.height, height, capabilities.height) frame_rate = _selected(constraint_set.frame_rate, frame_rate, capabilities.frame_rate) if (width, height, frame_rate) != camera: - self._native_obj._reconfigureCamera(int(width), int(height), float(frame_rate)) + _ = self._native_obj._reconfigureCamera(int(width), int(height), float(frame_rate)) self._native_obj._constraints = constraints def clone(self) -> webrtc.MediaStreamTrack: 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 97599a0..ce59e30 100644 --- a/python-webrtc/python/webrtc/interfaces/media_stream_track_processor.py +++ b/python-webrtc/python/webrtc/interfaces/media_stream_track_processor.py @@ -10,7 +10,9 @@ from __future__ import annotations from dataclasses import dataclass -from typing import TYPE_CHECKING, Any, ClassVar +from typing import TYPE_CHECKING, ClassVar + +from typing_extensions import override from webrtc import AudioData, MediaStreamTrack, MediaType, VideoFrame, WebRTCObject, wrtc from webrtc.models.dictionary import Dictionary @@ -26,6 +28,8 @@ #: How many 10 ms chunks of audio are queued for reads, as Chrome does DEFAULT_AUDIO_BUFFER_SIZE = 10 _MAX_BUFFER_SIZE = 65535 +# the members of a native item of video: the buffer, timestamp, rotation and RTP timestamp (6 for audio) +_VIDEO_ITEM_SIZE = 4 @dataclass @@ -52,12 +56,12 @@ class _TrackSource: def __init__(self, processor: MediaStreamTrackProcessor) -> None: self._processor = processor - self._controller: ReadableStreamDefaultController | None = None + self._controller: ReadableStreamDefaultController[VideoFrame | AudioData] | None = None - def start(self, controller: ReadableStreamDefaultController) -> None: + def start(self, controller: ReadableStreamDefaultController[VideoFrame | AudioData]) -> None: self._controller = controller - def pull(self, _controller: ReadableStreamDefaultController) -> None: + def pull(self, _controller: ReadableStreamDefaultController[VideoFrame | AudioData]) -> None: native = self._processor._native_obj created_outside_loop = native._listeners is None # the native events go to the loop reading @@ -75,7 +79,10 @@ def deliver(self) -> None: native = self._processor._native_obj stream = self._processor._readable controller = self._controller - while stream._state == 'readable' and stream._reader is not None and stream._reader._read_requests: + if controller is None: + msg = 'the stream of the processor has not started' + raise RuntimeError(msg) + while stream._state == 'readable' and stream._reader is not None and len(stream._reader._read_requests) > 0: item = native.read() if item is None: break @@ -123,26 +130,29 @@ def __init__(self, init: MediaStreamTrackProcessorInit) -> None: 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))) + super().__init__(wrtc.MediaStreamTrackProcessor(track._native_obj, max(1, max_buffer_size))) # the native processor doesn't keep the track, Python does self._track = track - self._video = video self._source = _TrackSource(self) - self._readable = ReadableStream(self._source, high_water_mark=0) + self._readable: ReadableStream[VideoFrame | AudioData] = ReadableStream(self._source, high_water_mark=0) self._attach() + @override 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[Any, ...]) -> VideoFrame | AudioData: - if self._video: + @staticmethod + def _wrap_media( + item: tuple[wrtc.VideoFrameBuffer, int, int, int] | tuple[bytes, int, int, int, int, int], + ) -> VideoFrame | AudioData: + if len(item) == _VIDEO_ITEM_SIZE: return VideoFrame._from_native(item) return AudioData._from_native(item) @property - def readable(self) -> ReadableStream: + def readable(self) -> ReadableStream[VideoFrame | AudioData]: """:obj:`webrtc.ReadableStream`: The media of the track.""" return self._readable diff --git a/python-webrtc/python/webrtc/interfaces/rtc_data_channel.py b/python-webrtc/python/webrtc/interfaces/rtc_data_channel.py index 09f58f5..7c646a3 100644 --- a/python-webrtc/python/webrtc/interfaces/rtc_data_channel.py +++ b/python-webrtc/python/webrtc/interfaces/rtc_data_channel.py @@ -10,7 +10,9 @@ from __future__ import annotations from dataclasses import dataclass -from typing import TYPE_CHECKING, ClassVar +from typing import TYPE_CHECKING, ClassVar, cast + +from typing_extensions import override from webrtc import ( BinaryType, @@ -22,12 +24,14 @@ WebRTCObject, wrtc, ) +from webrtc.exceptions import _event_error from webrtc.models.dictionary import Dictionary from webrtc.utils.events import EventTarget from webrtc.utils.names import Alias, alias if TYPE_CHECKING: import webrtc + from webrtc.enums import BinaryTypeValue, RTCPriorityTypeValue #: The maximum of an unsigned short, which limits the members of the init MAX_UNSIGNED_SHORT = 65535 @@ -68,7 +72,7 @@ class RTCDataChannelInit(Dictionary): protocol: str = '' negotiated: bool = False id: int | None = None - priority: RTCPriorityType | str = RTCPriorityType.low + priority: RTCPriorityType | RTCPriorityTypeValue = RTCPriorityType.low def _check(self) -> None: """Checks the members, as the specification requires. @@ -122,31 +126,33 @@ class RTCDataChannel(WebRTCObject[wrtc.RTCDataChannel], EventTarget): _class = wrtc.RTCDataChannel _events = ('open', 'message', 'bufferedamountlow', 'error', 'closing', 'close') + @override def _on_event(self, name: str, *args: object) -> None: # readyState changes along with the events if name in {'open', 'closing', 'close'}: - (state,) = args + (state,) = cast('tuple[RTCDataChannelState]', args) self._native_obj._surfaceState(state) elif name == '_sent': - (size,) = args + (size,) = cast('tuple[int]', args) if self._native_obj._decreaseBufferedAmount(size): # in the same task as the decrease, before anything that arrived meanwhile self._dispatch('bufferedamountlow') + @override 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 if name == 'message': - (message,) = args - data = message.data + (message,) = cast('tuple[wrtc.DataChannelMessage]', args) + data: str | bytes | Blob = message.data # binary_type as of delivery, per the specification if isinstance(data, bytes) and self._native_obj.binaryType == BinaryType.blob: data = Blob([data]) return MessageEvent(name, data, target=self) if name == 'error': - (error,) = args - return RTCErrorEvent(name, error.toPython(), target=self) + (error,) = cast('tuple[wrtc.RTCCallbackException]', args) + return RTCErrorEvent(name, _event_error(error), target=self) return super()._create_event(name, *args) @property @@ -220,7 +226,7 @@ def binary_type(self) -> BinaryType: return BinaryType(self._native_obj.binaryType) @binary_type.setter - def binary_type(self, value: BinaryType | str) -> None: + def binary_type(self, value: BinaryType | BinaryTypeValue) -> None: self._native_obj.binaryType = BinaryType(value).value def send(self, data: str | bytes | bytearray | memoryview | Blob) -> None: diff --git a/python-webrtc/python/webrtc/interfaces/rtc_dtls_transport.py b/python-webrtc/python/webrtc/interfaces/rtc_dtls_transport.py index 62e8b8a..66cebc4 100644 --- a/python-webrtc/python/webrtc/interfaces/rtc_dtls_transport.py +++ b/python-webrtc/python/webrtc/interfaces/rtc_dtls_transport.py @@ -9,8 +9,13 @@ from __future__ import annotations +from typing import cast + +from typing_extensions import override + import webrtc -from webrtc import RTCErrorEvent, WebRTCObject, wrtc +from webrtc import DtlsTransportState, RTCErrorEvent, WebRTCObject, wrtc +from webrtc.exceptions import _event_error from webrtc.utils.events import EventTarget @@ -28,16 +33,18 @@ class RTCDtlsTransport(WebRTCObject[wrtc.RTCDtlsTransport], EventTarget): _class = wrtc.RTCDtlsTransport _events = ('statechange', 'error') + @override def _on_event(self, name: str, *args: object) -> None: # the state changes along with its event if name == 'statechange': - (state,) = args + (state,) = cast('tuple[DtlsTransportState]', args) self._native_obj._surfaceState(state) + @override def _create_event(self, name: str, *args: object) -> webrtc.Event | None: if name == 'error': - (error,) = args - return RTCErrorEvent(name, error.toPython(), target=self) + (error,) = cast('tuple[wrtc.RTCCallbackException]', args) + return RTCErrorEvent(name, _event_error(error), target=self) return super()._create_event(name, *args) @property diff --git a/python-webrtc/python/webrtc/interfaces/rtc_dtmf_sender.py b/python-webrtc/python/webrtc/interfaces/rtc_dtmf_sender.py index ff294da..ac19993 100644 --- a/python-webrtc/python/webrtc/interfaces/rtc_dtmf_sender.py +++ b/python-webrtc/python/webrtc/interfaces/rtc_dtmf_sender.py @@ -10,7 +10,9 @@ from __future__ import annotations import re -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, cast + +from typing_extensions import override from webrtc import InvalidCharacterError, RTCDTMFToneChangeEvent, WebRTCObject, wrtc from webrtc.utils.events import EventTarget @@ -32,13 +34,15 @@ class RTCDTMFSender(WebRTCObject[wrtc.RTCDTMFSender], EventTarget): _class = wrtc.RTCDTMFSender _events = ('tonechange',) - def _on_event(self, _name: str, *args: object) -> None: - _, tone_buffer, insertion = args + @override + def _on_event(self, name: str, *args: object) -> None: + _, tone_buffer, insertion = cast('tuple[str, str, int]', args) # the tone buffer is shortened along with the event self._native_obj._surfaceBuffer(tone_buffer, insertion) + @override def _create_event(self, name: str, *args: object) -> webrtc.Event | None: - tone, _, _ = args + tone, _, _ = cast('tuple[str, str, int]', args) return RTCDTMFToneChangeEvent(name, tone, target=self) def insert_dtmf(self, tones: str, duration: int = 100, inter_tone_gap: int = 70) -> None: @@ -54,7 +58,7 @@ def insert_dtmf(self, tones: str, duration: int = 100, inter_tone_gap: int = 70) 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): + if _TONES.fullmatch(tones) is None: msg = f'{tones!r} has characters that are not DTMF tones' raise InvalidCharacterError(msg) duration = min(max(int(duration), 40), 6000) diff --git a/python-webrtc/python/webrtc/interfaces/rtc_ice_transport.py b/python-webrtc/python/webrtc/interfaces/rtc_ice_transport.py index 6f5c22e..bb51a31 100644 --- a/python-webrtc/python/webrtc/interfaces/rtc_ice_transport.py +++ b/python-webrtc/python/webrtc/interfaces/rtc_ice_transport.py @@ -11,7 +11,9 @@ import re import weakref -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, cast + +from typing_extensions import override from webrtc import ( CricketIceGatheringState, @@ -77,22 +79,26 @@ def _candidate_of(self, native: wrtc.IceCandidateInit) -> webrtc.RTCIceCandidate 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) + _ = _candidates.setdefault(self._native_obj, {}).setdefault(candidate.candidate, candidate) + @override def _on_event(self, name: str, *args: object) -> None: # the states change along with their events if name == 'statechange': - (state,) = args + (state,) = cast('tuple[RTCIceTransportState]', args) self._native_obj._surfaceState(state) elif name == 'gatheringstatechange': - (state,) = args - self._native_obj._surfaceGatheringState(state) - elif name == 'icecandidate' and args and args[0] is not None: + (gathering_state,) = cast('tuple[CricketIceGatheringState]', args) + self._native_obj._surfaceGatheringState(gathering_state) + elif name == 'icecandidate' and len(args) > 0 and args[0] is not None: self._native_obj._surfaceCandidate() + @override 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 + # the end of candidates has none + native = cast('wrtc.IceCandidateInit | None', args[0]) if len(args) > 0 else None + candidate = self._candidate_of(native) if native is not None else None return RTCPeerConnectionIceEvent(name, candidate, None, target=self) return super()._create_event(name, *args) @@ -109,7 +115,7 @@ def _check_open(self, operation: str) -> None: def gather( self, - gather_policy: webrtc.RTCIceTransportPolicy | str = 'all', + gather_policy: webrtc.RTCIceTransportPolicy | webrtc.RTCIceTransportPolicyValue = 'all', ice_servers: Sequence[webrtc.RTCIceServer] | None = None, ) -> None: """Gathers the candidates of a standalone transport, sent in ``icecandidate`` events. @@ -128,7 +134,9 @@ def gather( if self.gathering_state != CricketIceGatheringState.new: msg = 'The transport gathers its candidates already' raise InvalidStateError(msg) - self._native_obj.gather(gather_policy, RTCIceServer._to_native_list(ice_servers or ())) + self._native_obj.gather( + gather_policy, RTCIceServer._to_native_list(ice_servers if ice_servers is not None else ()) + ) def start( self, @@ -151,10 +159,10 @@ def start( ValueError: If the role is neither controlling nor controlled. """ self._check_open('start') - if not _UFRAG.fullmatch(remote_parameters.username_fragment): + if _UFRAG.fullmatch(remote_parameters.username_fragment) is None: msg = f'{remote_parameters.username_fragment!r} is not a valid ICE username fragment' raise InvalidSyntaxError(msg) - if not _PASSWORD.fullmatch(remote_parameters.password): + if _PASSWORD.fullmatch(remote_parameters.password) is None: msg = 'the ICE password is not valid' raise InvalidSyntaxError(msg) if role not in {RTCIceRole.controlling, RTCIceRole.controlled}: @@ -177,7 +185,10 @@ def add_remote_candidate(self, candidate: webrtc.RTCIceCandidate | webrtc.RTCIce if not isinstance(candidate, RTCIceCandidate): candidate = RTCIceCandidate(*RTCIceCandidate._members_of(candidate)) self._native_obj.addRemoteCandidate( - candidate.candidate, candidate.sdp_mid or '', candidate.sdp_m_line_index or 0, candidate.username_fragment + candidate.candidate, + candidate.sdp_mid if candidate.sdp_mid is not None else '', + candidate.sdp_m_line_index if candidate.sdp_m_line_index is not None else 0, + candidate.username_fragment, ) self._remember(candidate) diff --git a/python-webrtc/python/webrtc/interfaces/rtc_peer_connection.py b/python-webrtc/python/webrtc/interfaces/rtc_peer_connection.py index b470442..6cbfa0d 100644 --- a/python-webrtc/python/webrtc/interfaces/rtc_peer_connection.py +++ b/python-webrtc/python/webrtc/interfaces/rtc_peer_connection.py @@ -11,10 +11,13 @@ import asyncio import re -from typing import TYPE_CHECKING, ClassVar, Union +from typing import TYPE_CHECKING, ClassVar, Literal, Union, cast, overload + +from typing_extensions import override import webrtc from webrtc import ( + CricketIceGatheringState, Event, InvalidAccessError, InvalidStateError, @@ -47,9 +50,12 @@ if TYPE_CHECKING: from contextlib import AbstractAsyncContextManager + from typing import Callable from typing_extensions import Self + from webrtc.models.rtc_certificate import AlgorithmIdentifier + #: A description, as the methods that set one take it _Description = Union[RTCSessionDescription, RTCSessionDescriptionInit] @@ -133,10 +139,11 @@ class RTCPeerConnection(WebRTCObject[wrtc.RTCPeerConnection], EventTarget): _deferred_negotiation_id: int | None = None def __init__(self, configuration: webrtc.RTCConfiguration | None = None) -> None: - super().__init__(self._class(configuration._to_native() if configuration is not None else None)) + super().__init__(wrtc.RTCPeerConnection(configuration._to_native() if configuration is not None else None)) self._attach() @classmethod + @override 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 @@ -157,7 +164,7 @@ def _chain_emptied(self) -> None: # negotiationneeded fires now, if it's still needed if self._deferred_negotiation_id is not None: event_id, self._deferred_negotiation_id = self._deferred_negotiation_id, None - asyncio.get_running_loop().call_soon(self._dispatch, 'negotiationneeded', event_id) + _ = asyncio.get_running_loop().call_soon(self._dispatch, 'negotiationneeded', event_id) def _check_state(self, operation: str, *allowed: RTCSignalingState) -> None: """Checks the connection isn't closed, and is in one of the allowed states if they're given. @@ -169,13 +176,14 @@ def _check_state(self, operation: str, *allowed: RTCSignalingState) -> None: if state == RTCSignalingState.closed: msg = f"Can not {operation}: the RTCPeerConnection's signalingState is 'closed'" raise InvalidStateError(msg) - if allowed and state not in allowed: + if len(allowed) > 0 and state not in allowed: msg = f'Can not {operation} in the {state} signaling state' raise InvalidStateError(msg) + @override def _on_event(self, name: str, *args: object) -> None: if name == '_gatheringcomplete': - transports, state = args + transports, state = cast('tuple[list[wrtc.RTCIceTransport], webrtc.RTCIceGatheringState]', args) self._complete_gathering(transports, state) return # the state attributes change along with their events @@ -184,17 +192,17 @@ def _on_event(self, name: str, *args: object) -> None: state = args[0] getattr(self._native_obj, surface)(state) if name == 'signalingstatechange': - _, descriptions = args + _, descriptions = cast('tuple[RTCSignalingState, int]', args) # the descriptions as the change left them self._native_obj._applyDescriptions(descriptions) 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 - webrtc.RTCDataChannel._wrap(channel) + (channel,) = cast('tuple[wrtc.RTCDataChannel]', args) + _ = webrtc.RTCDataChannel._wrap(channel) # the events of the channel follow the handlers of this one - TaskQueue.post_to_running(channel._release) + _ = 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. @@ -203,7 +211,7 @@ def _complete_gathering(self, transports: list[wrtc.RTCIceTransport], state: web """ ice_transports = webrtc.RTCIceTransport._wrap_many(transports) for ice_transport in ice_transports: - ice_transport._native_obj._surfaceGatheringState(state) + ice_transport._native_obj._surfaceGatheringState(CricketIceGatheringState(state)) self._native_obj._surfaceIceGatheringState(state) self._native_obj._refreshDescriptions() for ice_transport in ice_transports: @@ -212,14 +220,16 @@ def _complete_gathering(self, transports: list[wrtc.RTCIceTransport], state: web # the end of candidates is an icecandidate event without a candidate self._dispatch('icecandidate') + @override 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: + creator_name = self._EVENT_CREATORS.get(name) + if creator_name is None: return super()._create_event(name, *args) - return getattr(self, creator)(*args) + creator: Callable[..., webrtc.Event | None] = getattr(self, creator_name) + return creator(*args) 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 @@ -237,9 +247,15 @@ def _ice_candidate_event(self, candidate: wrtc.IceCandidateInit | None = None) - return RTCPeerConnectionIceEvent('icecandidate', RTCIceCandidate(**kwargs), kwargs['url'], target=self) def _ice_candidate_error_event(self, *native: object) -> webrtc.Event: - address, port, url, error_code, error_text = native + address, port, url, error_code, error_text = cast('tuple[str, int, str, int, str]', native) return RTCPeerConnectionIceErrorEvent( - 'icecandidateerror', address or None, port or None, url, error_code, error_text, target=self + 'icecandidateerror', + address if address != '' else None, + port if port != 0 else None, + url, + error_code, + error_text, + target=self, ) def _data_channel_event(self, channel: wrtc.RTCDataChannel) -> webrtc.Event: @@ -274,7 +290,7 @@ def _apply_legacy_offer_option(self, kind: webrtc.MediaType, *, receive: bool | elif transceiver.direction == directions.recvonly: transceiver.direction = directions.inactive elif not any(t.direction in {directions.sendrecv, directions.recvonly} for t in transceivers): - self.add_transceiver(kind, RTCRtpTransceiverInit(direction=directions.recvonly)) + _ = self.add_transceiver(kind, RTCRtpTransceiverInit(direction=directions.recvonly)) def _completed_description(self) -> None: """The success task of setting a description.""" @@ -412,7 +428,7 @@ def add_track( :obj:`webrtc.RTCRtpSender`: The :obj:`webrtc.RTCRtpSender` object which will be used to transmit the media data. """ - if not stream: + if stream is None or (isinstance(stream, list) and len(stream) == 0): sender = self._native_obj.addTrack(track._native_obj, None) elif isinstance(stream, list): native_objects = [s._native_obj for s in stream] @@ -424,7 +440,7 @@ def add_track( def add_transceiver( self, - track_or_kind: webrtc.MediaStreamTrack | webrtc.MediaType, + track_or_kind: webrtc.MediaStreamTrack | webrtc.MediaType | webrtc.MediaTypeValue, init: webrtc.RTCRtpTransceiverInit | None = None, ) -> webrtc.RTCRtpTransceiver: """Creates a new :obj:`webrtc.RTCRtpTransceiver` and adds it to the transceivers of the connection. @@ -540,7 +556,7 @@ async def add_ice_candidate( 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: + if candidate_str != '' and sdp_mid is None and sdp_m_line_index is None: msg = 'sdp_mid and sdp_m_line_index are both None' raise TypeError(msg) @@ -616,7 +632,7 @@ async def get_stats(self, selector: webrtc.MediaStreamTrack | None = None) -> we @staticmethod async def generate_certificate( - algorithm: webrtc.models.rtc_certificate.AlgorithmIdentifier = 'ECDSA', expires: float | None = None + algorithm: AlgorithmIdentifier = 'ECDSA', expires: float | None = None ) -> webrtc.RTCCertificate: """Generates a certificate for :attr:`webrtc.RTCConfiguration.certificates`. @@ -805,6 +821,18 @@ def ice_gathering_state(self) -> webrtc.RTCIceGatheringState: setConfiguration = set_configuration +@overload +def _description_init( + description: _Description, *, allow_implicit: Literal[False] +) -> wrtc.RTCSessionDescriptionInit: ... + + +@overload +def _description_init( + description: _Description | RTCLocalSessionDescriptionInit | None, *, allow_implicit: Literal[True] +) -> wrtc.RTCSessionDescriptionInit | None: ... + + def _description_init( description: _Description | RTCLocalSessionDescriptionInit | None, *, allow_implicit: bool ) -> wrtc.RTCSessionDescriptionInit | None: @@ -812,7 +840,7 @@ def _description_init( if isinstance(description, RTCSessionDescription): return description._native_obj.init if allow_implicit and isinstance(description, RTCLocalSessionDescriptionInit): - if description.type is None and description.sdp: + if description.type is None and description.sdp != '': msg = 'the type of a description is required' raise TypeError(msg) description = None if description.type is None else RTCSessionDescriptionInit(description.type, description.sdp) @@ -837,7 +865,7 @@ def _check_send_encodings(encodings: list[webrtc.RTCRtpEncodingParameters], kind """ rids = [e.rid for e in encodings] for rid in rids: - if rid is not None and not _RID.fullmatch(rid): + if rid is not None and _RID.fullmatch(rid) is None: 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)): @@ -845,9 +873,9 @@ def _check_send_encodings(encodings: list[webrtc.RTCRtpEncodingParameters], kind raise ValueError(msg) codecs = [e.codec for e in encodings if e.codec is not None] - if codecs: + if len(codecs) > 0: capabilities = webrtc.RTCRtpSender.get_capabilities(kind) - supported = capabilities.codecs if capabilities is not None else [] + supported: list[RTCRtpCodec] = capabilities.codecs if capabilities is not None else [] for codec in codecs: if not any(RTCRtpCodec._matches(c, codec) for c in supported): msg = f'{codec.mime_type} can not be sent' diff --git a/python-webrtc/python/webrtc/interfaces/rtc_rtp_receiver.py b/python-webrtc/python/webrtc/interfaces/rtc_rtp_receiver.py index c31fbe2..6eddf9d 100644 --- a/python-webrtc/python/webrtc/interfaces/rtc_rtp_receiver.py +++ b/python-webrtc/python/webrtc/interfaces/rtc_rtp_receiver.py @@ -9,6 +9,8 @@ from __future__ import annotations +from typing import TypeVar + import webrtc from webrtc import ( InvalidRangeError, @@ -22,6 +24,8 @@ ) from webrtc.utils.native_calls import call_native +_SourceT = TypeVar('_SourceT', bound=RTCRtpContributingSource) + #: The maximum jitter_buffer_target, in milliseconds _MAX_JITTER_BUFFER_TARGET = 4000 @@ -31,8 +35,8 @@ class RTCRtpReceiver(WebRTCObject[wrtc.RTCRtpReceiver]): _class = wrtc.RTCRtpReceiver - def _sources(self, *, synchronization: bool) -> list[webrtc.RTCRtpContributingSource]: - cls = RTCRtpSynchronizationSource if synchronization else RTCRtpContributingSource + def _sources(self, cls: type[_SourceT]) -> list[_SourceT]: + synchronization = cls is RTCRtpSynchronizationSource # 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] @@ -70,7 +74,7 @@ def get_parameters(self) -> webrtc.RTCRtpReceiveParameters: return RTCRtpReceiveParameters._from_native(self._native_obj.getParameters()) @staticmethod - def get_capabilities(kind: webrtc.MediaType) -> webrtc.RTCRtpCapabilities | None: + def get_capabilities(kind: webrtc.MediaType | webrtc.MediaTypeValue) -> webrtc.RTCRtpCapabilities | None: """Returns the codecs and header extensions receivers of a kind support. Args: @@ -98,7 +102,7 @@ def get_synchronization_sources(self) -> list[webrtc.RTCRtpSynchronizationSource Returns: :obj:`list` of :obj:`webrtc.RTCRtpSynchronizationSource`: The sources, the most recent first. """ - return self._sources(synchronization=True) + return self._sources(RTCRtpSynchronizationSource) def get_contributing_sources(self) -> list[webrtc.RTCRtpContributingSource]: """Returns the contributing sources (CSRCs) of the media received in the last 10 seconds. @@ -108,7 +112,7 @@ def get_contributing_sources(self) -> list[webrtc.RTCRtpContributingSource]: Returns: :obj:`list` of :obj:`webrtc.RTCRtpContributingSource`: The sources, the most recent first. """ - return self._sources(synchronization=False) + return self._sources(RTCRtpContributingSource) #: Alias for :attr:`get_stats` getStats = get_stats diff --git a/python-webrtc/python/webrtc/interfaces/rtc_rtp_sender.py b/python-webrtc/python/webrtc/interfaces/rtc_rtp_sender.py index 17d58bd..241f276 100644 --- a/python-webrtc/python/webrtc/interfaces/rtc_rtp_sender.py +++ b/python-webrtc/python/webrtc/interfaces/rtc_rtp_sender.py @@ -72,7 +72,7 @@ def get_parameters(self) -> webrtc.RTCRtpSendParameters: if self.kind == MediaType.video: _default_scale_resolution_down_by(parameters.encodings) # they expire when the current task (with the code it resumed) is over, never without a loop - TaskQueue.post_to_running(self._native_obj._expireParameters, parameters.transaction_id, after_ready=True) + _ = TaskQueue.post_to_running(self._native_obj._expireParameters, parameters.transaction_id, after_ready=True) return parameters async def set_parameters( @@ -108,8 +108,8 @@ async def set_parameters( # a copy (pybind returns one): changed, then set back encodings = last.encodings for native, encoding in zip(encodings, parameters.encodings): - encoding._for_kind(kind)._apply(native) - for native, key_frame in zip(encodings, key_frames or ()): + _ = encoding._for_kind(kind)._apply(native) + for native, key_frame in zip(encodings, key_frames if key_frames is not None else ()): native.requestKeyFrame = bool(key_frame) last.encodings = encodings last.degradationPreference = parameters.degradation_preference @@ -155,14 +155,14 @@ def set_streams(self, *streams: webrtc.MediaStream) -> None: Args: *streams (:obj:`webrtc.MediaStream`): The streams, none to associate the track with no stream. """ - ids = [] + ids: list[str] = [] for stream in streams: if stream.id not in ids: ids.append(stream.id) self._native_obj.setStreams(ids) @staticmethod - def get_capabilities(kind: webrtc.MediaType) -> webrtc.RTCRtpCapabilities | None: + def get_capabilities(kind: webrtc.MediaType | webrtc.MediaTypeValue) -> webrtc.RTCRtpCapabilities | None: """Returns the codecs and header extensions senders of a kind support. Args: diff --git a/python-webrtc/python/webrtc/interfaces/rtc_rtp_transceiver.py b/python-webrtc/python/webrtc/interfaces/rtc_rtp_transceiver.py index 8ff4b06..638b9c8 100644 --- a/python-webrtc/python/webrtc/interfaces/rtc_rtp_transceiver.py +++ b/python-webrtc/python/webrtc/interfaces/rtc_rtp_transceiver.py @@ -56,7 +56,7 @@ def direction(self) -> webrtc.TransceiverDirection: return self._native_obj.direction @direction.setter - def direction(self, new_direction: webrtc.TransceiverDirection) -> None: + def direction(self, new_direction: webrtc.TransceiverDirection | webrtc.TransceiverDirectionValue) -> None: self._native_obj.direction = new_direction @property @@ -91,11 +91,11 @@ def set_codec_preferences(self, codecs: list[webrtc.RTCRtpCodec]) -> None: (like RTX or FEC) are given. """ kind = self.kind - natives = [] + natives: list[wrtc.RtpCodecCapability] = [] for source in (wrtc.RTCRtpReceiver.getCapabilities(kind), wrtc.RTCRtpSender.getCapabilities(kind)): natives.extend(source.codecs if source is not None else []) - preferences = [] + preferences: list[wrtc.RtpCodecCapability] = [] for codec in codecs: native = next((n for n in natives if RTCRtpCodec._from_native(n)._matches(codec)), None) if native is None: @@ -128,9 +128,9 @@ def set_header_extensions_to_negotiate(self, extensions: list[webrtc.RTCRtpHeade extension is stopped. """ current = {e.uri: e for e in self._native_obj.getHeaderExtensionsToNegotiate()} - natives = [] + natives: list[wrtc.RtpHeaderExtensionCapability] = [] for extension in extensions: - if not extension.uri: + if extension.uri == '': msg = 'the URI of a header extension must not be empty' raise ValueError(msg) native = wrtc.RtpHeaderExtensionCapability() diff --git a/python-webrtc/python/webrtc/interfaces/rtc_sctp_transport.py b/python-webrtc/python/webrtc/interfaces/rtc_sctp_transport.py index 31fd59d..b0682eb 100644 --- a/python-webrtc/python/webrtc/interfaces/rtc_sctp_transport.py +++ b/python-webrtc/python/webrtc/interfaces/rtc_sctp_transport.py @@ -9,6 +9,10 @@ from __future__ import annotations +from typing import cast + +from typing_extensions import override + import webrtc from webrtc import WebRTCObject, wrtc from webrtc.utils.events import EventTarget @@ -27,8 +31,9 @@ class RTCSctpTransport(WebRTCObject[wrtc.RTCSctpTransport], EventTarget): _class = wrtc.RTCSctpTransport _events = ('statechange',) - def _on_event(self, _name: str, *args: object) -> None: - (state,) = args + @override + def _on_event(self, name: str, *args: object) -> None: + (state,) = cast('tuple[webrtc.SctpTransportState]', args) # the state changes along with its event self._native_obj._surfaceState(state) diff --git a/python-webrtc/python/webrtc/interfaces/track_generator.py b/python-webrtc/python/webrtc/interfaces/track_generator.py index c8610ea..a94f7f7 100644 --- a/python-webrtc/python/webrtc/interfaces/track_generator.py +++ b/python-webrtc/python/webrtc/interfaces/track_generator.py @@ -19,6 +19,7 @@ from webrtc.streams import WritableStream if TYPE_CHECKING: + from webrtc.enums import MediaTypeValue from webrtc.streams import WritableStreamDefaultController @@ -52,12 +53,14 @@ def _write_audio(self, data: object) -> None: msg = 'The data is closed' raise TypeError(msg) audio = data._take() - if audio.format == AudioSampleFormat.s16: - samples = audio._data + data_bytes = audio._data + if audio.format == AudioSampleFormat.s16 and data_bytes is not None: + samples = data_bytes else: - samples = bytearray(audio.number_of_frames * audio.number_of_channels * 2) - audio.copy_to(samples, AudioDataCopyToOptions(plane_index=0, format=AudioSampleFormat.s16)) - samples = bytes(samples) + # closed, copy_to() raises + buffer = bytearray(audio.number_of_frames * audio.number_of_channels * 2) + audio.copy_to(buffer, AudioDataCopyToOptions(plane_index=0, format=AudioSampleFormat.s16)) + samples = bytes(buffer) # rates beyond an int are unsupported too: the native check rejects them rate = min(int(audio.sample_rate), 2**31 - 1) try: @@ -125,7 +128,7 @@ class MediaStreamTrackGeneratorInit(Dictionary): kind (:obj:`webrtc.MediaType`): ``audio`` or ``video``. """ - kind: MediaType + kind: MediaType | MediaTypeValue class MediaStreamTrackGenerator(MediaStreamTrack): @@ -144,7 +147,7 @@ class MediaStreamTrackGenerator(MediaStreamTrack): TypeError: If the kind isn't audio or video. """ - def __init__(self, kind: str | MediaType | MediaStreamTrackGeneratorInit) -> None: + def __init__(self, kind: MediaType | MediaTypeValue | MediaStreamTrackGeneratorInit) -> None: if isinstance(kind, MediaStreamTrackGeneratorInit): kind = kind.kind if kind not in {'audio', 'video'}: diff --git a/python-webrtc/python/webrtc/models/audio_data.py b/python-webrtc/python/webrtc/models/audio_data.py index 74a0712..ece2e0b 100644 --- a/python-webrtc/python/webrtc/models/audio_data.py +++ b/python-webrtc/python/webrtc/models/audio_data.py @@ -12,13 +12,23 @@ import math import warnings from dataclasses import dataclass -from typing import Any, ClassVar, NamedTuple - -from webrtc import AudioSampleFormat, InvalidRangeError, InvalidStateError, NotSupportedError, wrtc +from typing import TYPE_CHECKING, ClassVar, NamedTuple, cast + +from webrtc import ( + AudioSampleFormat, + AudioSampleFormatValue, + InvalidRangeError, + InvalidStateError, + NotSupportedError, + wrtc, +) from webrtc.models.closable import Closable from webrtc.models.dictionary import Dictionary from webrtc.utils.names import Alias, alias +if TYPE_CHECKING: + from typing_extensions import Buffer + _SAMPLE_BYTES = {'u8': 1, 's16': 2, 's32': 4, 'f32': 4} @@ -51,12 +61,12 @@ class AudioDataInit(Dictionary): data: A bytes-like buffer of the samples, which is copied. """ - format: AudioSampleFormat + format: AudioSampleFormat | AudioSampleFormatValue sample_rate: float number_of_frames: int number_of_channels: int timestamp: int - data: Any + data: Buffer #: Alias for :attr:`sample_rate` sampleRate: ClassVar[Alias[float]] = alias('sample_rate') @@ -80,7 +90,7 @@ class AudioDataCopyToOptions(Dictionary): plane_index: int frame_offset: int = 0 frame_count: int | None = None - format: AudioSampleFormat | None = None + format: AudioSampleFormat | AudioSampleFormatValue | None = None #: Alias for :attr:`plane_index` planeIndex: ClassVar[Alias[int]] = alias('plane_index') @@ -107,7 +117,8 @@ def _sample_rate(value: object) -> float: def _buffer(data: object, size: int) -> bytes: """The first bytes of a bytes-like buffer.""" try: - view = memoryview(data).cast('B') + # memoryview() is the check of the buffer protocol + view = memoryview(cast('Buffer', data)).cast('B') except TypeError: msg = 'data must be a bytes-like buffer' raise TypeError(msg) from None @@ -125,6 +136,7 @@ class _Layout(NamedTuple): class _CopyPlan(NamedTuple): + data: bytes format: AudioSampleFormat plane_index: int frame_offset: int @@ -157,6 +169,15 @@ class AudioData(Closable): ) """ + _data: bytes | None + _format: AudioSampleFormat + _sample_rate: float + _frames: int + _channels: int + _timestamp: int + #: Whether a data read from a track warns if garbage collected without being closed + _warn_unclosed: bool = False + def __init__(self, init: AudioDataInit) -> None: format = _sample_format(init.format) sample_rate = _sample_rate(init.sample_rate) @@ -172,7 +193,7 @@ def __init__(self, init: AudioDataInit) -> None: 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._data = data self._format, self._sample_rate, self._frames, self._channels = layout self._timestamp = timestamp @@ -232,7 +253,8 @@ def timestamp(self) -> int: return self._timestamp def _plan_copy(self, options: AudioDataCopyToOptions) -> _CopyPlan: - if self._data is None: + data = self._data + if data is None: msg = 'The data is closed' raise InvalidStateError(msg) plane_index = _unsigned(options.plane_index, 'plane_index') @@ -256,7 +278,9 @@ def _plan_copy(self, options: AudioDataCopyToOptions) -> _CopyPlan: 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)) + return _CopyPlan( + data, destination, plane_index, frame_offset, frame_count, elements * _sample_bytes(destination) + ) def allocation_size(self, options: AudioDataCopyToOptions) -> int: """Returns how many bytes :meth:`copy_to` needs. @@ -288,7 +312,7 @@ def copy_to(self, destination: bytearray | memoryview, options: AudioDataCopyToO raise InvalidRangeError(msg) try: wrtc.copyAudioSamples( - self._data, + plan.data, self._format.value, self._channels, self._frames, diff --git a/python-webrtc/python/webrtc/models/blob.py b/python-webrtc/python/webrtc/models/blob.py index 7fe49e9..61bffd5 100644 --- a/python-webrtc/python/webrtc/models/blob.py +++ b/python-webrtc/python/webrtc/models/blob.py @@ -16,6 +16,7 @@ from typing import TYPE_CHECKING, TypeVar, Union if TYPE_CHECKING: + import builtins from collections.abc import Iterable BlobPart = Union[str, bytes, bytearray, memoryview, 'Blob'] @@ -41,8 +42,8 @@ class Blob: """ def __init__(self, parts: Iterable[BlobPart] | None = None, type: str = '') -> None: - chunks = [] - for part in parts or (): + chunks: list[bytes] = [] + for part in parts if parts is not None else (): if isinstance(part, Blob): chunks.append(part._bytes) elif isinstance(part, str): @@ -89,7 +90,8 @@ def text(self) -> asyncio.Future[str]: """Returns a future of the bytes decoded as UTF-8.""" return _done(self._bytes.decode('utf-8', 'replace')) - def __bytes__(self) -> bytes: + # the bytes() method hides the builtin in the class + def __bytes__(self) -> builtins.bytes: return self._bytes def __len__(self) -> int: diff --git a/python-webrtc/python/webrtc/models/dictionary.py b/python-webrtc/python/webrtc/models/dictionary.py index 2f23497..069c50d 100644 --- a/python-webrtc/python/webrtc/models/dictionary.py +++ b/python-webrtc/python/webrtc/models/dictionary.py @@ -11,7 +11,7 @@ from collections.abc import Mapping from dataclasses import fields -from typing import TYPE_CHECKING, Any, ClassVar +from typing import TYPE_CHECKING, ClassVar from webrtc.utils.names import members @@ -28,10 +28,10 @@ class Dictionary: _dictionaries: ClassVar[Mapping[str, type[Dictionary]]] = {} if TYPE_CHECKING: # every subclass is a dataclass - __dataclass_fields__: ClassVar[dict[str, Field[Any]]] + __dataclass_fields__: ClassVar[dict[str, Field[object]]] @classmethod - def from_json(cls, value: Mapping[str, Any]) -> Self: + def from_json(cls, value: Mapping[str, object]) -> Self: """Creates the dictionary from its JSON form, like a message from the remote peer. Keys are the camelCase names of the specification or the snake_case ones, unknown keys are ignored, and diff --git a/python-webrtc/python/webrtc/models/events.py b/python-webrtc/python/webrtc/models/events.py index 3ec3daf..dfd9256 100644 --- a/python-webrtc/python/webrtc/models/events.py +++ b/python-webrtc/python/webrtc/models/events.py @@ -86,12 +86,13 @@ class MessageEvent(Event): Args: type (:obj:`str`): The name of the event. - data (:obj:`str` or :obj:`bytes`): The message, :obj:`bytes` if it was sent as binary. + data (:obj:`str`, :obj:`bytes` or :obj:`webrtc.Blob`): The message, :obj:`bytes` (or a :obj:`webrtc.Blob` + with the ``blob`` binary type) if it was sent as binary. target (:obj:`object`, optional): The object that emitted the event. """ type: str - data: str | bytes + data: str | bytes | webrtc.Blob target: webrtc.EventTarget | None = None diff --git a/python-webrtc/python/webrtc/models/rtc_certificate.py b/python-webrtc/python/webrtc/models/rtc_certificate.py index bf55f79..44253ae 100644 --- a/python-webrtc/python/webrtc/models/rtc_certificate.py +++ b/python-webrtc/python/webrtc/models/rtc_certificate.py @@ -122,7 +122,7 @@ def _key_params(algorithm: AlgorithmIdentifier) -> _KeyParams: return key_params(algorithm) -class RTCCertificate(WebRTCObject): +class RTCCertificate(WebRTCObject[wrtc.RTCCertificate]): """A certificate a connection uses to authenticate with DTLS. Generated with :meth:`generate` and set with :attr:`webrtc.RTCConfiguration.certificates`. Without one, @@ -155,7 +155,7 @@ async def generate(cls, algorithm: AlgorithmIdentifier = 'ECDSA', expires: float raise ValueError(msg) native = await asyncio.get_running_loop().run_in_executor( None, - cls._class.generate, + wrtc.RTCCertificate.generate, key_type, modulus_length, exponent, diff --git a/python-webrtc/python/webrtc/models/rtc_configuration.py b/python-webrtc/python/webrtc/models/rtc_configuration.py index bf05f40..3588c53 100644 --- a/python-webrtc/python/webrtc/models/rtc_configuration.py +++ b/python-webrtc/python/webrtc/models/rtc_configuration.py @@ -31,6 +31,13 @@ if TYPE_CHECKING: from collections.abc import Iterable + from webrtc.enums import ( + RTCBundlePolicyValue, + RTCIceTransportPolicyValue, + RTCRtcpMuxPolicyValue, + RTCRtpHeaderEncryptionPolicyValue, + ) + # the longest TURN username, as browsers limit it _MAX_USERNAME_LENGTH = 509 _MAX_PORT = 65535 @@ -51,22 +58,24 @@ def _check_url(url: str) -> str: 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'}: + if match is None 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'] + host: str = match['host'] + port: str | None = match['port'] + transport: str | None = match['transport'] if host.startswith('['): try: - ipaddress.IPv6Address(host[1:-1]) + _ = ipaddress.IPv6Address(host[1:-1]) except ValueError: msg = f'{url!r} has an invalid IPv6 address' raise InvalidSyntaxError(msg) from None - elif not _REG_NAME.fullmatch(host): + elif _REG_NAME.fullmatch(host) is None: msg = f'{url!r} has an invalid host' raise InvalidSyntaxError(msg) - if port is not None and (not port or int(port) > _MAX_PORT): + if port is not None and (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'}): @@ -121,7 +130,7 @@ def _to_native_list(cls, servers: Iterable[RTCIceServer]) -> list[wrtc.IceServer def _to_native(self) -> wrtc.IceServerInit: urls = [self.urls] if isinstance(self.urls, str) else list(self.urls) - if not urls: + if len(urls) == 0: msg = 'urls of an ICE server must not be empty' raise InvalidSyntaxError(msg) @@ -137,7 +146,7 @@ def _to_native(self) -> wrtc.IceServerInit: 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: + if self.username is None or self.credential is None or self.credential == '': msg = 'a TURN server needs a username and a credential' raise InvalidAccessError(msg) if len(self.username) > _MAX_USERNAME_LENGTH: @@ -147,9 +156,15 @@ def _to_native(self) -> wrtc.IceServerInit: native = wrtc.IceServerInit() native.urls = urls native.username = self.username - native.credential = self.credential + native.credential = self._password() return native + def _password(self) -> str | None: + if isinstance(self.credential, RTCOAuthCredential): + msg = 'the credential of a server other than an OAuth TURN one is a str' + raise TypeError(msg) + return self.credential + #: Alias for :attr:`credential_type` credentialType: ClassVar[Alias[str]] = alias('credential_type') @@ -181,14 +196,16 @@ class RTCConfiguration(Dictionary): """ ice_servers: list[RTCIceServer] = field(default_factory=list) - ice_transport_policy: RTCIceTransportPolicy = RTCIceTransportPolicy.all - bundle_policy: RTCBundlePolicy = RTCBundlePolicy.balanced - rtcp_mux_policy: RTCRtcpMuxPolicy = RTCRtcpMuxPolicy.require + ice_transport_policy: RTCIceTransportPolicy | RTCIceTransportPolicyValue = RTCIceTransportPolicy.all + bundle_policy: RTCBundlePolicy | RTCBundlePolicyValue = RTCBundlePolicy.balanced + rtcp_mux_policy: RTCRtcpMuxPolicy | RTCRtcpMuxPolicyValue = RTCRtcpMuxPolicy.require ice_candidate_pool_size: int = 0 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 + rtp_header_encryption_policy: RTCRtpHeaderEncryptionPolicy | RTCRtpHeaderEncryptionPolicyValue = ( + RTCRtpHeaderEncryptionPolicy.negotiate + ) _dictionaries: ClassVar = {'ice_servers': RTCIceServer} @@ -246,8 +263,8 @@ def _from_native(cls, native: wrtc.ConfigurationInit) -> RTCConfiguration: bundle_policy=native.bundlePolicy, rtcp_mux_policy=native.rtcpMuxPolicy, ice_candidate_pool_size=native.iceCandidatePoolSize, - port_range=tuple(native.portRange) if native.portRange else None, - certificates=RTCCertificate._wrap_many(native.certificates) if native.certificates else [], + port_range=tuple(native.portRange) if native.portRange is not None else None, + certificates=RTCCertificate._wrap_many(native.certificates) if native.certificates is not None else [], always_negotiate_data_channels=native.alwaysNegotiateDataChannels, rtp_header_encryption_policy=native.rtpHeaderEncryptionPolicy, ) @@ -255,11 +272,13 @@ def _from_native(cls, native: wrtc.ConfigurationInit) -> RTCConfiguration: #: Alias for :attr:`ice_servers` iceServers: ClassVar[Alias[list[RTCIceServer]]] = alias('ice_servers') #: Alias for :attr:`ice_transport_policy` - iceTransportPolicy: ClassVar[Alias[RTCIceTransportPolicy]] = alias('ice_transport_policy') + iceTransportPolicy: ClassVar[Alias[RTCIceTransportPolicy | RTCIceTransportPolicyValue]] = alias( + 'ice_transport_policy' + ) #: Alias for :attr:`bundle_policy` - bundlePolicy: ClassVar[Alias[RTCBundlePolicy]] = alias('bundle_policy') + bundlePolicy: ClassVar[Alias[RTCBundlePolicy | RTCBundlePolicyValue]] = alias('bundle_policy') #: Alias for :attr:`rtcp_mux_policy` - rtcpMuxPolicy: ClassVar[Alias[RTCRtcpMuxPolicy]] = alias('rtcp_mux_policy') + rtcpMuxPolicy: ClassVar[Alias[RTCRtcpMuxPolicy | RTCRtcpMuxPolicyValue]] = alias('rtcp_mux_policy') #: Alias for :attr:`ice_candidate_pool_size` iceCandidatePoolSize: ClassVar[Alias[int]] = alias('ice_candidate_pool_size') #: Alias for :attr:`port_range` @@ -267,4 +286,6 @@ def _from_native(cls, native: wrtc.ConfigurationInit) -> RTCConfiguration: #: Alias for :attr:`always_negotiate_data_channels` alwaysNegotiateDataChannels: ClassVar[Alias[bool]] = alias('always_negotiate_data_channels') #: Alias for :attr:`rtp_header_encryption_policy` - rtpHeaderEncryptionPolicy: ClassVar[Alias[RTCRtpHeaderEncryptionPolicy]] = alias('rtp_header_encryption_policy') + rtpHeaderEncryptionPolicy: ClassVar[Alias[RTCRtpHeaderEncryptionPolicy | RTCRtpHeaderEncryptionPolicyValue]] = ( + 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 cc60617..c4d5740 100644 --- a/python-webrtc/python/webrtc/models/rtc_ice_candidate.py +++ b/python-webrtc/python/webrtc/models/rtc_ice_candidate.py @@ -10,9 +10,11 @@ from __future__ import annotations import re -from dataclasses import dataclass +from dataclasses import dataclass, field from enum import Enum -from typing import TYPE_CHECKING, Any, ClassVar, TypeVar, Union +from typing import TYPE_CHECKING, ClassVar, TypeVar + +from typing_extensions import TypedDict from webrtc import ( RTCIceCandidateType, @@ -27,6 +29,9 @@ if TYPE_CHECKING: from collections.abc import Mapping + import wrtc + from webrtc.enums import RTCIceServerTransportProtocolValue + _FOUNDATION = re.compile(r'[A-Za-z0-9+/]{1,32}') _DIGITS = re.compile(r'[0-9]+') _TOKEN = re.compile(r"[!#$%&'*+\-.^_`{|}~A-Za-z0-9]+") @@ -42,8 +47,21 @@ _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 _CandidateFields(TypedDict, total=False): + """The fields parsed from a candidate-attribute, by the names of the properties of RTCIceCandidate.""" + + foundation: str + component: str | None + priority: int + address: str | None + protocol: str + port: int + type: str + tcp_type: str | None + related_address: str | None + related_port: int | None class _InvalidCandidateError(ValueError): @@ -51,7 +69,7 @@ class _InvalidCandidateError(ValueError): def _number(token: str, max_digits: int, valid: range) -> int | None: - if len(token) > max_digits or not _DIGITS.fullmatch(token): + if len(token) > max_digits or _DIGITS.fullmatch(token) is None: return None value = int(token) return value if value in valid else None @@ -77,7 +95,7 @@ def _parse_related(fields: _CandidateFields, rest: list[str], *, strict: bool) - fields['related_address'] = rest[1] fields['related_port'] = _required(_number(rest[3], 5, _PORTS)) return rest[_RELATED_TOKENS:] - if fields['type'] != 'host' and strict: + if fields.get('type') != 'host' and strict: raise _InvalidCandidateError return rest @@ -89,7 +107,7 @@ def _parse_tcp_type(fields: _CandidateFields, rest: list[str]) -> list[str]: 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': + if fields.get('protocol') == 'tcp' and fields.get('type') != 'relay': raise _InvalidCandidateError return rest @@ -99,7 +117,7 @@ def _base_fields(tokens: list[str]) -> _CandidateFields: 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), + 'foundation': _required(foundation if _FOUNDATION.fullmatch(foundation) is not None else None), 'component': _COMPONENTS.get(component_id), 'priority': _required(_number(priority, 10, range(1, 2**31))), 'address': address, @@ -122,7 +140,7 @@ def _parse_fields(value: str, *, strict: bool) -> _CandidateFields: 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]): + if len(rest) % 2 != 0 or not all(_TOKEN.fullmatch(t) for t in rest[::2]): raise _InvalidCandidateError return fields @@ -232,8 +250,11 @@ class RTCIceCandidate: sdp_mid: str | None = None sdp_m_line_index: int | None = None username_fragment: str | None = None - relay_protocol: RTCIceServerTransportProtocol | None = None + relay_protocol: RTCIceServerTransportProtocol | RTCIceServerTransportProtocolValue | None = None url: str | None = None + if TYPE_CHECKING: + # parsed from the candidate by __post_init__ + _parsed: _CandidateFields = field(init=False, repr=False, compare=False) def __post_init__(self) -> None: if self.sdp_mid is None and self.sdp_m_line_index is None: @@ -242,17 +263,19 @@ def __post_init__(self) -> None: 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 - object.__setattr__(self, '_parsed', _parse_candidate(self.candidate) or {}) + parsed = _parse_candidate(self.candidate) + object.__setattr__(self, '_parsed', parsed if parsed is not None else _CandidateFields()) @staticmethod def _members_of( candidate: RTCIceCandidate | RTCIceCandidateInit, ) -> 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 init.""" - return candidate.candidate or '', candidate.sdp_mid, candidate.sdp_m_line_index, candidate.username_fragment + text = candidate.candidate if candidate.candidate is not None else '' + return text, candidate.sdp_mid, candidate.sdp_m_line_index, candidate.username_fragment @classmethod - def _peer_reflexive(cls, kwargs: dict[str, Any]) -> RTCIceCandidate: + def _peer_reflexive(cls, kwargs: wrtc._IceCandidateKwargs) -> 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. @@ -263,16 +286,23 @@ def _peer_reflexive(cls, kwargs: dict[str, Any]) -> RTCIceCandidate: Returns: :obj:`webrtc.RTCIceCandidate`: The candidate. """ - fields = _parse_candidate(kwargs.get('candidate', ''), strict=False) or {} - candidate = cls(**{**kwargs, 'candidate': ''}) + fields = _parse_candidate(kwargs['candidate'], strict=False) + init = kwargs.copy() + init['candidate'] = '' + candidate = cls(**init) # libwebrtc knows no related address of it, which is port 0 - parsed = {**fields, 'address': None, 'related_address': None, 'related_port': 0} + parsed: _CandidateFields = { + **(fields if fields is not None else _CandidateFields()), + 'address': None, + 'related_address': None, + 'related_port': 0, + } # past the frozen __setattr__, as __post_init__ parsed the empty candidate string vars(candidate)['_parsed'] = parsed return candidate @classmethod - def from_json(cls, value: Mapping[str, Any]) -> RTCIceCandidate: + def from_json(cls, value: Mapping[str, object]) -> RTCIceCandidate: """Creates a candidate from its JSON form, as :meth:`to_json` returns it. Args: @@ -336,7 +366,7 @@ 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, str | int | None]: """The candidate as a JSON-serializable dictionary, to send to the remote peer. Returns: diff --git a/python-webrtc/python/webrtc/models/rtc_rtp_transceiver_init.py b/python-webrtc/python/webrtc/models/rtc_rtp_transceiver_init.py index e3bce1e..611e696 100644 --- a/python-webrtc/python/webrtc/models/rtc_rtp_transceiver_init.py +++ b/python-webrtc/python/webrtc/models/rtc_rtp_transceiver_init.py @@ -13,7 +13,7 @@ from typing import TYPE_CHECKING, ClassVar from webrtc import wrtc -from webrtc.enums import TransceiverDirection +from webrtc.enums import TransceiverDirection, TransceiverDirectionValue from webrtc.models.dictionary import Dictionary from webrtc.models.rtp_parameters import RTCRtpEncodingParameters from webrtc.utils.names import Alias, alias @@ -38,7 +38,7 @@ class RTCRtpTransceiverInit(Dictionary): ValueError: If the direction isn't a member of :obj:`webrtc.TransceiverDirection`. """ - direction: TransceiverDirection = TransceiverDirection.sendrecv + direction: TransceiverDirection | TransceiverDirectionValue = TransceiverDirection.sendrecv streams: list[webrtc.MediaStream] = field(default_factory=list) send_encodings: list[RTCRtpEncodingParameters] = field(default_factory=list) diff --git a/python-webrtc/python/webrtc/models/rtc_session_description.py b/python-webrtc/python/webrtc/models/rtc_session_description.py index 6187836..1c90a1c 100644 --- a/python-webrtc/python/webrtc/models/rtc_session_description.py +++ b/python-webrtc/python/webrtc/models/rtc_session_description.py @@ -17,7 +17,7 @@ import webrtc -class RTCSessionDescription(WebRTCObject): +class RTCSessionDescription(WebRTCObject[wrtc.RTCSessionDescription]): """One end of a connection or potential connection and how it's configured. Each :obj:`webrtc.RTCSessionDescription` consists of @@ -51,11 +51,11 @@ class RTCSessionDescription(WebRTCObject): def __init__( self, - type: webrtc.RTCSdpType | webrtc.RTCSessionDescriptionInit, + type: webrtc.RTCSdpType | webrtc.RTCSdpTypeValue | webrtc.RTCSessionDescriptionInit, sdp: str = '', ) -> None: init = type if isinstance(type, RTCSessionDescriptionInit) else RTCSessionDescriptionInit(type, sdp) - super().__init__(self._class(init._to_native())) + super().__init__(wrtc.RTCSessionDescription(init._to_native())) def to_json(self) -> dict[str, str]: """The description as a JSON-serializable dictionary, to send to the remote peer. 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 78f625a..c371471 100644 --- a/python-webrtc/python/webrtc/models/rtc_session_description_init.py +++ b/python-webrtc/python/webrtc/models/rtc_session_description_init.py @@ -13,7 +13,7 @@ from typing import ClassVar from webrtc import wrtc -from webrtc.enums import RTCSdpType +from webrtc.enums import RTCSdpType, RTCSdpTypeValue from webrtc.models.dictionary import Dictionary @@ -32,7 +32,7 @@ class RTCSessionDescriptionInit(Dictionary): TypeError: If the SDP is :obj:`None`. """ - type: RTCSdpType + type: RTCSdpType | RTCSdpTypeValue sdp: str = '' def __post_init__(self) -> None: @@ -50,10 +50,10 @@ def to_json(self) -> dict[str, str]: Returns: :obj:`dict`: ``type`` (like ``'offer'``) and ``sdp``. """ - return {'type': self.type.value, 'sdp': self.sdp} + return {'type': RTCSdpType(self.type).value, 'sdp': self.sdp} def __repr__(self) -> str: - return f'RTCSessionDescriptionInit(type={self.type.value!r}, sdp={len(self.sdp)} characters)' + return f'RTCSessionDescriptionInit(type={RTCSdpType(self.type).value!r}, sdp={len(self.sdp)} characters)' #: Alias for :attr:`to_json` toJSON: ClassVar = to_json @@ -74,7 +74,7 @@ class RTCLocalSessionDescriptionInit(Dictionary): TypeError: If the SDP is :obj:`None`. """ - type: RTCSdpType | None = None + type: RTCSdpType | RTCSdpTypeValue | None = None sdp: str = '' def __post_init__(self) -> None: diff --git a/python-webrtc/python/webrtc/models/rtc_stats.py b/python-webrtc/python/webrtc/models/rtc_stats.py index 967a0d6..12004e9 100644 --- a/python-webrtc/python/webrtc/models/rtc_stats.py +++ b/python-webrtc/python/webrtc/models/rtc_stats.py @@ -23,6 +23,10 @@ StatsValue = Union[str, int, float, bool, None, list['StatsValue'], dict[str, 'StatsValue']] +def _is_empty(value: StatsValue) -> bool: + return value is None or value == '' + + class RTCStats(dict[str, StatsValue]): """Stats of one object, like an outbound RTP stream. @@ -44,20 +48,31 @@ def __getattr__(self, name: str) -> StatsValue: msg = f'{type(self).__name__} of type {self.get("type")!r} has no {name!r}' raise AttributeError(msg) from None + def _string(self, key: str) -> str: + value = self[key] + if not isinstance(value, str): + msg = f'{key} of stats is a {type(value).__name__}, not a str' + raise TypeError(msg) + return value + @property def id(self) -> str: """:obj:`str`: Identifies the stats in its report.""" - return self['id'] + return self._string('id') @property def type(self) -> str: """:obj:`str`: The type of the stats, like ``'outbound-rtp'``.""" - return self['type'] + return self._string('type') @property def timestamp(self) -> float: """:obj:`float`: When the stats were collected, in milliseconds since the epoch.""" - return self['timestamp'] + value = self['timestamp'] + if isinstance(value, bool) or not isinstance(value, (int, float)): + msg = f'timestamp of stats is a {type(value).__name__}, not a number' + raise TypeError(msg) + return value class RTCStatsReport(Mapping[str, RTCStats]): @@ -74,20 +89,21 @@ def _from_native(cls, report: str, receivers: Iterable[webrtc.RTCRtpReceiver] = """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 '[]')] + entries: list[dict[str, StatsValue]] = json.loads(report if report != '' else '[]') + stats = [RTCStats(entry) for entry in entries] for entry in stats: # libwebrtc serializes microseconds - entry['timestamp'] /= 1000 + entry['timestamp'] = 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 'address' in entry - and not entry['address'] + and _is_empty(entry['address']) ): entry['address'] = None - return cls({entry['id']: entry for entry in stats}) + return cls({entry.id: entry for entry in stats}) def __getitem__(self, stats_id: str) -> RTCStats: return self._stats[stats_id] diff --git a/python-webrtc/python/webrtc/models/rtp_parameters.py b/python-webrtc/python/webrtc/models/rtp_parameters.py index 4ca1b9d..f13dc58 100644 --- a/python-webrtc/python/webrtc/models/rtp_parameters.py +++ b/python-webrtc/python/webrtc/models/rtp_parameters.py @@ -12,12 +12,22 @@ import dataclasses import math from dataclasses import dataclass, field -from typing import Any, ClassVar, TypeVar +from typing import TYPE_CHECKING, ClassVar, TypeVar + +from typing_extensions import TypedDict from webrtc import MediaType, RTCDegradationPreference, RTCPriorityType, TransceiverDirection, wrtc from webrtc.models.dictionary import Dictionary from webrtc.utils.names import Alias, alias +if TYPE_CHECKING: + from webrtc.enums import ( + MediaTypeValue, + RTCDegradationPreferenceValue, + RTCPriorityTypeValue, + TransceiverDirectionValue, + ) + _NativeCodecT = TypeVar('_NativeCodecT', bound='wrtc.RtpCodec') # the bitrate priorities of libwebrtc for RTCRtpEncodingParameters.priority, as Chromium maps them @@ -35,10 +45,10 @@ def _is_unsigned_long(value: object) -> bool: def _parse_fmtp(line: str | None) -> dict[str, str]: - parameters = {} - for raw_item in (line or '').split(';'): + parameters: dict[str, str] = {} + for raw_item in (line if line is not None else '').split(';'): item = raw_item.strip() - if not item: + if item == '': continue if '=' in item: key, _, value = item.partition('=') @@ -50,16 +60,32 @@ def _parse_fmtp(line: str | None) -> dict[str, str]: def _format_fmtp(parameters: dict[str, str]) -> str | None: - if not parameters: + if len(parameters) == 0: return None - return ';'.join(f'{key}={value}' if key else value for key, value in parameters.items()) + return ';'.join(f'{key}={value}' if key != '' else value for key, value in parameters.items()) + + +class _CodecMembers(TypedDict, closed=True): + mime_type: str + clock_rate: int + channels: int | None + sdp_fmtp_line: str | None + + +def _channels_or_one(channels: int | None) -> int: + return channels if channels is not None and channels != 0 else 1 -def _codec_members(native: wrtc.RtpCodec) -> dict[str, Any]: +def _codec_members(native: wrtc.RtpCodec) -> _CodecMembers: """The members RTCRtpCodec and RTCRtpCodecParameters share.""" + clock_rate = native.clockRate + # libwebrtc sets it for every codec it reports, while the specification requires it + if clock_rate is None: + msg = f'{native.mimeType} has no clock rate' + raise ValueError(msg) return { 'mime_type': native.mimeType, - 'clock_rate': native.clockRate, + 'clock_rate': clock_rate, 'channels': native.numChannels, 'sdp_fmtp_line': _format_fmtp(native.parameters), } @@ -89,7 +115,7 @@ def _from_native(cls, native: wrtc.RtpCodec) -> RTCRtpCodec: 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: + if name == '': msg = f'{self.mime_type!r} is not a valid MIME type of a codec' raise ValueError(msg) native = native_class() @@ -105,7 +131,7 @@ def _matches(self, other: RTCRtpCodec) -> bool: return ( self.mime_type.lower() == other.mime_type.lower() and self.clock_rate == other.clock_rate - and (self.channels or 1) == (other.channels or 1) + and _channels_or_one(self.channels) == _channels_or_one(other.channels) and _parse_fmtp(self.sdp_fmtp_line) == _parse_fmtp(other.sdp_fmtp_line) ) @@ -206,8 +232,8 @@ class RTCRtpEncodingParameters(Dictionary): 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 + priority: RTCPriorityType | RTCPriorityTypeValue = RTCPriorityType.low + network_priority: RTCPriorityType | RTCPriorityTypeValue = RTCPriorityType.low scalability_mode: str | None = None adaptive_ptime: bool = False codec: RTCRtpCodec | None = None @@ -222,7 +248,7 @@ def _from_native(cls, native: wrtc.RtpEncodingParameters) -> RTCRtpEncodingParam max_bitrate=native.maxBitrate, max_framerate=native.maxFramerate, scale_resolution_down_by=native.scaleResolutionDownBy, - rid=native.rid or None, + rid=native.rid if native.rid != '' else None, priority=priority, network_priority=native.networkPriority, scalability_mode=native.scalabilityMode, @@ -265,7 +291,7 @@ def _apply(self, native: wrtc.RtpEncodingParameters) -> wrtc.RtpEncodingParamete def _to_native(self) -> wrtc.RtpEncodingParameters: native = self._apply(wrtc.RtpEncodingParameters()) - native.rid = self.rid or '' + native.rid = self.rid if self.rid is not None else '' return native #: Alias for :attr:`max_bitrate` @@ -335,7 +361,7 @@ class RTCRtpSendParameters(Dictionary): codecs: list[RTCRtpCodecParameters] = field(default_factory=list) header_extensions: list[RTCRtpHeaderExtensionParameters] = field(default_factory=list) rtcp: RTCRtcpParameters = field(default_factory=RTCRtcpParameters) - degradation_preference: RTCDegradationPreference | None = None + degradation_preference: RTCDegradationPreference | RTCDegradationPreferenceValue | None = None _dictionaries: ClassVar = { 'encodings': RTCRtpEncodingParameters, @@ -374,7 +400,7 @@ class RTCRtpHeaderExtensionCapability(Dictionary): """ uri: str - direction: TransceiverDirection = TransceiverDirection.sendrecv + direction: TransceiverDirection | TransceiverDirectionValue = TransceiverDirection.sendrecv @classmethod def _from_native(cls, native: wrtc.RtpHeaderExtensionCapability) -> RTCRtpHeaderExtensionCapability: @@ -405,7 +431,7 @@ def _from_native(cls, native: wrtc.RtpCapabilities) -> RTCRtpCapabilities: @classmethod def _supported( - cls, native_class: type[wrtc.RTCRtpSender | wrtc.RTCRtpReceiver], kind: MediaType + cls, native_class: type[wrtc.RTCRtpSender | wrtc.RTCRtpReceiver], kind: MediaType | MediaTypeValue ) -> RTCRtpCapabilities | None: """The capabilities of ``wrtc.RTCRtpSender`` or ``wrtc.RTCRtpReceiver`` for a kind. diff --git a/python-webrtc/python/webrtc/models/rtp_source.py b/python-webrtc/python/webrtc/models/rtp_source.py index dafdd89..23d852f 100644 --- a/python-webrtc/python/webrtc/models/rtp_source.py +++ b/python-webrtc/python/webrtc/models/rtp_source.py @@ -9,6 +9,11 @@ from __future__ import annotations +from typing import TYPE_CHECKING + +if TYPE_CHECKING: + from typing_extensions import Self + from dataclasses import dataclass from typing import ClassVar @@ -38,7 +43,7 @@ class RTCRtpContributingSource(Dictionary): audio_level: float | None = None @classmethod - def _from_native(cls, native: tuple[bool, int, float, int, int | None]) -> RTCRtpContributingSource: + def _from_native(cls, native: tuple[bool, int, float, int, int | None]) -> Self: """A source from the native one: whether it's an SSRC, the source, timestamp, RTP timestamp and level.""" _, source, timestamp, rtp_timestamp, level = native if level is not None: diff --git a/python-webrtc/python/webrtc/models/video_frame.py b/python-webrtc/python/webrtc/models/video_frame.py index dd0b294..8e59a02 100644 --- a/python-webrtc/python/webrtc/models/video_frame.py +++ b/python-webrtc/python/webrtc/models/video_frame.py @@ -14,17 +14,22 @@ import warnings from dataclasses import dataclass from enum import Enum -from typing import TYPE_CHECKING, Any, ClassVar, NamedTuple, TypeVar +from typing import TYPE_CHECKING, ClassVar, NamedTuple, TypeVar, cast, overload from webrtc import ( AlphaOption, + AlphaOptionValue, InvalidStateError, NotSupportedError, RTCException, VideoColorPrimaries, + VideoColorPrimariesValue, VideoMatrixCoefficients, + VideoMatrixCoefficientsValue, VideoPixelFormat, + VideoPixelFormatValue, VideoTransferCharacteristics, + VideoTransferCharacteristicsValue, wrtc, ) from webrtc.models.closable import Closable @@ -32,10 +37,9 @@ from webrtc.utils.names import Alias, alias if TYPE_CHECKING: - from typing_extensions import Buffer + from typing_extensions import Buffer, TypeGuard _EnumT = TypeVar('_EnumT', bound=Enum) -_InitT = TypeVar('_InitT', 'VideoFrameBufferInit', 'VideoFrameInit') _MAX_UNSIGNED_LONG = 2**32 - 1 _RGB_FORMATS = (VideoPixelFormat.RGBA, VideoPixelFormat.RGBX, VideoPixelFormat.BGRA, VideoPixelFormat.BGRX) @@ -119,12 +123,12 @@ class VideoColorSpace: full_range (:obj:`bool`, optional): Whether the samples use the full range of their bits. """ - primaries: VideoColorPrimaries | None = None - transfer: VideoTransferCharacteristics | None = None - matrix: VideoMatrixCoefficients | None = None + primaries: VideoColorPrimaries | VideoColorPrimariesValue | None = None + transfer: VideoTransferCharacteristics | VideoTransferCharacteristicsValue | None = None + matrix: VideoMatrixCoefficients | VideoMatrixCoefficientsValue | None = None full_range: bool | None = None - def to_json(self) -> dict[str, Any]: + def to_json(self) -> dict[str, str | bool | None]: """Returns the members as a dictionary with the camelCase names, like ``toJSON()``.""" return { 'primaries': self.primaries, @@ -150,9 +154,9 @@ class VideoColorSpaceInit(Dictionary): full_range (:obj:`bool`, optional): Whether the samples use the full range of their bits. """ - primaries: VideoColorPrimaries | None = None - transfer: VideoTransferCharacteristics | None = None - matrix: VideoMatrixCoefficients | None = None + primaries: VideoColorPrimaries | VideoColorPrimariesValue | None = None + transfer: VideoTransferCharacteristics | VideoTransferCharacteristicsValue | None = None + matrix: VideoMatrixCoefficients | VideoMatrixCoefficientsValue | None = None full_range: bool | None = None #: Alias for :attr:`full_range` @@ -210,7 +214,7 @@ class VideoFrameBufferInit(Dictionary): _dictionaries: ClassVar = {'layout': PlaneLayout, 'visible_rect': DOMRectInit, 'color_space': VideoColorSpaceInit} - format: VideoPixelFormat + format: VideoPixelFormat | VideoPixelFormatValue coded_width: int coded_height: int timestamp: int @@ -256,7 +260,7 @@ class VideoFrameInit(Dictionary): timestamp: int | None = None duration: int | None = None - alpha: AlphaOption = AlphaOption.keep + alpha: AlphaOption | AlphaOptionValue = AlphaOption.keep visible_rect: DOMRectInit | DOMRectReadOnly | None = None rotation: float = 0 flip: bool = False @@ -286,7 +290,7 @@ class VideoFrameCopyToOptions(Dictionary): rect: DOMRectInit | DOMRectReadOnly | None = None layout: list[PlaneLayout] | None = None - format: VideoPixelFormat | None = None + format: VideoPixelFormat | VideoPixelFormatValue | None = None class _Plane(NamedTuple): @@ -342,9 +346,10 @@ def _optional_enum(cls: type[_EnumT], value: object) -> _EnumT | None: return None if value is None else _enum(cls, value) -def _is_buffer(value: object) -> bool: +def _is_buffer(value: object) -> TypeGuard[Buffer]: try: - memoryview(value) + # the buffer protocol has no runtime check of its own before Python 3.12: memoryview() is the check + _ = memoryview(cast('Buffer', value)) except TypeError: return False return True @@ -401,7 +406,7 @@ def _rotation(value: float) -> int: if not math.isfinite(value): msg = f'The rotation must be finite, not {value!r}' raise TypeError(msg) - return int(math.floor(value / 90 + 0.5) * 90) % 360 + return math.floor(value / 90 + 0.5) * 90 % 360 def _checked_rect(rect: DOMRectReadOnly, coded_size: tuple[int, int]) -> DOMRectReadOnly: @@ -428,7 +433,7 @@ def _parse_visible_rect( ) -> DOMRectReadOnly: 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: + if rect.x % plane.subsampling_x != 0 or rect.y % plane.subsampling_y != 0: msg = f'The rect must be aligned to the subsampling of {format.value}' raise TypeError(msg) return rect @@ -449,6 +454,7 @@ def end(self) -> int: class _CopyPlan(NamedTuple): + resource: wrtc.VideoFrameBuffer format: VideoPixelFormat rect: DOMRectReadOnly size: int @@ -504,16 +510,6 @@ def _color_space(value: VideoColorSpaceInit | VideoColorSpace | None) -> VideoCo ) -def _init_of(init: _InitT | None, cls: type[_InitT]) -> _InitT: - """The init of a constructor, which a frame of a buffer requires.""" - if init is not None: - return init - if cls is VideoFrameInit: - return cls() - msg = f'A VideoFrame of a buffer needs 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: @@ -578,6 +574,23 @@ class VideoFrame(Closable): ) """ + _resource: wrtc.VideoFrameBuffer | None + _format: VideoPixelFormat + _visible_rect: DOMRectReadOnly + _display: tuple[int, int] + _rotation: int + _flip: bool + _timestamp: int + _duration: int | None + _color_space: VideoColorSpace + _metadata: VideoFrameMetadata + + @overload + def __init__(self, source: VideoFrame, init: VideoFrameInit | None = None) -> None: ... + + @overload + def __init__(self, source: Buffer, init: VideoFrameBufferInit) -> None: ... + def __init__( self, source: Buffer | VideoFrame, @@ -585,9 +598,15 @@ def __init__( ) -> None: self._resource = None if isinstance(source, VideoFrame): - self._init_from_frame(source, _init_of(init, VideoFrameInit)) + if isinstance(init, VideoFrameBufferInit): + msg = 'A VideoFrame of a VideoFrame takes a VideoFrameInit' + raise TypeError(msg) + self._init_from_frame(source, init if init is not None else VideoFrameInit()) elif _is_buffer(source): - self._init_from_buffer(source, _init_of(init, VideoFrameBufferInit)) + if not isinstance(init, VideoFrameBufferInit): + msg = 'A VideoFrame of a buffer needs a VideoFrameBufferInit' + raise TypeError(msg) + self._init_from_buffer(source, init) else: msg = f'A VideoFrame is created from a buffer or a VideoFrame, not {type(source).__name__}' raise TypeError(msg) @@ -602,13 +621,15 @@ def _init_from_buffer(self, data: Buffer, init: VideoFrameBufferInit) -> None: 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) + color_space = _color_space(init.color_space) + if color_space is None: + color_space = _SRGB if format in _RGB_FORMATS else _REC709 self._set( resource, format, geometry=_Geometry( DOMRectReadOnly(0, 0, width, height), - display or _oriented(width, height, rotation), + display if display is not None else _oriented(width, height, rotation), rotation, flip=bool(init.flip), ), @@ -678,7 +699,9 @@ def _from_native(cls, native: tuple[wrtc.VideoFrameBuffer, int, int, int]) -> Vi geometry=_Geometry( DOMRectReadOnly(0, 0, width, height), _oriented(width, height, rotation), rotation, flip=False ), - info=_FrameInfo(timestamp, None, _REC601, VideoFrameMetadata(rtp_timestamp or None)), + info=_FrameInfo( + timestamp, None, _REC601, VideoFrameMetadata(rtp_timestamp if rtp_timestamp != 0 else None) + ), ) return frame @@ -769,7 +792,8 @@ def metadata(self) -> VideoFrameMetadata: return VideoFrameMetadata(self._metadata.rtp_timestamp) def _plan_copy(self, options: VideoFrameCopyToOptions | None) -> _CopyPlan: - if self._resource is None: + resource = self._resource + if resource is None: msg = 'The frame is closed' raise InvalidStateError(msg) if options is None: @@ -783,7 +807,7 @@ def _plan_copy(self, options: VideoFrameCopyToOptions | None) -> _CopyPlan: 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) + return _CopyPlan(resource, format, rect, size, planes) def allocation_size(self, options: VideoFrameCopyToOptions | None = None) -> int: """Returns how many bytes :meth:`copy_to` needs. @@ -824,11 +848,11 @@ def _copy_to( 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]) + plan.resource.copyPlanes(destination, [tuple(plane) for plane in plan.planes]) else: (plane,) = plan.planes rect, matrix = plan.rect, self._color_space.matrix - self._resource.convertTo( + plan.resource.convertTo( destination, plan.format.value, int(rect.x), diff --git a/python-webrtc/python/webrtc/streams.py b/python-webrtc/python/webrtc/streams.py index 26410d9..0aafb79 100644 --- a/python-webrtc/python/webrtc/streams.py +++ b/python-webrtc/python/webrtc/streams.py @@ -17,7 +17,9 @@ import collections import inspect from dataclasses import dataclass -from typing import TYPE_CHECKING, Any, Callable, NamedTuple, Protocol +from typing import TYPE_CHECKING, Callable, Generic, NamedTuple, Protocol, cast + +from typing_extensions import TypeVar if TYPE_CHECKING: from collections.abc import AsyncIterator @@ -35,6 +37,13 @@ ] +#: The type of the chunks of a stream, any object by default +_T = TypeVar('_T', default=object) +#: The type of the chunks a transform stream outputs +_O = TypeVar('_O', default=object) +_R = TypeVar('_R') + + def _loop() -> asyncio.AbstractEventLoop: try: return asyncio.get_running_loop() @@ -43,33 +52,37 @@ def _loop() -> asyncio.AbstractEventLoop: raise RuntimeError(msg) from None -def _pending() -> asyncio.Future: +def _pending() -> asyncio.Future[_R]: return _loop().create_future() -def _resolved(value: object = None) -> asyncio.Future: - future = _pending() +def _resolved(value: _R) -> asyncio.Future[_R]: + future: asyncio.Future[_R] = _pending() future.set_result(value) return future # pipes running, see ReadableStream.pipe_to -_running_pipes: set[asyncio.Future] = set() +_running_pipes: set[asyncio.Future[None]] = set() -def _rejected(error: BaseException) -> asyncio.Future: - future = _pending() +def _rejected(error: BaseException) -> asyncio.Future[_R]: + future: asyncio.Future[_R] = _pending() future.set_exception(error) return future -def _handled(future: asyncio.Future) -> asyncio.Future: +def _retrieve(future: asyncio.Future[_R]) -> None: + _ = future.cancelled() or future.exception() + + +def _handled(future: asyncio.Future[_R]) -> asyncio.Future[_R]: """Marks a future whose exception may be left unretrieved (like the closed promise of a reader).""" - future.add_done_callback(lambda f: f.cancelled() or f.exception()) + future.add_done_callback(_retrieve) return future -def _settle(future: asyncio.Future | None, value: object = None, error: BaseException | None = None) -> None: +def _settle(future: asyncio.Future[_R] | None, value: _R, error: BaseException | None = None) -> None: if future is None or future.done(): return if error is not None: @@ -78,10 +91,16 @@ def _settle(future: asyncio.Future | None, value: object = None, error: BaseExce future.set_result(value) -def _reject(future: asyncio.Future, error: BaseException) -> asyncio.Future: +def _fail(future: asyncio.Future[_R] | None, error: BaseException) -> None: + if future is not None and not future.done(): + future.set_exception(error) + + +def _reject(future: asyncio.Future[_R], error: BaseException) -> asyncio.Future[_R]: """Rejects a pending future, or returns a new rejected one in place of a settled one.""" if future.done(): - return _handled(_rejected(error)) + rejected: asyncio.Future[_R] = _rejected(error) + return _handled(rejected) future.set_exception(error) return future @@ -96,9 +115,12 @@ def _reason_error(reason: object) -> BaseException: 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.""" + method: Callable[..., object] | None if isinstance(obj, dict): - return obj.get(name) - return getattr(obj, name, None) if obj is not None else None + method = cast('dict[str, Callable[..., object]]', obj).get(name) + else: + method = getattr(obj, name, None) if obj is not None else None + return method def _call(obj: object, name: str, *args: object) -> object: @@ -108,7 +130,10 @@ def _call(obj: object, name: str, *args: object) -> object: async def _await(result: object) -> object: - return await result if inspect.isawaitable(result) else result + if inspect.isawaitable(result): + awaited: object = await result + return awaited + return result def _then(result: object, on_done: Callable[[], None], on_error: Callable[[BaseException], None]) -> None: @@ -117,20 +142,21 @@ def _then(result: object, on_done: Callable[[], None], on_error: Callable[[BaseE on_done() return - def done(future: asyncio.Future) -> None: + def done(future: asyncio.Future[object]) -> None: error = asyncio.CancelledError() if future.cancelled() else future.exception() if error is not None: on_error(error) else: on_done() - asyncio.ensure_future(result).add_done_callback(done) + task: asyncio.Future[object] = asyncio.ensure_future(result) + task.add_done_callback(done) -def _future_of(result: object) -> asyncio.Future: +def _future_of(result: object) -> asyncio.Future[None]: """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)) + future: asyncio.Future[None] = _pending() + _then(result, lambda: _settle(future, None), lambda e: _fail(future, e)) return future @@ -151,12 +177,12 @@ def _run( _then(result, on_done, on_error) -class _ReadableWritablePair(Protocol): +class _ReadableWritablePair(Protocol[_T, _O]): @property - def readable(self) -> ReadableStream: ... + def readable(self) -> ReadableStream[_O]: ... @property - def writable(self) -> WritableStream: ... + def writable(self) -> WritableStream[_T]: ... class _PipeOptions(NamedTuple): @@ -166,7 +192,7 @@ class _PipeOptions(NamedTuple): @dataclass -class ReadableStreamReadResult: +class ReadableStreamReadResult(Generic[_T]): """The result of :meth:`ReadableStreamDefaultReader.read`. Args: @@ -174,18 +200,22 @@ class ReadableStreamReadResult: done (:obj:`bool`): Whether the stream is closed and has no more chunks. """ - value: Any + value: _T | None done: bool -class ReadableStreamDefaultController: +#: The result of a read, by another name: tests/idl/expectations.json expects read() not to name it yet +_ReadResult = ReadableStreamReadResult + + +class ReadableStreamDefaultController(Generic[_T]): """Lets an underlying source enqueue chunks, close or error its stream.""" - def __init__(self, stream: ReadableStream, source: object, high_water_mark: float) -> None: + def __init__(self, stream: ReadableStream[_T], source: object, high_water_mark: float) -> None: self._stream = stream self._source = source self._high_water_mark = high_water_mark - self._queue: collections.deque[Any] = collections.deque() + self._queue: collections.deque[_T] = collections.deque() self._close_requested = False self._started = False self._pulling = False @@ -201,7 +231,7 @@ def desired_size(self) -> float | None: return 0 return self._high_water_mark - len(self._queue) - def enqueue(self, chunk: object) -> None: + def enqueue(self, chunk: _T) -> None: """Enqueues a chunk, which fulfills a pending read if there's one. Raises: @@ -210,8 +240,9 @@ def enqueue(self, chunk: object) -> None: if self._close_requested or self._stream._state != 'readable': 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)) + reader = self._stream._reader + if reader is not None and len(reader._read_requests) > 0: + _settle(reader._read_requests.popleft(), ReadableStreamReadResult(chunk, done=False)) else: self._queue.append(chunk) self._call_pull_if_needed() @@ -226,7 +257,7 @@ def close(self) -> None: msg = 'The stream is closed or closing' raise TypeError(msg) self._close_requested = True - if not self._queue: + if len(self._queue) == 0: self._stream._close() def error(self, error: BaseException | None = None) -> None: @@ -249,7 +280,8 @@ def _should_call_pull(self) -> bool: return False if stream._has_read_requests(): return True - return self.desired_size > 0 + desired_size = self.desired_size + return desired_size is not None and desired_size > 0 def _call_pull_if_needed(self) -> None: if not self._should_call_pull(): @@ -267,16 +299,16 @@ def pulled() -> None: _run(self._source, 'pull', self, on_done=pulled, on_error=self.error) - def _read(self, request: asyncio.Future) -> None: - if self._queue: + def _read(self, reader: ReadableStreamDefaultReader[_T], request: asyncio.Future[_ReadResult[_T]]) -> None: + if len(self._queue) > 0: chunk = self._queue.popleft() - if self._close_requested and not self._queue: + if self._close_requested and len(self._queue) == 0: self._stream._close() else: self._call_pull_if_needed() _settle(request, ReadableStreamReadResult(chunk, done=False)) else: - self._stream._reader._read_requests.append(request) + reader._read_requests.append(request) self._call_pull_if_needed() def _cancel(self, reason: object) -> object: @@ -287,7 +319,7 @@ def _cancel(self, reason: object) -> object: desiredSize = desired_size -class ReadableStream: +class ReadableStream(Generic[_T]): """A stream of chunks to read (https://developer.mozilla.org/en-US/docs/Web/API/ReadableStream). Args: @@ -299,7 +331,7 @@ class ReadableStream: def __init__(self, underlying_source: object = None, high_water_mark: float = 1) -> None: self._state = 'readable' self._stored_error: BaseException | None = None - self._reader: ReadableStreamDefaultReader | None = None + self._reader: ReadableStreamDefaultReader[_T] | None = None self._controller = ReadableStreamDefaultController(self, underlying_source, high_water_mark) self._controller._start() @@ -308,14 +340,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[_T]: """Returns a reader, which holds the stream until it's released. Raises :obj:`TypeError` if the stream is locked. """ return ReadableStreamDefaultReader(self) - def cancel(self, reason: object = None) -> asyncio.Future: + def cancel(self, reason: object = None) -> asyncio.Future[None]: """Cancels the stream: its source stops and its chunks are dropped. Returns: @@ -327,12 +359,12 @@ def cancel(self, reason: object = None) -> asyncio.Future: def pipe_to( self, - destination: WritableStream, + destination: WritableStream[_T], *, prevent_close: bool = False, prevent_abort: bool = False, prevent_cancel: bool = False, - ) -> asyncio.Future: + ) -> asyncio.Future[None]: """Writes every chunk of the stream to a writable stream, waiting for it when it's full. Args: @@ -357,7 +389,7 @@ def pipe_to( @staticmethod async def _pipe( - reader: ReadableStreamDefaultReader, writer: WritableStreamDefaultWriter, options: _PipeOptions + reader: ReadableStreamDefaultReader[_T], writer: WritableStreamDefaultWriter[_T], options: _PipeOptions ) -> None: try: await _pipe_chunks(reader, writer, prevent_close=options.prevent_close) @@ -371,7 +403,7 @@ async def _pipe( reader.release_lock() writer.release_lock() - def pipe_through(self, transform: _ReadableWritablePair, **options: bool) -> ReadableStream: + def pipe_through(self, transform: _ReadableWritablePair[_T, _O], **options: bool) -> ReadableStream[_O]: """Pipes the stream into the writable side of a transform (like :obj:`TransformStream`). Args: @@ -381,10 +413,10 @@ def pipe_through(self, transform: _ReadableWritablePair, **options: bool) -> Rea Returns: :obj:`ReadableStream`: The readable side of the transform. """ - _handled(self.pipe_to(transform.writable, **options)) + _ = _handled(self.pipe_to(transform.writable, **options)) return transform.readable - def values(self, *, prevent_cancel: bool = False) -> AsyncIterator[Any]: + def values(self, *, prevent_cancel: bool = False) -> AsyncIterator[_T]: """Iterates over the chunks, like ``async for``. Stopping early cancels the stream once the iterator is finalized, right away with @@ -398,7 +430,7 @@ def values(self, *, prevent_cancel: bool = False) -> AsyncIterator[Any]: """ return _iterate(self.get_reader(), prevent_cancel=prevent_cancel) - def __aiter__(self) -> AsyncIterator[Any]: + def __aiter__(self) -> AsyncIterator[_T]: return self.values() def _close(self) -> None: @@ -407,7 +439,7 @@ def _close(self) -> None: self._state = 'closed' if self._reader is not None: self._reader._settle_read_requests(ReadableStreamReadResult(None, done=True)) - _settle(self._reader._closed) + _settle(self._reader._closed, None) def _error(self, error: BaseException) -> None: if self._state != 'readable': @@ -415,17 +447,21 @@ def _error(self, error: BaseException) -> None: self._state = 'errored' self._stored_error = error if self._reader is not None: - self._reader._settle_read_requests(error=error) - _settle(self._reader._closed, error=error) + self._reader._fail_read_requests(error) + _fail(self._reader._closed, error) def _has_read_requests(self) -> bool: - return self._reader is not None and bool(self._reader._read_requests) + return self._reader is not None and len(self._reader._read_requests) > 0 + + def _error_stored(self) -> BaseException: + """The error of an errored stream.""" + return _error_or_default(self._stored_error) - def _cancel(self, reason: object) -> asyncio.Future: + def _cancel(self, reason: object) -> asyncio.Future[None]: if self._state == 'closed': - return _resolved() + return _resolved(None) if self._state == 'errored': - return _rejected(self._stored_error) + return _rejected(self._error_stored()) self._close() return _future_of(self._controller._cancel(reason)) @@ -437,7 +473,7 @@ def _cancel(self, reason: object) -> asyncio.Future: pipeThrough = pipe_through -class ReadableStreamDefaultReader: +class ReadableStreamDefaultReader(Generic[_T]): """Reads the chunks of a stream, which it locks until :meth:`release_lock`. Args: @@ -447,25 +483,25 @@ class ReadableStreamDefaultReader: TypeError: If the stream is locked. """ - def __init__(self, stream: ReadableStream) -> None: + def __init__(self, stream: ReadableStream[_T]) -> None: if stream.locked: 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()) + self._stream: ReadableStream[_T] | None = stream + self._read_requests: collections.deque[asyncio.Future[_ReadResult[_T]]] = collections.deque() + self._closed: asyncio.Future[None] = _handled(_pending()) stream._reader = self if stream._state == 'closed': - _settle(self._closed) + _settle(self._closed, None) elif stream._state == 'errored': - _settle(self._closed, error=stream._stored_error) + _settle(self._closed, None, stream._stored_error) @property - def closed(self) -> asyncio.Future: + def closed(self) -> asyncio.Future[None]: """:obj:`asyncio.Future`: Done once the stream is closed, failed if it errors or the lock is released.""" return self._closed - def read(self) -> asyncio.Future: + def read(self) -> asyncio.Future[_ReadResult[_T]]: """Reads the next chunk. Returns: @@ -478,12 +514,12 @@ def read(self) -> asyncio.Future: if stream._state == 'closed': return _resolved(ReadableStreamReadResult(None, done=True)) if stream._state == 'errored': - return _rejected(stream._stored_error) - request = _pending() - stream._controller._read(request) + return _rejected(stream._error_stored()) + request: asyncio.Future[_ReadResult[_T]] = _pending() + stream._controller._read(self, request) return request - def cancel(self, reason: object = None) -> asyncio.Future: + def cancel(self, reason: object = None) -> asyncio.Future[None]: """Cancels the stream (see :meth:`ReadableStream.cancel`).""" if self._stream is None: return _rejected(TypeError('The reader is released')) @@ -495,24 +531,33 @@ def release_lock(self) -> None: if stream is None: return error = TypeError('The reader is released') - self._settle_read_requests(error=error) + self._fail_read_requests(error) if stream._state == 'readable': - _settle(self._closed, error=error) + _fail(self._closed, error) else: self._closed = _handled(_rejected(error)) stream._reader = None self._stream = 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) + def _settle_read_requests(self, value: ReadableStreamReadResult[_T]) -> None: + while len(self._read_requests) > 0: + _settle(self._read_requests.popleft(), value) + + def _fail_read_requests(self, error: BaseException) -> None: + while len(self._read_requests) > 0: + _fail(self._read_requests.popleft(), error) #: Alias for :meth:`release_lock` releaseLock = release_lock +def _chunk(result: ReadableStreamReadResult[_T]) -> _T: + """The chunk of a result that isn't done, which may be :obj:`None` too.""" + return cast('_T', result.value) + + async def _pipe_chunks( - reader: ReadableStreamDefaultReader, writer: WritableStreamDefaultWriter, *, prevent_close: bool + reader: ReadableStreamDefaultReader[_T], writer: WritableStreamDefaultWriter[_T], *, prevent_close: bool ) -> None: while True: await writer.ready @@ -522,25 +567,29 @@ async def _pipe_chunks( await writer.close() return # writes aren't awaited, like in the specification - _handled(writer.write(result.value)) + _ = _handled(writer.write(_chunk(result))) async def _stop_pipe( - reader: ReadableStreamDefaultReader, - writer: WritableStreamDefaultWriter, + reader: ReadableStreamDefaultReader[_T], + writer: WritableStreamDefaultWriter[_T], 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'}: + stream = writer._stream + if stream is None: + msg = 'The writer is released' + raise TypeError(msg) + if 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]: +async def _iterate(reader: ReadableStreamDefaultReader[_T], *, prevent_cancel: bool) -> AsyncIterator[_T]: done = False try: while True: @@ -548,7 +597,7 @@ async def _iterate(reader: ReadableStreamDefaultReader, *, prevent_cancel: bool) if result.done: done = True return - yield result.value + yield _chunk(result) finally: # runs once the generator is finalized, so an early stop cancels then if reader._stream is not None: @@ -561,15 +610,15 @@ async def _iterate(reader: ReadableStreamDefaultReader, *, prevent_cancel: bool) _CLOSE = object() -class WritableStreamDefaultController: +class WritableStreamDefaultController(Generic[_T]): """Lets an underlying sink error its stream.""" - def __init__(self, stream: WritableStream, sink: object, high_water_mark: float) -> None: + def __init__(self, stream: WritableStream[_T], 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: collections.deque[tuple[object, asyncio.Future]] = collections.deque() + self._queue: collections.deque[tuple[object, asyncio.Future[None]]] = collections.deque() self._started = False self._in_flight = False @@ -592,19 +641,19 @@ def failed(error: BaseException) -> None: _then(_call(self._sink, 'start', self), started, failed) - def _write(self, chunk: object, future: asyncio.Future) -> None: + def _write(self, chunk: _T, future: asyncio.Future[None]) -> None: self._queue.append((chunk, future)) # advance first: a sink done right away leaves no backpressure to signal self._advance() self._stream._update_backpressure() - def _close(self, future: asyncio.Future) -> None: + def _close(self, future: asyncio.Future[None]) -> None: self._queue.append((_CLOSE, future)) self._advance() def _advance(self) -> None: stream = self._stream - if not self._started or self._in_flight or not self._queue: + if not self._started or self._in_flight or len(self._queue) == 0: return if stream._state == 'erroring': stream._finish_erroring() @@ -623,34 +672,34 @@ def failed(error: BaseException) -> None: 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: + def _settle_in_flight(self, future: asyncio.Future[None], error: BaseException | None = None) -> None: self._in_flight = False - self._queue.popleft() - _settle(future, error=error) + _ = self._queue.popleft() + _settle(future, None, error) - def _written(self, future: asyncio.Future) -> None: + def _written(self, future: asyncio.Future[None]) -> None: self._settle_in_flight(future) self._stream._update_backpressure() self._advance() - def _closed(self, future: asyncio.Future) -> None: + def _closed(self, future: asyncio.Future[None]) -> None: self._settle_in_flight(future) stream = self._stream stream._state = 'closed' if stream._writer is not None: - _settle(stream._writer._closed) + _settle(stream._writer._closed, None) def _reject_queue(self, error: BaseException) -> None: """Rejects the queued requests but the one in flight.""" in_flight = self._queue.popleft() if self._in_flight else None for _, future in self._queue: - _settle(future, error=error) + _fail(future, error) self._queue.clear() if in_flight is not None: self._queue.append(in_flight) -class WritableStream: +class WritableStream(Generic[_T]): """A stream to write chunks to (https://developer.mozilla.org/en-US/docs/Web/API/WritableStream). Args: @@ -663,7 +712,7 @@ class WritableStream: def __init__(self, underlying_sink: object = None, high_water_mark: float = 1) -> None: self._state = 'writable' self._stored_error: BaseException | None = None - self._writer: WritableStreamDefaultWriter | None = None + self._writer: WritableStreamDefaultWriter[_T] | None = None self._close_requested = False self._controller = WritableStreamDefaultController(self, underlying_sink, high_water_mark) self._controller._start() @@ -673,14 +722,14 @@ 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[_T]: """Returns a writer, which holds the stream until it's released. Raises :obj:`TypeError` if the stream is locked. """ return WritableStreamDefaultWriter(self) - def close(self) -> asyncio.Future: + def close(self) -> asyncio.Future[None]: """Closes the stream once the chunks written before are. Returns: @@ -690,7 +739,7 @@ def close(self) -> asyncio.Future: return _rejected(TypeError('The stream is locked')) return self._close() - def abort(self, reason: object = None) -> asyncio.Future: + def abort(self, reason: object = None) -> asyncio.Future[None]: """Aborts the stream: queued chunks are dropped and the sink is aborted. Returns: @@ -700,17 +749,17 @@ def abort(self, reason: object = None) -> asyncio.Future: return _rejected(TypeError('The stream is locked')) return self._abort(reason) - def _close(self) -> asyncio.Future: + def _close(self) -> asyncio.Future[None]: 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() + future: asyncio.Future[None] = _pending() self._controller._close(future) return future - def _abort(self, reason: object) -> asyncio.Future: + def _abort(self, reason: object) -> asyncio.Future[None]: if self._state in {'closed', 'errored'}: - return _resolved() + return _resolved(None) error = reason if isinstance(reason, BaseException) else asyncio.CancelledError(reason) self._controller._reject_queue(error) self._state = 'errored' @@ -722,16 +771,18 @@ def _start_erroring(self, error: BaseException) -> None: self._state = 'erroring' self._stored_error = error if self._writer is not None: - _settle(self._writer._ready, error=error) - self._writer._ready = _handled(_rejected(error)) + _fail(self._writer._ready, error) + rejected: asyncio.Future[None] = _rejected(error) + self._writer._ready = _handled(rejected) if not self._controller._in_flight and self._controller._started: self._finish_erroring() def _finish_erroring(self) -> None: self._state = 'errored' - self._controller._reject_queue(self._stored_error) - self._reject_writer(self._stored_error) - _call(self._controller._sink, 'abort', self._stored_error) + error = _error_or_default(self._stored_error) + self._controller._reject_queue(error) + self._reject_writer(error) + _ = _call(self._controller._sink, 'abort', self._stored_error) def _deal_with_rejection(self, error: BaseException) -> None: if self._state == 'writable': @@ -743,7 +794,7 @@ def _reject_writer(self, error: BaseException) -> None: writer = self._writer if writer is None: return - _settle(writer._closed, error=error) + _fail(writer._closed, error) writer._ready = _reject(writer._ready, error) def _update_backpressure(self) -> None: @@ -752,15 +803,16 @@ def _update_backpressure(self) -> None: return backpressure = self._controller._desired_size() <= 0 if backpressure and writer._ready.done(): - writer._ready = _handled(_pending()) + pending: asyncio.Future[None] = _pending() + writer._ready = _handled(pending) elif not backpressure: - _settle(writer._ready) + _settle(writer._ready, None) #: Alias for :meth:`get_writer` getWriter = get_writer -class WritableStreamDefaultWriter: +class WritableStreamDefaultWriter(Generic[_T]): """Writes chunks to a stream, which it locks until :meth:`release_lock`. Args: @@ -770,31 +822,31 @@ class WritableStreamDefaultWriter: TypeError: If the stream is locked. """ - def __init__(self, stream: WritableStream) -> None: + def __init__(self, stream: WritableStream[_T]) -> None: if stream.locked: msg = 'The stream is locked' raise TypeError(msg) - self._stream: WritableStream | None = stream + self._stream: WritableStream[_T] | None = stream stream._writer = self - self._closed = _handled(_pending()) - self._ready = _handled(_pending()) + self._closed: asyncio.Future[None] = _handled(_pending()) + self._ready: asyncio.Future[None] = _handled(_pending()) if stream._state == 'writable': if stream._controller._desired_size() > 0 or stream._close_requested: - _settle(self._ready) + _settle(self._ready, None) elif stream._state == 'closed': - _settle(self._ready) - _settle(self._closed) + _settle(self._ready, None) + _settle(self._closed, None) else: - _settle(self._ready, error=stream._stored_error) - _settle(self._closed, error=stream._stored_error) + _settle(self._ready, None, stream._stored_error) + _settle(self._closed, None, stream._stored_error) @property - def closed(self) -> asyncio.Future: + def closed(self) -> asyncio.Future[None]: """:obj:`asyncio.Future`: Done once the stream is closed, failed if it errors or the lock is released.""" return self._closed @property - def ready(self) -> asyncio.Future: + def ready(self) -> asyncio.Future[None]: """:obj:`asyncio.Future`: Done when the stream can take a chunk without queuing it beyond its limit.""" return self._ready @@ -815,7 +867,7 @@ def desired_size(self) -> float | None: return 0 return stream._controller._desired_size() - def write(self, chunk: object) -> asyncio.Future: + def write(self, chunk: _T) -> asyncio.Future[None]: """Writes a chunk. Returns: @@ -825,20 +877,20 @@ def write(self, chunk: object) -> asyncio.Future: if stream is None: return _rejected(TypeError('The writer is released')) if stream._state in {'errored', 'erroring'}: - return _rejected(stream._stored_error) + return _rejected(_error_or_default(stream._stored_error)) if stream._close_requested or stream._state == 'closed': return _rejected(TypeError('The stream is closed or closing')) - future = _pending() + future: asyncio.Future[None] = _pending() stream._controller._write(chunk, future) return future - def close(self) -> asyncio.Future: + def close(self) -> asyncio.Future[None]: """Closes the stream (see :meth:`WritableStream.close`).""" if self._stream is None: return _rejected(TypeError('The writer is released')) return self._stream._close() - def abort(self, reason: object = None) -> asyncio.Future: + def abort(self, reason: object = None) -> asyncio.Future[None]: """Aborts the stream (see :meth:`WritableStream.abort`).""" if self._stream is None: return _rejected(TypeError('The writer is released')) @@ -861,10 +913,10 @@ def release_lock(self) -> None: releaseLock = release_lock -class TransformStreamDefaultController: +class TransformStreamDefaultController(Generic[_T, _O]): """Lets a transformer enqueue chunks to the readable side, error or terminate its stream.""" - def __init__(self, stream: TransformStream) -> None: + def __init__(self, stream: TransformStream[_T, _O]) -> None: self._stream = stream @property @@ -872,7 +924,7 @@ 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: object) -> None: + def enqueue(self, chunk: _O) -> None: """Enqueues a chunk to the readable side.""" self._stream._readable._controller.enqueue(chunk) @@ -893,7 +945,7 @@ def terminate(self) -> None: desiredSize = desired_size -class TransformStream: +class TransformStream(Generic[_T, _O]): """A pair of streams where what's written is transformed and read. See https://developer.mozilla.org/en-US/docs/Web/API/TransformStream. @@ -906,48 +958,50 @@ class TransformStream: def __init__(self, transformer: object = None) -> None: self._transformer = transformer - self._controller = TransformStreamDefaultController(self) + self._controller: TransformStreamDefaultController[_T, _O] = TransformStreamDefaultController(self) # settled by a pull of the readable side, which relieves backpressure - 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) + self._pull_waiter: asyncio.Future[None] | None = None + self._readable: ReadableStream[_O] = ReadableStream(_TransformSource(self), high_water_mark=0) + self._writable: WritableStream[_T] = WritableStream(_TransformSink(self), high_water_mark=1) + _ = _call(transformer, 'start', self._controller) @property - def readable(self) -> ReadableStream: + def readable(self) -> ReadableStream[_O]: """:obj:`ReadableStream`: The transformed chunks.""" return self._readable @property - def writable(self) -> WritableStream: + def writable(self) -> WritableStream[_T]: """:obj:`WritableStream`: The chunks to transform.""" return self._writable -class _TransformSink: - def __init__(self, stream: TransformStream) -> None: +class _TransformSink(Generic[_T, _O]): + def __init__(self, stream: TransformStream[_T, _O]) -> None: self._stream = stream - async def write(self, chunk: object, _controller: WritableStreamDefaultController) -> None: + def _has_room(self) -> bool: + readable = self._stream._readable + desired_size = readable._controller.desired_size + return desired_size is None or desired_size > 0 or readable._has_read_requests() + + async def write(self, chunk: _T, _controller: WritableStreamDefaultController[_T]) -> 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._has_read_requests() - ): - stream._pull_waiter = _pending() - await stream._pull_waiter + while stream._readable._state == 'readable' and not self._has_room(): + waiter: asyncio.Future[None] = _pending() + stream._pull_waiter = waiter + await waiter transform = _member(stream._transformer, 'transform') if transform is None: - stream._controller.enqueue(chunk) + # chunks pass unchanged without a transform + stream._controller.enqueue(cast('_O', chunk)) else: await _await(transform(chunk, stream._controller)) async def close(self) -> None: stream = self._stream - await _await(_call(stream._transformer, 'flush', stream._controller)) + _ = await _await(_call(stream._transformer, 'flush', stream._controller)) if stream._readable._state == 'readable' and not stream._readable._controller._close_requested: stream._readable._controller.close() @@ -955,12 +1009,12 @@ def abort(self, reason: object) -> None: self._stream._readable._controller.error(_reason_error(reason)) -class _TransformSource: - def __init__(self, stream: TransformStream) -> None: +class _TransformSource(Generic[_T, _O]): + def __init__(self, stream: TransformStream[_T, _O]) -> None: self._stream = stream - def pull(self, _controller: ReadableStreamDefaultController) -> None: - _settle(self._stream._pull_waiter) + def pull(self, _controller: ReadableStreamDefaultController[_O]) -> None: + _settle(self._stream._pull_waiter, None) def cancel(self, reason: object) -> None: self._stream._writable._controller.error(_reason_error(reason)) diff --git a/python-webrtc/python/webrtc/utils/events.py b/python-webrtc/python/webrtc/utils/events.py index a00c216..ce38249 100644 --- a/python-webrtc/python/webrtc/utils/events.py +++ b/python-webrtc/python/webrtc/utils/events.py @@ -11,13 +11,16 @@ import asyncio import inspect -from typing import Callable, NamedTuple, TypeVar, overload +from typing import TYPE_CHECKING, Callable, NamedTuple, Protocol, TypeVar, cast, overload + +from typing_extensions import Never import webrtc from webrtc.utils.task_queue import TaskQueue Handler = Callable[['webrtc.Event'], object] -_H = TypeVar('_H', bound=Handler) +# any one-argument callable: a handler takes the event subclass of its event, like RTCTrackEvent +_H = TypeVar('_H', bound=Callable[[Never], object]) #: The tasks of the coroutine handlers, referenced until they're done _handler_tasks: set[asyncio.Future[object]] = set() @@ -48,7 +51,7 @@ def __init__(self, target: EventTarget) -> None: def __call__(self, name: str, *args: object) -> None: # a libwebrtc thread, with the GIL held: only schedule - registrations = self.__dict__.get('registrations') + registrations: dict[str, list[_Registration]] | None = self.__dict__.get('registrations') if registrations is None: # the garbage collector cleared this object (in a cycle with its target) before the native one let go return @@ -58,7 +61,7 @@ def __call__(self, name: str, *args: object) -> None: primary_loop = self.primary_loop = next( (r.loop for regs in list(registrations.values()) for r in list(regs) if not r.loop.is_closed()), None ) - loops = [primary_loop] if primary_loop else [] + loops: list[asyncio.AbstractEventLoop] = [primary_loop] if primary_loop is not None else [] for registration in registrations.get(name, ()): if registration.loop not in loops: loops.append(registration.loop) @@ -77,7 +80,7 @@ def deliver(self, loop: asyncio.AbstractEventLoop, name: str, args: tuple[object 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] - if not registrations: + if len(registrations) == 0: return event = self.target._create_event(name, *args) if event is None: @@ -110,14 +113,20 @@ def add(self, name: str, handler: Handler, *, once: bool) -> None: return # like addEventListener, a handler is registered once registrations.append(_Registration(handler, loop, once)) - def remove(self, name: str, handler: Handler | None) -> None: + def remove(self, name: str, handler: Callable[[Never], object] | None) -> None: if handler is None: - self.registrations.pop(name, None) + _ = self.registrations.pop(name, None) return registrations = self.registrations.get(name, []) self.registrations[name] = [r for r in registrations if r.handler != handler] +class _NativeEventTarget(Protocol): + """A native object that emits events, through the listeners it holds.""" + + _listeners: _Listeners | None + + class EventTarget: """Mixin of :obj:`webrtc.WebRTCObject` subclasses that emit events. @@ -137,13 +146,21 @@ async def on_candidate(event): #: Names of the events the object emits _events: tuple[str, ...] = () - def _listeners(self, *, create: bool) -> _Listeners | None: + if TYPE_CHECKING: + + @property + def _native_obj(self) -> _NativeEventTarget: ... + + def _listeners(self) -> _Listeners | None: + return self._native_obj._listeners + + def _created_listeners(self) -> _Listeners: native = self._native_obj listeners = native._listeners - if listeners is None and create: + if listeners is None: listeners = _Listeners(self) # before the native object gets it: it delivers the events it held right away - listeners.ensure_primary_loop() + _ = listeners.ensure_primary_loop() native._listeners = listeners return listeners @@ -154,7 +171,7 @@ def _attach(self) -> None: a loop. """ if _running_loop() is not None: - self._listeners(create=True).ensure_primary_loop() + _ = self._created_listeners().ensure_primary_loop() def _check_event(self, name: str) -> None: if name not in self._events: @@ -168,12 +185,13 @@ def _add(self, name: str, handler: _H | None, *, once: bool) -> _H | Callable[[_ 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) + # the handler takes the event of its name, which only the native object knows + self._created_listeners().add(name, cast('Handler', handler), once=once) return handler 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) + listeners = self._listeners() if listeners is not None: listeners.deliver(asyncio.get_running_loop(), name, args) @@ -233,7 +251,7 @@ def once(self, name: str, handler: _H | None = None) -> _H | Callable[[_H], _H]: """ return self._add(name, handler, once=True) - def off(self, name: str, handler: Handler | None = None) -> None: + def off(self, name: str, handler: Callable[[Never], object] | None = None) -> None: """Removes a handler of an event, or every handler of the event. Args: @@ -244,6 +262,6 @@ def off(self, name: str, handler: Handler | None = None) -> None: ValueError: If the object has no such event. """ self._check_event(name) - listeners = self._listeners(create=False) + listeners = self._listeners() if listeners is not None: listeners.remove(name, handler) diff --git a/python-webrtc/python/webrtc/utils/names.py b/python-webrtc/python/webrtc/utils/names.py index 2dfbe95..b1ba310 100644 --- a/python-webrtc/python/webrtc/utils/names.py +++ b/python-webrtc/python/webrtc/utils/names.py @@ -10,7 +10,7 @@ from __future__ import annotations import re -from typing import TYPE_CHECKING, Any, Generic, TypeVar, overload +from typing import TYPE_CHECKING, Generic, TypeVar, overload if TYPE_CHECKING: from collections.abc import Iterable, Mapping @@ -29,7 +29,7 @@ def snake_case(name: str) -> str: return ''.join(f'_{c.lower()}' if c.isupper() else c for c in name) -def members(value: Mapping[str, Any], names: Iterable[str]) -> dict[str, Any]: +def members(value: Mapping[str, object], names: Iterable[str]) -> dict[str, object]: """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} @@ -62,7 +62,7 @@ def __set__(self, obj: object, value: _T) -> None: setattr(obj, self.name, value) -def alias(name: str) -> Alias[Any]: +def alias(name: str) -> Alias[_T]: """The :obj:`Alias` of an attribute, like :func:`dataclasses.field` for a field. Args: diff --git a/python-webrtc/python/webrtc/utils/native_calls.py b/python-webrtc/python/webrtc/utils/native_calls.py index c59b482..da21b0b 100644 --- a/python-webrtc/python/webrtc/utils/native_calls.py +++ b/python-webrtc/python/webrtc/utils/native_calls.py @@ -10,7 +10,7 @@ from __future__ import annotations import asyncio -from typing import TYPE_CHECKING, Callable, TypeVar +from typing import TYPE_CHECKING, Callable, TypeVar, overload from webrtc.utils.task_queue import TaskQueue @@ -20,15 +20,29 @@ import wrtc _P = ParamSpec('_P') + _OnFailure = Callable[[wrtc.RTCCallbackException], None] _T = TypeVar('_T') +@overload async def call_native( - method: Callable[Concatenate[Callable[[_T], None], Callable[[wrtc.RTCCallbackException], None], _P], None], + method: Callable[Concatenate[Callable[[], None], _OnFailure, _P], None], *args: _P.args, **kwargs: _P.kwargs +) -> None: ... + + +@overload +async def call_native( + method: Callable[Concatenate[Callable[[_T], None], _OnFailure, _P], None], *args: _P.args, **kwargs: _P.kwargs +) -> _T: ... + + +async def call_native( + method: Callable[Concatenate[Callable[[], None], _OnFailure, _P], None] + | Callable[Concatenate[Callable[[_T], None], _OnFailure, _P], None], *args: _P.args, **kwargs: _P.kwargs, -) -> _T: +) -> _T | None: """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 @@ -46,7 +60,7 @@ async def call_native( The error passed to ``on_failure``, as a Python exception. """ loop = asyncio.get_running_loop() - future = loop.create_future() + future: asyncio.Future[_T | None] = loop.create_future() def settle(result: _T | None, error: wrtc.RTCCallbackException | None) -> None: # the caller may have been canceled meanwhile diff --git a/python-webrtc/python/webrtc/utils/operations.py b/python-webrtc/python/webrtc/utils/operations.py index 366fad2..39a1c18 100644 --- a/python-webrtc/python/webrtc/utils/operations.py +++ b/python-webrtc/python/webrtc/utils/operations.py @@ -22,7 +22,7 @@ from typing import TYPE_CHECKING, Callable if TYPE_CHECKING: - from collections.abc import AsyncIterator + from collections.abc import AsyncGenerator class OperationsChain: @@ -45,7 +45,7 @@ def busy(self) -> bool: return self._last is not None and not self._last.done() @contextlib.asynccontextmanager - async def operation(self) -> AsyncIterator[None]: + async def operation(self) -> AsyncGenerator[None, None]: """Chains an operation after the ones that are running.""" previous = self._last done = asyncio.get_running_loop().create_future() diff --git a/python-webrtc/python/webrtc/utils/task_queue.py b/python-webrtc/python/webrtc/utils/task_queue.py index 1d7286f..d6adfbe 100644 --- a/python-webrtc/python/webrtc/utils/task_queue.py +++ b/python-webrtc/python/webrtc/utils/task_queue.py @@ -13,7 +13,10 @@ import collections import threading import weakref -from typing import Callable, NamedTuple +from typing import TYPE_CHECKING, Callable, NamedTuple + +if TYPE_CHECKING: + from collections.abc import Iterable class _Item(NamedTuple): @@ -54,6 +57,15 @@ def __init__(self, loop: asyncio.AbstractEventLoop) -> None: def _loop(self) -> asyncio.AbstractEventLoop | None: return self._loop_ref() + @property + def _live_loop(self) -> asyncio.AbstractEventLoop: + """The loop, from a callback it runs, so alive.""" + loop = self._loop_ref() + if loop is None: + msg = 'the loop of the queue is gone' + raise RuntimeError(msg) + return loop + @classmethod def of(cls, loop: asyncio.AbstractEventLoop) -> TaskQueue: """Returns the queue of a loop. @@ -70,8 +82,10 @@ def of(cls, loop: asyncio.AbstractEventLoop) -> TaskQueue: except TypeError: attributes = None if attributes is not None: - queue = attributes.get(cls._ATTRIBUTE) - return queue if queue is not None else attributes.setdefault(cls._ATTRIBUTE, cls(loop)) + queue: TaskQueue | None = attributes.get(cls._ATTRIBUTE) + if queue is None: + queue = attributes.setdefault(cls._ATTRIBUTE, cls(loop)) + return queue queue = cls._queues.get(loop) if queue is None: # the loops closed meanwhile won't run what's queued @@ -130,7 +144,7 @@ def post( self._items.clear() return try: - loop.call_soon_threadsafe(self._run) + _ = loop.call_soon_threadsafe(self._run) except RuntimeError: self._items.clear() @@ -138,8 +152,8 @@ def _others_ready(self) -> bool: """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: + ready: Iterable[object] | None = getattr(self._loop, '_ready', None) + for handle in ready if ready is not None else (): callback = getattr(handle, '_callback', None) if isinstance(handle, asyncio.TimerHandle) or getattr(callback, '__self__', None) in {self, self._loop}: continue @@ -155,13 +169,13 @@ 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): - self._loop.call_soon(self._settle, deferred + 1) + _ = self._live_loop.call_soon(self._settle, deferred + 1) else: self._resumed = False def _run(self, deferred: int = 0) -> None: if self._defers(deferred, resumed=self._resumed): - self._loop.call_soon(self._run, deferred + 1) + _ = self._live_loop.call_soon(self._run, deferred + 1) return more = True @@ -174,7 +188,7 @@ def _run(self, deferred: int = 0) -> None: raise finally: if more: - self._loop.call_soon(self._run) + _ = self._live_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.""" @@ -196,8 +210,8 @@ def _run_next(self) -> bool: item = self._items.popleft() if item.resumes: self._resumed = True - self._loop.call_soon(self._settle) - item.callback(*item.args) + _ = self._live_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 diff --git a/stubs/wrtc/__init__.pyi b/stubs/wrtc/__init__.pyi index 1a93513..a8ff505 100644 --- a/stubs/wrtc/__init__.pyi +++ b/stubs/wrtc/__init__.pyi @@ -1,10 +1,14 @@ from __future__ import annotations import collections.abc import typing +import typing_extensions import webrtc.enums +import webrtc.exceptions +import webrtc.models.media_track_constraints +import webrtc.utils.events __all__: list[str] = ['ConfigurationInit', 'DataChannelMessage', 'IceCandidateInit', 'IceServerInit', 'MediaStream', 'MediaStreamTrack', 'MediaStreamTrackProcessor', 'PeerConnectionFactory', 'PythonWebRTCException', 'PythonWebRTCExceptionBase', 'RTCCallbackException', 'RTCCertificate', 'RTCDTMFSender', 'RTCDataChannel', 'RTCDtlsTransport', 'RTCIceTransport', 'RTCPeerConnection', 'RTCRtpReceiver', 'RTCRtpSender', 'RTCRtpTransceiver', 'RTCSctpTransport', 'RTCSessionDescription', 'RTCSessionDescriptionInit', 'RtcpParameters', 'RtpCapabilities', 'RtpCodec', 'RtpCodecCapability', 'RtpCodecParameters', 'RtpEncodingParameters', 'RtpExtension', 'RtpHeaderExtensionCapability', 'RtpParameters', 'RtpTransceiverInit', 'SdpParseException', 'TrackGenerator', 'VideoFrameBuffer', 'copyAudioSamples', 'getUserMedia', 'ping'] class RTCCallbackException: - def toPython(self) -> typing.Any: + def toPython(self) -> webrtc.exceptions.RTCException: ... class PythonWebRTCExceptionBase(Exception): pass @@ -14,8 +18,13 @@ class SdpParseException(PythonWebRTCExceptionBase): pass class RTCSessionDescriptionInit: sdp: str - type: webrtc.enums.RTCSdpType - def __init__(self, arg0: webrtc.enums.RTCSdpType, arg1: str) -> None: + @property + def type(self) -> webrtc.enums.RTCSdpType: + ... + @type.setter + def type(self, arg0: webrtc.enums.RTCSdpType | webrtc.enums.RTCSdpTypeValue) -> None: + ... + def __init__(self, arg0: webrtc.enums.RTCSdpType | webrtc.enums.RTCSdpTypeValue, arg1: str) -> None: ... class RTCSessionDescription: def __init__(self, arg0: RTCSessionDescriptionInit) -> None: @@ -30,7 +39,7 @@ class RTCSessionDescription: def type(self) -> webrtc.enums.RTCSdpType: ... class IceCandidateInit: - def kwargs(self) -> dict: + def kwargs(self) -> _IceCandidateKwargs: ... class RTCCertificate: @staticmethod @@ -54,10 +63,30 @@ class IceServerInit: ... class ConfigurationInit: alwaysNegotiateDataChannels: bool - bundlePolicy: webrtc.enums.RTCBundlePolicy - iceTransportPolicy: webrtc.enums.RTCIceTransportPolicy - rtcpMuxPolicy: webrtc.enums.RTCRtcpMuxPolicy - rtpHeaderEncryptionPolicy: webrtc.enums.RTCRtpHeaderEncryptionPolicy + @property + def bundlePolicy(self) -> webrtc.enums.RTCBundlePolicy: + ... + @bundlePolicy.setter + def bundlePolicy(self, arg0: webrtc.enums.RTCBundlePolicy | webrtc.enums.RTCBundlePolicyValue) -> None: + ... + @property + def iceTransportPolicy(self) -> webrtc.enums.RTCIceTransportPolicy: + ... + @iceTransportPolicy.setter + def iceTransportPolicy(self, arg0: webrtc.enums.RTCIceTransportPolicy | webrtc.enums.RTCIceTransportPolicyValue) -> None: + ... + @property + def rtcpMuxPolicy(self) -> webrtc.enums.RTCRtcpMuxPolicy: + ... + @rtcpMuxPolicy.setter + def rtcpMuxPolicy(self, arg0: webrtc.enums.RTCRtcpMuxPolicy | webrtc.enums.RTCRtcpMuxPolicyValue) -> None: + ... + @property + def rtpHeaderEncryptionPolicy(self) -> webrtc.enums.RTCRtpHeaderEncryptionPolicy: + ... + @rtpHeaderEncryptionPolicy.setter + def rtpHeaderEncryptionPolicy(self, arg0: webrtc.enums.RTCRtpHeaderEncryptionPolicy | webrtc.enums.RTCRtpHeaderEncryptionPolicyValue) -> None: + ... def __init__(self) -> None: ... @property @@ -85,7 +114,12 @@ class ConfigurationInit: def portRange(self, arg0: tuple[typing.SupportsInt | typing.SupportsIndex, typing.SupportsInt | typing.SupportsIndex] | None) -> None: ... class RtpCodec: - kind: webrtc.enums.MediaType + @property + def kind(self) -> webrtc.enums.MediaType: + ... + @kind.setter + def kind(self, arg0: webrtc.enums.MediaType | webrtc.enums.MediaTypeValue) -> None: + ... name: str def __init__(self) -> None: ... @@ -140,7 +174,12 @@ class RtpExtension: def id(self, arg1: typing.SupportsInt | typing.SupportsIndex) -> None: ... class RtpHeaderExtensionCapability: - direction: webrtc.enums.TransceiverDirection + @property + def direction(self) -> webrtc.enums.TransceiverDirection: + ... + @direction.setter + def direction(self, arg0: webrtc.enums.TransceiverDirection | webrtc.enums.TransceiverDirectionValue) -> None: + ... uri: str def __init__(self) -> None: ... @@ -159,7 +198,12 @@ class RtpEncodingParameters: active: bool adaptivePtime: bool codec: RtpCodec | None - networkPriority: webrtc.enums.RTCPriorityType + @property + def networkPriority(self) -> webrtc.enums.RTCPriorityType: + ... + @networkPriority.setter + def networkPriority(self, arg0: webrtc.enums.RTCPriorityType | webrtc.enums.RTCPriorityTypeValue) -> None: + ... requestKeyFrame: bool rid: str scalabilityMode: str | None @@ -196,7 +240,12 @@ class RtpEncodingParameters: def ssrc(self, arg0: typing.SupportsInt | typing.SupportsIndex | None) -> None: ... class RtpParameters: - degradationPreference: webrtc.enums.RTCDegradationPreference | None + @property + def degradationPreference(self) -> webrtc.enums.RTCDegradationPreference | None: + ... + @degradationPreference.setter + def degradationPreference(self, arg0: webrtc.enums.RTCDegradationPreference | webrtc.enums.RTCDegradationPreferenceValue | None) -> None: + ... mid: str rtcp: RtcpParameters transactionId: str @@ -236,7 +285,12 @@ class RtpCapabilities: def headerExtensions(self, arg0: collections.abc.Sequence[RtpHeaderExtensionCapability]) -> None: ... class RtpTransceiverInit: - direction: webrtc.enums.TransceiverDirection + @property + def direction(self) -> webrtc.enums.TransceiverDirection: + ... + @direction.setter + def direction(self, arg0: webrtc.enums.TransceiverDirection | webrtc.enums.TransceiverDirectionValue) -> None: + ... def __init__(self) -> None: ... @property @@ -261,15 +315,15 @@ class PeerConnectionFactory: def __init__(self) -> None: ... class MediaStreamTrack: - _constraints: typing.Any - _listeners: typing.Any + _constraints: webrtc.models.media_track_constraints.MediaTrackConstraints | None + _listeners: webrtc.utils.events._Listeners | None contentHint: str enabled: bool def _camera(self) -> tuple[int, int, float] | None: ... def _reconfigureCamera(self, width: typing.SupportsInt | typing.SupportsIndex, height: typing.SupportsInt | typing.SupportsIndex, frameRate: typing.SupportsFloat | typing.SupportsIndex) -> bool: ... - def _settings(self) -> dict: + def _settings(self) -> _TrackSettings: ... def _surfaceEnded(self) -> None: ... @@ -298,7 +352,7 @@ class MediaStreamTrack: def readyState(self) -> webrtc.enums.MediaStreamTrackState: ... class MediaStream: - _listeners: typing.Any + _listeners: webrtc.utils.events._Listeners | None @staticmethod def create(tracks: collections.abc.Sequence[MediaStreamTrack]) -> MediaStream: ... @@ -323,7 +377,7 @@ class MediaStream: def id(self) -> str: ... class RTCIceTransport: - _listeners: typing.Any + _listeners: webrtc.utils.events._Listeners | None def __init__(self) -> None: ... def _surfaceCandidate(self) -> None: @@ -334,7 +388,7 @@ class RTCIceTransport: ... def addRemoteCandidate(self, candidate: str, sdpMid: str, sdpMLineIndex: typing.SupportsInt | typing.SupportsIndex, usernameFragment: str | None) -> None: ... - def gather(self, policy: webrtc.enums.RTCIceTransportPolicy, iceServers: collections.abc.Sequence[IceServerInit]) -> None: + def gather(self, policy: webrtc.enums.RTCIceTransportPolicy | webrtc.enums.RTCIceTransportPolicyValue, iceServers: collections.abc.Sequence[IceServerInit]) -> None: ... def getLocalCandidates(self) -> list[IceCandidateInit]: ... @@ -366,10 +420,10 @@ class RTCIceTransport: def state(self) -> webrtc.enums.RTCIceTransportState: ... class RTCDtlsTransport: - _listeners: typing.Any + _listeners: webrtc.utils.events._Listeners | None def _surfaceState(self, state: webrtc.enums.DtlsTransportState) -> None: ... - def getRemoteCertificates(self) -> list: + def getRemoteCertificates(self) -> list[bytes]: ... @property def iceTransport(self) -> RTCIceTransport: @@ -378,7 +432,7 @@ class RTCDtlsTransport: def state(self) -> webrtc.enums.DtlsTransportState: ... class RTCSctpTransport: - _listeners: typing.Any + _listeners: webrtc.utils.events._Listeners | None def _surfaceState(self, state: webrtc.enums.SctpTransportState) -> None: ... @property @@ -394,7 +448,7 @@ class RTCSctpTransport: def transport(self) -> RTCDtlsTransport: ... class RTCDTMFSender: - _listeners: typing.Any + _listeners: webrtc.utils.events._Listeners | None def _surfaceBuffer(self, buffer: str, insertion: typing.SupportsInt | typing.SupportsIndex) -> None: ... def insertDTMF(self, tones: str, duration: typing.SupportsInt | typing.SupportsIndex, interToneGap: typing.SupportsInt | typing.SupportsIndex) -> None: @@ -462,7 +516,12 @@ class RTCRtpReceiver: def transport(self) -> RTCDtlsTransport | None: ... class RTCRtpTransceiver: - direction: webrtc.enums.TransceiverDirection + @property + def direction(self) -> webrtc.enums.TransceiverDirection: + ... + @direction.setter + def direction(self, arg0: webrtc.enums.TransceiverDirection | webrtc.enums.TransceiverDirectionValue) -> None: + ... def getCodecPreferences(self) -> list[RtpCodecCapability]: ... def getHeaderExtensionsToNegotiate(self) -> list[RtpHeaderExtensionCapability]: @@ -498,10 +557,10 @@ class RTCRtpTransceiver: ... class DataChannelMessage: @property - def data(self) -> typing.Any: + def data(self) -> str | bytes: ... class RTCDataChannel: - _listeners: typing.Any + _listeners: webrtc.utils.events._Listeners | None binaryType: str def _decreaseBufferedAmount(self, sent: typing.SupportsInt | typing.SupportsIndex) -> bool: ... @@ -511,7 +570,7 @@ class RTCDataChannel: ... def close(self) -> None: ... - def send(self, data: str, binary: bool) -> None: + def send(self, data: str | bytes, binary: bool) -> None: ... @property def bufferedAmount(self) -> int: @@ -550,7 +609,7 @@ class RTCDataChannel: def readyState(self) -> webrtc.enums.RTCDataChannelState: ... class RTCPeerConnection: - _listeners: typing.Any + _listeners: webrtc.utils.events._Listeners | None @staticmethod def _connectionOf(sender: RTCRtpSender) -> RTCPeerConnection | None: ... @@ -579,7 +638,7 @@ class RTCPeerConnection: def addTrack(self, track: MediaStreamTrack, streams: collections.abc.Sequence[MediaStream]) -> RTCRtpSender: ... @typing.overload - def addTransceiver(self, kind: webrtc.enums.MediaType, init: RtpTransceiverInit | None) -> RTCRtpTransceiver: + def addTransceiver(self, kind: webrtc.enums.MediaType | webrtc.enums.MediaTypeValue, init: RtpTransceiverInit | None) -> RTCRtpTransceiver: ... @typing.overload def addTransceiver(self, track: MediaStreamTrack, init: RtpTransceiverInit | None) -> RTCRtpTransceiver: @@ -588,7 +647,7 @@ class RTCPeerConnection: ... def createAnswer(self, onSuccess: collections.abc.Callable[[RTCSessionDescription], None], onFailure: collections.abc.Callable[[RTCCallbackException], None], voiceActivityDetection: bool) -> None: ... - def createDataChannel(self, label: str, ordered: bool, maxPacketLifeTime: typing.SupportsInt | typing.SupportsIndex | None, maxRetransmits: typing.SupportsInt | typing.SupportsIndex | None, protocol: str, negotiated: bool, id: typing.SupportsInt | typing.SupportsIndex | None, priority: webrtc.enums.RTCPriorityType) -> RTCDataChannel: + def createDataChannel(self, label: str, ordered: bool, maxPacketLifeTime: typing.SupportsInt | typing.SupportsIndex | None, maxRetransmits: typing.SupportsInt | typing.SupportsIndex | None, protocol: str, negotiated: bool, id: typing.SupportsInt | typing.SupportsIndex | None, priority: webrtc.enums.RTCPriorityType | webrtc.enums.RTCPriorityTypeValue) -> RTCDataChannel: ... def createOffer(self, onSuccess: collections.abc.Callable[[RTCSessionDescription], None], onFailure: collections.abc.Callable[[RTCCallbackException], None], iceRestart: bool, voiceActivityDetection: bool) -> None: ... @@ -650,11 +709,11 @@ class RTCPeerConnection: ... class VideoFrameBuffer: @staticmethod - def fromData(format: str, width: typing.SupportsInt | typing.SupportsIndex, height: typing.SupportsInt | typing.SupportsIndex, data: collections.abc.Buffer, layout: collections.abc.Sequence[tuple[typing.SupportsInt | typing.SupportsIndex, typing.SupportsInt | typing.SupportsIndex]]) -> VideoFrameBuffer: + def fromData(format: str, width: typing.SupportsInt | typing.SupportsIndex, height: typing.SupportsInt | typing.SupportsIndex, data: typing_extensions.Buffer, layout: collections.abc.Sequence[tuple[typing.SupportsInt | typing.SupportsIndex, typing.SupportsInt | typing.SupportsIndex]]) -> VideoFrameBuffer: ... - def convertTo(self, destination: collections.abc.Buffer, format: str, x: typing.SupportsInt | typing.SupportsIndex, y: typing.SupportsInt | typing.SupportsIndex, width: typing.SupportsInt | typing.SupportsIndex, height: typing.SupportsInt | typing.SupportsIndex, offset: typing.SupportsInt | typing.SupportsIndex, stride: typing.SupportsInt | typing.SupportsIndex, matrix: str, fullRange: bool) -> None: + def convertTo(self, destination: typing_extensions.Buffer, format: str, x: typing.SupportsInt | typing.SupportsIndex, y: typing.SupportsInt | typing.SupportsIndex, width: typing.SupportsInt | typing.SupportsIndex, height: typing.SupportsInt | typing.SupportsIndex, offset: typing.SupportsInt | typing.SupportsIndex, stride: typing.SupportsInt | typing.SupportsIndex, matrix: str, fullRange: bool) -> None: ... - def copyPlanes(self, destination: collections.abc.Buffer, planes: collections.abc.Sequence[tuple[typing.SupportsInt | typing.SupportsIndex, typing.SupportsInt | typing.SupportsIndex, typing.SupportsInt | typing.SupportsIndex, typing.SupportsInt | typing.SupportsIndex, typing.SupportsInt | typing.SupportsIndex, typing.SupportsInt | typing.SupportsIndex]]) -> None: + def copyPlanes(self, destination: typing_extensions.Buffer, planes: collections.abc.Sequence[tuple[typing.SupportsInt | typing.SupportsIndex, typing.SupportsInt | typing.SupportsIndex, typing.SupportsInt | typing.SupportsIndex, typing.SupportsInt | typing.SupportsIndex, typing.SupportsInt | typing.SupportsIndex, typing.SupportsInt | typing.SupportsIndex]]) -> None: ... def withoutAlpha(self) -> VideoFrameBuffer: ... @@ -668,14 +727,14 @@ class VideoFrameBuffer: def width(self) -> int: ... class MediaStreamTrackProcessor: - _listeners: typing.Any + _listeners: webrtc.utils.events._Listeners | None def __init__(self, track: MediaStreamTrack, maxBufferSize: typing.SupportsInt | typing.SupportsIndex) -> None: ... def _ackWakeup(self) -> None: ... def cancel(self) -> None: ... - def read(self) -> typing.Any: + def read(self) -> tuple[VideoFrameBuffer, int, int, int] | tuple[bytes, int, int, int, int, int] | None: ... @property def discardedFrames(self) -> int: @@ -709,10 +768,25 @@ def _alive() -> dict[str, int]: ... def _alive_factories() -> int: ... -def copyAudioSamples(source: collections.abc.Buffer, sourceFormat: str, channels: typing.SupportsInt | typing.SupportsIndex, frames: typing.SupportsInt | typing.SupportsIndex, destination: collections.abc.Buffer, destinationFormat: str, planeIndex: typing.SupportsInt | typing.SupportsIndex, frameOffset: typing.SupportsInt | typing.SupportsIndex, frameCount: typing.SupportsInt | typing.SupportsIndex) -> None: +def copyAudioSamples(source: typing_extensions.Buffer, sourceFormat: str, channels: typing.SupportsInt | typing.SupportsIndex, frames: typing.SupportsInt | typing.SupportsIndex, destination: typing_extensions.Buffer, destinationFormat: str, planeIndex: typing.SupportsInt | typing.SupportsIndex, frameOffset: typing.SupportsInt | typing.SupportsIndex, frameCount: typing.SupportsInt | typing.SupportsIndex) -> None: ... def getUserMedia(audio: bool, video: bool, width: typing.SupportsInt | typing.SupportsIndex, height: typing.SupportsInt | typing.SupportsIndex, frameRate: typing.SupportsFloat | typing.SupportsIndex) -> MediaStream: ... def ping() -> None: ... _sanitized: bool = False +class _IceCandidateKwargs(typing_extensions.TypedDict, closed=True): + candidate: str + sdp_mid: str + sdp_m_line_index: int + username_fragment: str | None + url: str | None + relay_protocol: typing.Literal['udp', 'tcp', 'tls'] | None +class _TrackSettings(typing.TypedDict, total=False): + width: int + height: int + frame_rate: float + device: typing.Literal['camera', 'microphone'] + sample_rate: int + channel_count: int + sample_size: int diff --git a/tests/chaos.py b/tests/chaos.py index 3205ec2..1b53131 100644 --- a/tests/chaos.py +++ b/tests/chaos.py @@ -29,7 +29,7 @@ from tests.helpers import connect if TYPE_CHECKING: - from collections.abc import Callable + from collections.abc import Callable, Coroutine log = logging.getLogger('chaos') @@ -50,16 +50,21 @@ def __init__(self, seed: int) -> None: self.connections: list[webrtc.RTCPeerConnection] = [] self.channels: list[webrtc.RTCDataChannel] = [] self.tracks: list[webrtc.MediaStreamTrack] = [] - self.processors: list[tuple[webrtc.MediaStreamTrackProcessor, webrtc.ReadableStreamDefaultReader]] = [] + self.processors: list[ + tuple[ + webrtc.MediaStreamTrackProcessor, + webrtc.ReadableStreamDefaultReader[webrtc.VideoFrame | webrtc.AudioData], + ] + ] = [] 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 + return self.random.choice(pool) if len(pool) > 0 else None - def drop(self, pool: list[object]) -> None: - if pool: + def drop(self, pool: list[T]) -> None: + if len(pool) > 0: pool.pop(self.random.randrange(len(pool))) def handler(self) -> Callable[[webrtc.Event], str | None]: @@ -69,7 +74,7 @@ def handler(self) -> Callable[[webrtc.Event], str | None]: def handle(_event: webrtc.Event) -> str | None: if action == 0 and target is not None: - target.close() if hasattr(target, 'close') else target.stop() + target.stop() if isinstance(target, webrtc.MediaStreamTrack) else target.close() elif action == 1: msg = 'a handler raises' raise RuntimeError(msg) @@ -94,7 +99,7 @@ async def new_connection(self) -> None: async def close_connection(self) -> None: pc = self.pick(self.connections) - if pc: + if pc is not None: pc.close() async def drop_connection(self) -> None: @@ -107,17 +112,17 @@ async def connect_two(self) -> None: async def add_track(self) -> None: pc, track = self.pick(self.connections), self.pick(self.tracks) - if pc and track: + if pc is not None and track is not None: pc.add_track(track) async def remove_track(self) -> None: pc = self.pick(self.connections) - if pc and pc.get_senders(): + if pc is not None and len(pc.get_senders()) > 0: pc.remove_track(self.random.choice(pc.get_senders())) async def add_transceiver(self) -> None: pc = self.pick(self.connections) - if pc: + if pc is not None: transceiver = pc.add_transceiver(self.random.choice(['audio', 'video'])) if self.random.random() < 0.3: transceiver.stop() @@ -126,49 +131,49 @@ async def add_transceiver(self) -> None: async def negotiate(self) -> None: pc = self.pick(self.connections) - if pc: + if pc is not None: await pc.set_local_description() async def create_channel(self) -> None: pc = self.pick(self.connections) - if pc: + if pc is not None: 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) -> None: channel = self.pick(self.channels) - if channel: + if channel is not None: channel.send(self.random.choice(['text', b'\x00' * self.random.randrange(70000), bytearray(10)])) async def close_channel(self) -> None: channel = self.pick(self.channels) - if channel: + if channel is not None: channel.close() async def stats(self) -> None: pc = self.pick(self.connections) - if pc: + if pc is not None: await pc.get_stats() async def restart_ice(self) -> None: pc = self.pick(self.connections) - if pc: + if pc is not None: pc.restart_ice() async def handle_connection_event(self) -> None: pc = self.pick(self.connections) - if pc: + if pc is not None: pc.on(self.random.choice(['connectionstatechange', 'icecandidate', 'track', 'datachannel']), self.handler()) async def replace_track(self) -> None: pc = self.pick(self.connections) - if pc and pc.get_senders(): + if pc is not None and len(pc.get_senders()) > 0: 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(): + if pc is not None and len(pc.get_senders()) > 0: sender = self.random.choice(pc.get_senders()) parameters = sender.get_parameters() for encoding in parameters.encodings: @@ -185,17 +190,17 @@ async def get_user_media(self) -> None: async def stop_track(self) -> None: track = self.pick(self.tracks) - if track: + if track is not None: track.stop() async def clone_track(self) -> None: track = self.pick(self.tracks) - if track: + if track is not None: self.tracks.append(track.clone()) async def toggle_track(self) -> None: track = self.pick(self.tracks) - if track: + if track is not None: track.enabled = not track.enabled async def drop_track(self) -> None: @@ -203,7 +208,7 @@ async def drop_track(self) -> None: async def new_processor(self) -> None: track = self.pick(self.tracks) - if track: + if track is not None: processor = webrtc.MediaStreamTrackProcessor( webrtc.MediaStreamTrackProcessorInit(track, max_buffer_size=self.random.randrange(4)) ) @@ -211,18 +216,20 @@ async def new_processor(self) -> None: self.processors.append((processor, processor.readable.get_reader())) async def read(self) -> None: - if self.processors: - _, reader = self.pick(self.processors) + if len(self.processors) > 0: + _, reader = self.random.choice(self.processors) try: result = await asyncio.wait_for(reader.read(), 0.2) except asyncio.TimeoutError: return if not result.done: + # the processor's readable is typed ReadableStream[object] + assert isinstance(result.value, (webrtc.VideoFrame, webrtc.AudioData)) result.value.close() async def cancel_processor(self) -> None: - if self.processors: - _, reader = self.pick(self.processors) + if len(self.processors) > 0: + _, reader = self.random.choice(self.processors) await reader.cancel() async def drop_processor(self) -> None: @@ -239,9 +246,9 @@ async def new_generator(self) -> None: self.tracks.append(generator) async def write(self) -> None: - if not self.generators: + if len(self.generators) == 0: return - writer, kind = self.pick(self.generators) + writer, kind = self.random.choice(self.generators) if kind == 'video': width, height = self.random.choice([(2, 2), (33, 17), (320, 240)]) chunk = webrtc.VideoFrame( @@ -264,8 +271,8 @@ async def write(self) -> None: await writer.write(chunk) async def close_generator(self) -> None: - if self.generators: - writer, _ = self.pick(self.generators) + if len(self.generators) > 0: + writer, _ = self.random.choice(self.generators) await writer.close() async def drop_generator(self) -> None: @@ -284,12 +291,12 @@ async def frame(self) -> None: async def use_frame(self) -> None: frame = self.pick(self.frames) - if frame: + if frame is not None: self.random.choice([frame.close, lambda: self.frames.append(frame.clone())])() async def pipe(self) -> None: - track = self.pick([track for track in self.tracks if track.kind == 'video']) - if track: + track = self.pick([video for video in self.tracks if video.kind == 'video']) + if track is not None: processor = webrtc.MediaStreamTrackProcessor(webrtc.MediaStreamTrackProcessorInit(track)) generator = webrtc.VideoTrackGenerator() self.tracks.append(generator.track) @@ -297,7 +304,7 @@ async def pipe(self) -> None: async def constraints(self) -> None: track = self.pick(self.tracks) - if track: + if track is not None: track.get_settings() track.apply_constraints( self.random.choice([ @@ -310,7 +317,7 @@ async def constraints(self) -> None: 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: + if len(tracks) > 0 and self.random.random() < 0.5: stream.remove_track(tracks[0]) stream.get_tracks() @@ -366,7 +373,8 @@ 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) + run: Callable[[], Coroutine[object, object, None]] = getattr(self, name) + await asyncio.wait_for(run(), 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: @@ -387,10 +395,10 @@ def main() -> None: asyncio.run(Chaos(args.seed).run(args.steps)) # the last references may be released on helper threads deadline = time.monotonic() + 1 - while wrtc._alive_factories() and time.monotonic() < deadline: + while wrtc._alive_factories() != 0 and time.monotonic() < deadline: gc.collect() time.sleep(0.05) - alive = {name: count for name, count in wrtc._alive().items() if count} + alive = {name: count for name, count in wrtc._alive().items() if count != 0} log.info('done, %d factories alive, native objects alive: %s', wrtc._alive_factories(), alive) diff --git a/tests/fuzz/fuzz_audio_data.py b/tests/fuzz/fuzz_audio_data.py index c614311..9313b15 100644 --- a/tests/fuzz/fuzz_audio_data.py +++ b/tests/fuzz/fuzz_audio_data.py @@ -9,6 +9,7 @@ from __future__ import annotations +import pathlib import sys import atheris @@ -18,6 +19,9 @@ import webrtc +sys.path.insert(0, str(pathlib.Path(__file__).parent.parent.parent)) +from tests.helpers import mistyped + FORMATS = list(webrtc.AudioSampleFormat) EXPECTED = (TypeError, ValueError, BufferError, webrtc.NotSupportedError, webrtc.InvalidStateError) SAMPLE_BYTES = {'u8': 1, 's16': 2, 's32': 4, 'f32': 4} @@ -25,6 +29,7 @@ def check_identity(audio: webrtc.AudioData, data: bytes) -> None: """Interleaved samples copied out in their own format are the same bytes.""" + assert audio.format is not None if audio.format.value.endswith('-planar'): return out = bytearray(audio.allocation_size(webrtc.AudioDataCopyToOptions(plane_index=0))) @@ -48,9 +53,9 @@ def test_one_input(data: bytes) -> None: webrtc.AudioDataInit( format=format, sample_rate=sample_rate, - number_of_frames=frames, - number_of_channels=channels, - timestamp=inp.integer(), + number_of_frames=mistyped(frames), + number_of_channels=mistyped(channels), + timestamp=mistyped(inp.integer()), data=source, ) ) @@ -62,16 +67,16 @@ def test_one_input(data: bytes) -> None: def exercise(inp: Input, audio: webrtc.AudioData) -> None: for _ in range(inp.small(4)): - options = webrtc.AudioDataCopyToOptions(plane_index=inp.integer(8)) + options = webrtc.AudioDataCopyToOptions(plane_index=mistyped(inp.integer(8))) if inp.flag(): - options.frame_offset = inp.integer(512) + options.frame_offset = mistyped(inp.integer(512)) if inp.flag(): - options.frame_count = inp.integer(512) + options.frame_count = mistyped(inp.integer(512)) if inp.flag(): options.format = inp.choice(FORMATS) try: size = audio.allocation_size(options) - audio.copy_to(inp.destination(min(size, 1 << 20)), options) + audio.copy_to(mistyped(inp.destination(min(size, 1 << 20))), options) except EXPECTED: pass if inp.small(8) == 0: diff --git a/tests/fuzz/fuzz_generator.py b/tests/fuzz/fuzz_generator.py index 032818a..f604202 100644 --- a/tests/fuzz/fuzz_generator.py +++ b/tests/fuzz/fuzz_generator.py @@ -25,7 +25,7 @@ import webrtc sys.path.insert(0, str(pathlib.Path(__file__).parent.parent.parent)) -from tests.helpers import connect +from tests.helpers import connect, mistyped EXPECTED = (TypeError, ValueError, BufferError, webrtc.NotSupportedError, webrtc.InvalidStateError) PIXEL_FORMATS = list(webrtc.VideoPixelFormat) @@ -40,6 +40,12 @@ class Session: """A connected pair sending a generator of each kind, whose writers are replaced once they fail.""" + caller: webrtc.RTCPeerConnection + callee: webrtc.RTCPeerConnection + processors: list[webrtc.MediaStreamTrackProcessor] + senders: dict[webrtc.MediaTypeValue, webrtc.RTCRtpSender] + writers: dict[webrtc.MediaTypeValue, webrtc.WritableStreamDefaultWriter[object]] + async def start(self) -> None: self.caller, self.callee = webrtc.RTCPeerConnection(), webrtc.RTCPeerConnection() self.processors = [] @@ -49,14 +55,15 @@ async def start(self) -> None: webrtc.MediaStreamTrackProcessor(webrtc.MediaStreamTrackProcessorInit(event.track)) ), ) - self.senders, self.writers = {}, {} + self.senders = {} + self.writers = {} for kind in ('audio', 'video'): generator = webrtc.MediaStreamTrackGenerator(kind) self.senders[kind] = self.caller.add_track(generator) self.writers[kind] = generator.writable.get_writer() await connect(self.caller, self.callee) - async def write(self, kind: str, chunk: webrtc.AudioData | webrtc.VideoFrame) -> None: + async def write(self, kind: webrtc.MediaTypeValue, chunk: webrtc.AudioData | webrtc.VideoFrame) -> None: try: await self.writers[kind].write(chunk) except EXPECTED: @@ -73,7 +80,14 @@ def audio_data(inp: Input) -> webrtc.AudioData: 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 v > 0 for v in (channels, frames)): + if not ( + isinstance(channels, int) + and not isinstance(channels, bool) + and channels > 0 + and isinstance(frames, int) + and not isinstance(frames, bool) + and frames > 0 + ): channels, frames = 1, 480 size = min(frames * channels * SAMPLE_BYTES[format.value.split('-')[0]], 1 << 20) return webrtc.AudioData( @@ -82,7 +96,7 @@ def audio_data(inp: Input) -> webrtc.AudioData: sample_rate=rate, number_of_frames=frames, number_of_channels=channels, - timestamp=inp.integer(), + timestamp=mistyped(inp.integer()), data=bytes(size), ) ) @@ -92,7 +106,9 @@ 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]) - init = webrtc.VideoFrameBufferInit(format=format, coded_width=width, coded_height=height, timestamp=inp.integer()) + init = webrtc.VideoFrameBufferInit( + format=format, coded_width=width, coded_height=height, timestamp=mistyped(inp.integer()) + ) if inp.flag(): init.rotation = inp.choice([0, 90, 180, 270]) # enough for every format: 4 planes of 16-bit samples at most diff --git a/tests/fuzz/fuzz_native_buffers.py b/tests/fuzz/fuzz_native_buffers.py index d6df1bd..1de3f47 100644 --- a/tests/fuzz/fuzz_native_buffers.py +++ b/tests/fuzz/fuzz_native_buffers.py @@ -46,7 +46,8 @@ def video(inp: Input) -> None: for _ in range(inp.small(4)): action = inp.small(2) if action == 0: - copies = [tuple(inp.unsigned(256) for _ in range(6)) for _ in range(inp.small(4))] + u = inp.unsigned + copies = [(u(256), u(256), u(256), u(256), u(256), u(256)) for _ in range(inp.small(4))] buffer.copyPlanes(inp.destination(inp.small(1 << 15)), copies) elif action == 1: buffer.convertTo( @@ -81,7 +82,7 @@ def audio(inp: Input) -> None: def test_one_input(data: bytes) -> None: inp = Input(data) - with contextlib.suppress(EXPECTED): + with contextlib.suppress(*EXPECTED): video(inp) if inp.flag() else audio(inp) diff --git a/tests/fuzz/fuzz_video_frame.py b/tests/fuzz/fuzz_video_frame.py index 673d498..70fb272 100644 --- a/tests/fuzz/fuzz_video_frame.py +++ b/tests/fuzz/fuzz_video_frame.py @@ -10,6 +10,7 @@ from __future__ import annotations import asyncio +import pathlib import sys import atheris @@ -19,6 +20,9 @@ import webrtc +sys.path.insert(0, str(pathlib.Path(__file__).parent.parent.parent)) +from tests.helpers import mistyped + FORMATS = list(webrtc.VideoPixelFormat) # 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) @@ -29,7 +33,7 @@ async def _copy_to( frame: webrtc.VideoFrame, destination: Buffer, options: webrtc.VideoFrameCopyToOptions | None ) -> None: - await frame.copy_to(destination, options) + await frame.copy_to(mistyped(destination), options) def copy_to(frame: webrtc.VideoFrame, destination: Buffer, options: webrtc.VideoFrameCopyToOptions | None) -> None: @@ -41,7 +45,7 @@ def rect(inp: Input) -> webrtc.DOMRectInit: def layout(inp: Input) -> list[webrtc.PlaneLayout]: - return [webrtc.PlaneLayout(inp.integer(4096), inp.integer(256)) for _ in range(inp.small(4))] + return [webrtc.PlaneLayout(mistyped(inp.integer(4096)), mistyped(inp.integer(256))) for _ in range(inp.small(4))] def copy_options(inp: Input) -> webrtc.VideoFrameCopyToOptions: @@ -63,7 +67,7 @@ def frame_of_frame(inp: Input, frame: webrtc.VideoFrame) -> webrtc.VideoFrame: init.rotation = inp.number(360) init.flip = inp.flag() if inp.flag(): - init.display_width, init.display_height = inp.integer(), inp.integer() + init.display_width, init.display_height = mistyped(inp.integer()), mistyped(inp.integer()) return webrtc.VideoFrame(frame, init) @@ -88,12 +92,12 @@ def check_identity(format: webrtc.VideoPixelFormat, size: tuple[int, int], data: """A packed frame copied out as it is gives the same bytes.""" width, height = size init = webrtc.VideoFrameBufferInit(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: + allocation = webrtc.VideoFrame(bytes(1 << 16), init).allocation_size() if width * height <= 1024 else 0 + if allocation == 0: return - packed = bytes(data)[:size] + bytes(max(0, size - len(data))) + packed = bytes(data)[:allocation] + bytes(max(0, allocation - len(data))) with webrtc.VideoFrame(packed, init) as frame: - out = bytearray(size) + out = bytearray(allocation) copy_to(frame, out, None) assert bytes(out) == packed, f'{format} {width}x{height} copied out differently' @@ -106,7 +110,9 @@ def test_one_input(data: bytes) -> None: check_identity(format, (width, height), inp.buffer(width * height * 8)) return width, height = inp.integer(), inp.integer() - init = webrtc.VideoFrameBufferInit(format=format, coded_width=width, coded_height=height, timestamp=inp.integer()) + init = webrtc.VideoFrameBufferInit( + format=format, coded_width=mistyped(width), coded_height=mistyped(height), timestamp=mistyped(inp.integer()) + ) if inp.flag(): init.layout = layout(inp) if inp.flag(): @@ -114,7 +120,7 @@ def test_one_input(data: bytes) -> None: if inp.flag(): init.rotation = inp.number(360) if inp.flag(): - init.display_width, init.display_height = inp.integer(), inp.integer() + init.display_width, init.display_height = mistyped(inp.integer()), mistyped(inp.integer()) try: frame = webrtc.VideoFrame(inp.buffer(inp.small(1 << 15)), init) exercise(inp, frame) diff --git a/tests/fuzz/inputs.py b/tests/fuzz/inputs.py index 07d7001..28e5b85 100644 --- a/tests/fuzz/inputs.py +++ b/tests/fuzz/inputs.py @@ -28,7 +28,7 @@ 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()), + lambda data: memoryview(bytearray(data)).cast('B', [len(data), 1]) if len(data) > 0 else memoryview(bytearray()), ] diff --git a/tests/helpers.py b/tests/helpers.py index a5595a0..b0b869e 100644 --- a/tests/helpers.py +++ b/tests/helpers.py @@ -16,7 +16,7 @@ import subprocess import sys import textwrap -from typing import TYPE_CHECKING, Callable +from typing import TYPE_CHECKING, Callable, TypeVar, cast import pytest @@ -24,11 +24,18 @@ import wrtc if TYPE_CHECKING: - from collections.abc import AsyncIterator, Awaitable + from collections.abc import AsyncGenerator, Awaitable #: The fixture creating connections with a configuration CreatePC = Callable[..., webrtc.RTCPeerConnection] +_T = TypeVar('_T') + + +def mistyped(value: object) -> _T: + """A value of the wrong type, passed where a test checks that the library rejects it at runtime.""" + return cast('_T', value) + async def exchange_offer(caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection) -> None: offer = await caller.create_offer() @@ -75,7 +82,7 @@ async def wait_until(predicate: Callable[[], object], what: str, timeout: float """ loop = asyncio.get_running_loop() deadline = loop.time() + timeout - while not await _called(predicate): + while not bool(await _called(predicate)): if loop.time() > deadline: msg = f'Timed out waiting for {what}' raise TimeoutError(msg) @@ -210,7 +217,7 @@ async def write_video( @contextlib.asynccontextmanager -async def writing(write: Callable[..., Awaitable[None]], *args: object, **kwargs: object) -> AsyncIterator[None]: +async def writing(write: Callable[..., Awaitable[None]], *args: object, **kwargs: object) -> AsyncGenerator[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)) @@ -252,7 +259,7 @@ class Counters(ctypes.Structure): counters.cb = ctypes.sizeof(counters) process = ctypes.windll.kernel32.GetCurrentProcess() ctypes.windll.psapi.GetProcessMemoryInfo(process, ctypes.byref(counters), counters.cb) - return counters.WorkingSetSize + return int(counters.WorkingSetSize) # macOS and other BSDs return int(subprocess.check_output(['/bin/ps', '-o', 'rss=', '-p', str(os.getpid())])) * 1024 diff --git a/tests/idl/child.py b/tests/idl/child.py index 5d3f17e..3c7a918 100644 --- a/tests/idl/child.py +++ b/tests/idl/child.py @@ -23,9 +23,13 @@ def main() -> None: pm.eval((WPT_ROOT / 'resources' / 'webidl2' / 'lib' / 'webidl2.js').read_text()) parse = pm.eval('(text) => JSON.stringify(globalThis.WebIDL2.parse(text))') - definitions = [] + definitions: list[dict[str, object]] = [] for name, text in json.load(sys.stdin).items(): - for definition in json.loads(parse(text)): + ast: object = parse(text) + if not isinstance(ast, str): + msg = f'JSON.stringify returned {ast!r}' + raise TypeError(msg) + for definition in json.loads(ast): definition['file'] = name definitions.append(definition) json.dump(definitions, sys.stdout) diff --git a/tests/idl/compare.py b/tests/idl/compare.py index 2af220c..2206c5e 100644 --- a/tests/idl/compare.py +++ b/tests/idl/compare.py @@ -21,12 +21,24 @@ import inspect import re import textwrap -from typing import TYPE_CHECKING, Any +from typing import TYPE_CHECKING -if TYPE_CHECKING: - from collections.abc import Callable +from typing_extensions import TypeGuard - from tests.idl.spec import Definition, Node, Spec +if TYPE_CHECKING: + from collections.abc import Callable, Sequence + + from tests.idl.spec import ( + Argument, + Attribute, + Constructor, + Definition, + Field, + IdlType, + Member, + Operation, + Spec, + ) _BOUNDARY = re.compile(r'(?<=[a-z0-9])(?=[A-Z])|(?<=[A-Z])(?=[A-Z][a-z])') _AWAITABLE = re.compile(r'\b(Future|Awaitable|Coroutine|Task)\b') @@ -78,19 +90,25 @@ def _text(annotation: object) -> str | None: return annotation if isinstance(annotation, str) else getattr(annotation, '__name__', str(annotation)) -def _is_async(function: Callable[..., Any]) -> bool: +def _mentions(pattern: str | re.Pattern[str], annotation: object) -> bool: + """Whether the text of an annotation matches a pattern.""" + text = _text(annotation) + return re.search(pattern, text if text is not None else '') is not None + + +def _is_async(function: Callable[..., object]) -> bool: """Whether a function is a coroutine function or returns an awaitable, like a future.""" if inspect.iscoroutinefunction(function): return True - return _AWAITABLE.search(_text(inspect.signature(function).return_annotation) or '') is not None + return _mentions(_AWAITABLE, inspect.signature(function).return_annotation) -def _type_names(idl_type: Node) -> str: +def _type_names(idl_type: IdlType) -> str: inner = idl_type['idlType'] return inner if isinstance(inner, str) else ' '.join(_type_names(item) for item in inner) -def _is_handler(member: Node) -> bool: +def _is_handler(member: Member) -> TypeGuard[Attribute]: """Whether a member is an event handler attribute, which the package replaces with ``on(name)``.""" return ( member['type'] == 'attribute' @@ -111,7 +129,7 @@ class _Types: def __init__(self, spec: Spec) -> None: self.spec = spec - def lacks(self, label: str, idl_type: Node, annotation: object) -> list[str]: + def lacks(self, label: str, idl_type: IdlType, annotation: object) -> list[str]: """Checks that an annotation names the definitions the IDL type refers to.""" text = _text(annotation) if text is None: @@ -121,7 +139,7 @@ def named(name: str) -> bool: return re.search(rf'\b{name}\b', text) is not None missing = sorted(name for name in self.spec.named_types(idl_type, named) if not named(name)) - return [f'{label}: type lacks {", ".join(missing)}'] if missing else [] + return [f'{label}: type lacks {", ".join(missing)}'] if len(missing) > 0 else [] class _Arguments(_Types): @@ -135,20 +153,21 @@ def __init__(self, spec: Spec, label: str, parameters: list[inspect.Parameter]) self.order: list[int] = [] self.found: list[str] = [] - def check(self, arguments: list[Node]) -> list[str]: + def check(self, arguments: list[Argument]) -> list[str]: for index, argument in enumerate(arguments): self.argument(index, argument) if self.order != sorted(self.order): self.found.append(f'{self.label}: arguments out of order') for parameter in self.unused.values(): - prefix = {_P.VAR_POSITIONAL: '*', _P.VAR_KEYWORD: '**'}.get(parameter.kind, '') + prefix = '*' if parameter.kind is _P.VAR_POSITIONAL else '**' if parameter.kind is _P.VAR_KEYWORD else '' self.found.append(f'{self.label}({prefix}{parameter.name}): extra argument') return self.found def take(self, name: str) -> inspect.Parameter | None: - return self.unused.pop(snake_case(name), None) or self.unused.pop(name, None) + parameter = self.unused.pop(snake_case(name), None) + return parameter if parameter is not None else self.unused.pop(name, None) - def argument(self, index: int, argument: Node) -> None: + def argument(self, index: int, argument: Argument) -> None: path = f'{self.label}({argument["name"]})' parameter = self.take(argument['name']) if parameter is None: @@ -182,12 +201,12 @@ def positional(self, index: int) -> inspect.Parameter | None: return None return self.unused.pop(parameter.name) - def flattened(self, argument: Node) -> list[Node] | None: + def flattened(self, argument: Argument) -> list[Field] | None: """The members of a dictionary argument, if the signature takes them as parameters of their own.""" dictionary = self.spec.dictionary(argument['idlType']) if dictionary is None: return None - members = self.spec.members(dictionary.name) + members = self.spec.fields(dictionary.name) matched = [ self.unused[name] for member in members @@ -195,20 +214,18 @@ def flattened(self, argument: Node) -> list[Node] | None: ] # a lone positional parameter named like a member but annotated as a dictionary is the argument, renamed lone = len(matched) == 1 and matched[0].kind is not _P.KEYWORD_ONLY - if not matched or ( - lone and re.search(rf'\b(dict|Mapping|{dictionary.name})\b', _text(matched[0].annotation) or '') - ): + if len(matched) == 0 or (lone and _mentions(rf'\b(dict|Mapping|{dictionary.name})\b', matched[0].annotation)): return None return members - def members(self, prefix: str, members: list[Node]) -> None: + def members(self, prefix: str, members: list[Field]) -> None: for member in members: path = f'{prefix}.{member["name"]})' parameter = self.take(member['name']) if parameter is None: self.found.append(f'{path}: missing member') continue - has_default = parameter.default is not _P.empty + has_default: bool = parameter.default is not _P.empty self.found += _default(path, required=member['required'], has_default=has_default) self.found += self.lacks(path, member['idlType'], parameter.annotation) @@ -263,10 +280,10 @@ def check_enum(self) -> list[str]: ] def check_dictionary(self) -> list[str]: - fields = ( + fields: dict[str, dataclasses.Field[object]] = ( {field.name: field for field in dataclasses.fields(self.cls)} if dataclasses.is_dataclass(self.cls) else {} ) - for member in self.spec.members(self.definition.name): + for member in self.spec.fields(self.definition.name): name = self.resolve(member['name'], 'member') if name is None: continue @@ -281,33 +298,34 @@ def check_dictionary(self) -> list[str]: def check_interface(self) -> list[str]: members = self.spec.members(self.definition.name) - operations: dict[str, list[Node]] = {} + operations: dict[str, list[Operation]] = {} for member in members: - if member['type'] == 'operation' and member['name']: + if member['type'] == 'operation' and member['name'] != '': operations.setdefault(member['name'], []).append(member) else: self.check_member(member) self.check_events({member['name'][2:] for member in members if _is_handler(member)}) constructors = [member for member in members if member['type'] == 'constructor'] - if constructors: + if len(constructors) > 0: self.check_overloads('constructor', constructors, lambda: inspect.signature(self.cls)) for name, overloads in operations.items(): self.check_operation(name, overloads) return self.found + self.extras() - def check_member(self, member: Node) -> None: - kind, name = member['type'], member.get('name') + def check_member(self, member: Member) -> None: + kind = member['type'] if kind in _PROTOCOLS: self.expected |= set(_PROTOCOLS[kind]) self.found += [f'{kind}: missing {method}' for method in _PROTOCOLS[kind] if not hasattr(self.cls, method)] - elif kind == 'const': + elif member['type'] == 'const': + name = member['name'] self.expected.add(name) if name not in self.names: self.found.append(f'{name}: missing constant') elif _is_handler(member): - self.expected |= {name, snake_case(name)} - elif kind == 'attribute': + self.expected |= {member['name'], snake_case(member['name'])} + elif member['type'] == 'attribute': self.check_attribute(member) def check_events(self, events: set[str]) -> None: @@ -315,7 +333,7 @@ def check_events(self, events: set[str]) -> None: self.found += [f'on{event}: missing event' for event in events - actual] self.found += [f'on{event}: extra event' for event in actual - events] - def check_attribute(self, member: Node) -> None: + def check_attribute(self, member: Attribute) -> None: idl_name = member['name'] name = self.resolve(idl_name, 'attribute') if name is None: @@ -330,10 +348,11 @@ def check_attribute(self, member: Node) -> None: self.found.append(f'{idl_name}: should be read-only') elif not member['readonly'] and attribute.fset is None: self.found.append(f'{idl_name}: should be writable') + assert attribute.fget is not None # a property without a getter has no type annotation = inspect.get_annotations(attribute.fget).get('return', _P.empty) self.found += self.lacks(idl_name, member['idlType'], annotation) - def check_operation(self, idl_name: str, overloads: list[Node]) -> None: + def check_operation(self, idl_name: str, overloads: list[Operation]) -> None: name = self.resolve(idl_name, 'method') if name is None: return @@ -359,7 +378,9 @@ def signature() -> inspect.Signature: self.check_overloads(idl_name, overloads, signature) self.found += self.lacks(f'{idl_name}()', overloads[0]['idlType'], signature().return_annotation) - def check_overloads(self, label: str, overloads: list[Node], signature: Callable[[], inspect.Signature]) -> None: + def check_overloads( + self, label: str, overloads: Sequence[Operation | Constructor], signature: Callable[[], inspect.Signature] + ) -> None: """Checks the arguments against the overload the signature matches best.""" try: parameters = list(signature().parameters.values()) @@ -382,9 +403,9 @@ def _differences(spec: Spec, definition: Definition, module: object) -> list[str def compare(spec: Spec, module: object) -> dict[str, list[str]]: """The differences of every definition that has some, by the name of the definition.""" - differences = {} + differences: dict[str, list[str]] = {} for name, definition in spec.definitions.items(): found = sorted(set(_differences(spec, definition, module))) - if found: + if len(found) > 0: differences[name] = found return differences diff --git a/tests/idl/spec.py b/tests/idl/spec.py index 55cbfb7..825ec54 100644 --- a/tests/idl/spec.py +++ b/tests/idl/spec.py @@ -15,7 +15,9 @@ from dataclasses import dataclass, field from importlib.util import find_spec from pathlib import Path -from typing import TYPE_CHECKING, Any +from typing import TYPE_CHECKING, Literal, Union + +from typing_extensions import TypedDict, TypeGuard if TYPE_CHECKING: from collections.abc import Callable @@ -52,7 +54,112 @@ 'webcrypto.idl': set(), } -Node = dict[str, Any] + +# the JSON AST of webidl2.js, as far as the comparison reads it +class IdlType(TypedDict): + type: str | None + generic: str + nullable: bool + union: bool + idlType: str | list[IdlType] + + +class Argument(TypedDict): + type: Literal['argument'] + name: str + idlType: IdlType + optional: bool + variadic: bool + + +class Attribute(TypedDict): + type: Literal['attribute'] + name: str + idlType: IdlType + special: str + readonly: bool + + +class Operation(TypedDict): + type: Literal['operation'] + name: str + idlType: IdlType + arguments: list[Argument] + special: str + + +class Constructor(TypedDict): + type: Literal['constructor'] + arguments: list[Argument] + + +class Const(TypedDict): + type: Literal['const'] + name: str + idlType: IdlType + + +class Field(TypedDict): + type: Literal['field'] + name: str + idlType: IdlType + required: bool + + +class Declaration(TypedDict): + type: Literal['iterable', 'async_iterable', 'maplike', 'setlike'] + idlType: list[IdlType] + arguments: list[Argument] + readonly: bool + + +Member = Union[Attribute, Operation, Constructor, Const, Field, Declaration] + + +class Container(TypedDict): + type: Literal['interface', 'interface mixin', 'dictionary', 'namespace', 'callback interface'] + name: str + inheritance: str | None + members: list[Member] + partial: bool + file: str + + +class EnumValue(TypedDict): + type: Literal['enum-value'] + value: str + + +class Enum(TypedDict): + type: Literal['enum'] + name: str + values: list[EnumValue] + file: str + + +class Typedef(TypedDict): + type: Literal['typedef'] + name: str + idlType: IdlType + file: str + + +class Callback(TypedDict): + type: Literal['callback'] + name: str + idlType: IdlType + arguments: list[Argument] + file: str + + +class Includes(TypedDict): + type: Literal['includes'] + target: str + includes: str + file: str + + +Node = Union[Container, Enum, Typedef, Callback, Includes] @dataclass @@ -60,46 +167,55 @@ class Definition: name: str kind: str # interface, dictionary or enum parent: str | None = None - members: list[Node] = field(default_factory=list) + members: list[Member] = field(default_factory=list) values: list[str] = field(default_factory=list) @dataclass class Spec: definitions: dict[str, Definition] - typedefs: dict[str, Node] + typedefs: dict[str, IdlType] def lineage(self, name: str) -> list[Definition]: """The definition and its ancestors that are part of the spec, nearest first.""" - chain = [] - while name in self.definitions: - chain.append(self.definitions[name]) - name = self.definitions[name].parent + chain: list[Definition] = [] + current: str | None = name + while current is not None and current in self.definitions: + chain.append(self.definitions[current]) + current = self.definitions[current].parent return chain - def members(self, name: str) -> list[Node]: + def members(self, name: str) -> list[Member]: """The members of a definition, inherited ones included.""" return [member for definition in self.lineage(name) for member in definition.members] - def named_types(self, idl_type: Node | list[Node] | str, known: Callable[[str], bool] | None = None) -> set[str]: + def fields(self, name: str) -> list[Field]: + """The members of a dictionary, inherited ones included.""" + return [member for member in self.members(name) if member['type'] == 'field'] + + def named_types( + self, idl_type: IdlType | list[IdlType] | str, known: Callable[[str], bool] | None = None + ) -> set[str]: """The definitions a type refers to through unions, generics and typedefs, but not ``known`` typedefs.""" if isinstance(idl_type, list): - return set().union(*(self.named_types(item, known) for item in idl_type)) + none: set[str] = set() + return none.union(*(self.named_types(item, known) for item in idl_type)) if isinstance(idl_type, dict): return self.named_types(idl_type['idlType'], known) if idl_type in self.typedefs: - return set() if known and known(idl_type) else self.named_types(self.typedefs[idl_type], known) + skip = known is not None and known(idl_type) + return set() if skip else self.named_types(self.typedefs[idl_type], known) return {idl_type} if idl_type in self.definitions else set() - def dictionary(self, idl_type: Node) -> Definition | None: + def dictionary(self, idl_type: IdlType) -> Definition | None: """The dictionary a type is, unless it's a union or a generic.""" - while not idl_type['union'] and not idl_type['generic'] and isinstance(idl_type['idlType'], str): + while not idl_type['union'] and idl_type['generic'] == '' and isinstance(idl_type['idlType'], str): name = idl_type['idlType'] if name in self.typedefs: idl_type = self.typedefs[name] continue definition = self.definitions.get(name) - return definition if definition and definition.kind == 'dictionary' else None + return definition if definition is not None and definition.kind == 'dictionary' else None return None @@ -110,7 +226,7 @@ def load() -> Spec: def parse(texts: dict[str, str], files: dict[str, set[str] | None]) -> Spec: """Parses IDL texts by file name and keeps the definitions each file takes, as in :data:`FILES`.""" - nodes = json.loads( + nodes: list[Node] = json.loads( subprocess.run( [sys.executable, '-m', 'tests.idl.child'], input=json.dumps(texts), @@ -120,55 +236,75 @@ def parse(texts: dict[str, str], files: dict[str, set[str] | None]) -> Spec: ).stdout ) spec = Spec(_merge(nodes), {node['name']: node['idlType'] for node in nodes if node['type'] == 'typedef'}) - taken = { - node['name'] - for node in nodes - if node.get('name') in spec.definitions - and not node.get('partial') - and (files[node['file']] is None or node['name'] in files[node['file']]) - } + taken: set[str] = set() + for node in nodes: + if node['type'] == 'includes' or node['name'] not in spec.definitions or _is_partial(node): + continue + wanted = files[node['file']] + if wanted is None or node['name'] in wanted: + taken.add(node['name']) return Spec({name: spec.definitions[name] for name in sorted(_used(spec, taken))}, spec.typedefs) +def _is_partial(node: Node) -> TypeGuard[Container]: + return ( + node['type'] in {'interface', 'interface mixin', 'dictionary', 'namespace', 'callback interface'} + and node['partial'] + ) + + +def _definitions(nodes: list[Node]) -> dict[str, Definition]: + """The interfaces, dictionaries and enums, without their partials.""" + definitions: dict[str, Definition] = {} + for node in nodes: + if node['type'] in {'interface', 'dictionary'} and not node['partial']: + definitions[node['name']] = Definition( + node['name'], node['type'], node['inheritance'], list(node['members']) + ) + elif node['type'] == 'enum': + values = [value['value'] for value in node['values']] + definitions[node['name']] = Definition(node['name'], node['type'], values=values) + return definitions + + def _merge(nodes: list[Node]) -> dict[str, Definition]: """The interfaces, dictionaries and enums, with the members of their partials and mixins.""" - definitions = { - node['name']: Definition( - node['name'], - node['type'], - node.get('inheritance'), - list(node.get('members', [])), - [value['value'] for value in node.get('values', [])], - ) - for node in nodes - if node['type'] in {'interface', 'dictionary', 'enum'} and not node.get('partial') - } - mixins: dict[str, list[Node]] = {} + definitions = _definitions(nodes) + mixins: dict[str, list[Member]] = {} for node in nodes: if node['type'] == 'interface mixin': mixins.setdefault(node['name'], []).extend(node['members']) # partials and mixins extend definitions of any file, but only those that exist for node in nodes: - if node.get('partial') and node['name'] in definitions: + if node['type'] == 'includes': + if node['target'] in definitions: + definitions[node['target']].members.extend(mixins.get(node['includes'], [])) + elif node['name'] in definitions and _is_partial(node): definitions[node['name']].members.extend(node['members']) - elif node['type'] == 'includes' and node['target'] in definitions: - definitions[node['target']].members.extend(mixins.get(node['includes'], [])) return definitions +def _types(member: Member) -> list[IdlType | list[IdlType]]: + """The type of a member and of its arguments.""" + types: list[IdlType | list[IdlType]] = [] if member['type'] == 'constructor' else [member['idlType']] + if member['type'] not in {'attribute', 'const', 'field'}: + types += [argument['idlType'] for argument in member['arguments']] + return types + + def _used(spec: Spec, taken: set[str]) -> set[str]: """The taken definitions, their ancestors and the dictionaries and enums they use.""" wanted = set(taken) queue = list(taken) - while queue: + while len(queue) > 0: definition = spec.definitions[queue.pop()] - used = {definition.parent} & spec.definitions.keys() + used: set[str] = set() + if definition.parent is not None and definition.parent in spec.definitions: + used.add(definition.parent) for member in definition.members: - types = [member.get('idlType')] + [argument['idlType'] for argument in member.get('arguments') or []] used |= { name - for idl_type in types - if idl_type + for idl_type in _types(member) for name in spec.named_types(idl_type) if spec.definitions[name].kind != 'interface' } diff --git a/tests/idl/test_compare.py b/tests/idl/test_compare.py index 043869f..2a41cc7 100644 --- a/tests/idl/test_compare.py +++ b/tests/idl/test_compare.py @@ -89,14 +89,15 @@ def add(self, items: list[str]) -> None: ... def send_dtmf(self, tones: str) -> None: ... - def create(self) -> Thing: ... + def create(self) -> Thing: + raise NotImplementedError Thing.sharedName = Thing.shared_name Thing.sendDTMF = Thing.send_dtmf -class Report(UserDict): +class Report(UserDict[str, object]): pass @@ -161,6 +162,7 @@ def __getitem__(self, key: str) -> object: ... def test_future_counts_as_async() -> None: class Thing: - def start(self) -> asyncio.Future[None]: ... + def start(self) -> asyncio.Future[None]: + raise NotImplementedError assert 'start: should be async' not in differences(Thing=Thing)['Thing'] diff --git a/tests/idl/test_idl.py b/tests/idl/test_idl.py index 4732bc0..ae333da 100644 --- a/tests/idl/test_idl.py +++ b/tests/idl/test_idl.py @@ -28,10 +28,11 @@ @pytest.fixture(scope='module') def differences() -> dict[str, list[str]]: + assert SPEC is not None return compare(SPEC, webrtc) -@pytest.mark.parametrize('name', sorted(set(SPEC.definitions) | set(EXPECTED)) if SPEC else []) +@pytest.mark.parametrize('name', sorted(set(SPEC.definitions) | set(EXPECTED)) if SPEC is not None else []) def test_definition(name: str, differences: dict[str, list[str]]) -> None: problems = expectations.mismatches(EXPECTED.get(name, []), differences.get(name, [])) - assert not problems, '\n'.join([*problems, 'run `python -m tests.idl update` if the change is intended']) + assert len(problems) == 0, '\n'.join([*problems, 'run `python -m tests.idl update` if the change is intended']) diff --git a/tests/rtc_peer_connection/test_add_track.py b/tests/rtc_peer_connection/test_add_track.py index 898926e..3ca703d 100644 --- a/tests/rtc_peer_connection/test_add_track.py +++ b/tests/rtc_peer_connection/test_add_track.py @@ -180,14 +180,18 @@ async def test_10( await wait_for_ice_gathering_complete(callee) second_track, *_ = audio_stream2.get_tracks() - candidates = [] - caller.on('icecandidate', lambda event: candidates.append(event.candidate)) + candidates: list[webrtc.RTCIceCandidate | None] = [] + + def on_candidate(event: webrtc.RTCPeerConnectionIceEvent) -> None: + candidates.append(event.candidate) + + caller.on('icecandidate', on_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' + assert len(candidates) == 0, '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 diff --git a/tests/rtc_peer_connection/test_add_transceiver.py b/tests/rtc_peer_connection/test_add_transceiver.py index 2ba5628..fdb5949 100644 --- a/tests/rtc_peer_connection/test_add_transceiver.py +++ b/tests/rtc_peer_connection/test_add_transceiver.py @@ -10,6 +10,7 @@ import pytest import webrtc +from tests.helpers import mistyped def test_1(pc: webrtc.RTCPeerConnection) -> None: @@ -17,7 +18,7 @@ def test_1(pc: webrtc.RTCPeerConnection) -> None: assert hasattr(pc, 'add_transceiver') with pytest.raises(TypeError): - pc.add_transceiver('invalid') + pc.add_transceiver(mistyped('invalid')) def _create_and_test_transceiver(pc: webrtc.RTCPeerConnection, kind: webrtc.MediaType) -> None: @@ -79,7 +80,7 @@ def test_4(pc: webrtc.RTCPeerConnection) -> None: def test_5() -> None: """An init with an invalid direction can't be created, so add_transceiver can't get one.""" with pytest.raises(ValueError, match='not a valid TransceiverDirection'): - webrtc.RTCRtpTransceiverInit(direction='invalid') + webrtc.RTCRtpTransceiverInit(direction=mistyped('invalid')) def test_6(pc: webrtc.RTCPeerConnection, audio_stream: webrtc.MediaStream) -> None: diff --git a/tests/test_audio_data.py b/tests/test_audio_data.py index de4a5a6..de300d7 100644 --- a/tests/test_audio_data.py +++ b/tests/test_audio_data.py @@ -11,19 +11,29 @@ import array import struct +from typing import TYPE_CHECKING import pytest import webrtc +from tests.helpers import mistyped from webrtc import AudioSampleFormat +if TYPE_CHECKING: + from collections.abc import Callable + def f32(*values: float) -> bytes: return array.array('f', values).tobytes() def audio_data( - *, format: str = 'f32-planar', channels: int = 2, frames: int = 5, data: bytes | None = None, **init: object + *, + format: webrtc.AudioSampleFormatValue = '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( @@ -52,7 +62,7 @@ def test_construct() -> None: def test_init_from_json() -> None: """The init comes from its JSON form with camelCase names, a dictionary isn't taken itself.""" - init = { + init: dict[str, object] = { 'format': 's16', 'sampleRate': 48000, 'numberOfFrames': 480, @@ -69,18 +79,19 @@ def test_init_from_json() -> None: @pytest.mark.parametrize( - 'change', + 'create', [ - {'format': 'x32'}, - {'frames': 0}, - {'channels': 0}, - {'data': bytes(3)}, + lambda: audio_data(format=mistyped('x32')), + lambda: audio_data(frames=0), + lambda: audio_data(channels=0), + lambda: audio_data(data=bytes(3)), ], + ids=['change0', 'change1', 'change2', 'change3'], ) -def test_invalid_init(change: dict[str, object]) -> None: +def test_invalid_init(create: Callable[[], webrtc.AudioData]) -> None: """An invalid init, or data too small for it, is a TypeError.""" with pytest.raises(TypeError): - audio_data(**change) + create() def test_close_and_clone() -> None: @@ -157,7 +168,7 @@ def test_destination_too_small() -> None: @pytest.mark.parametrize('source', VALUES) @pytest.mark.parametrize('destination', VALUES) -def test_sample_conversions(source: str, destination: str) -> None: +def test_sample_conversions(source: webrtc.AudioSampleFormatValue, destination: webrtc.AudioSampleFormatValue) -> 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()) @@ -172,7 +183,7 @@ def test_sample_conversions(source: str, destination: str) -> None: @pytest.mark.parametrize('destination', ['u8', 's16', 's32']) -def test_non_finite_f32_samples_convert(destination: str) -> None: +def test_non_finite_f32_samples_convert(destination: webrtc.AudioSampleFormatValue) -> 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()) diff --git a/tests/test_configuration.py b/tests/test_configuration.py index 5d30979..c06faca 100644 --- a/tests/test_configuration.py +++ b/tests/test_configuration.py @@ -18,7 +18,7 @@ import pytest import webrtc -from tests.helpers import exchange_offer +from tests.helpers import exchange_offer, mistyped if TYPE_CHECKING: from tests.helpers import CreatePC @@ -159,9 +159,12 @@ async def test_always_negotiate_data_channels_and_header_encryption(create_pc: C assert 'a=cryptex' in offer.sdp pc.set_configuration(configuration) - for changed in ({'always_negotiate_data_channels': False}, {'rtp_header_encryption_policy': 'negotiate'}): + for changed in ( + dataclasses.replace(configuration, always_negotiate_data_channels=False), + dataclasses.replace(configuration, rtp_header_encryption_policy='negotiate'), + ): with pytest.raises(webrtc.InvalidModificationError): - pc.set_configuration(dataclasses.replace(configuration, **changed)) + pc.set_configuration(changed) @pytest.mark.asyncio @@ -216,7 +219,9 @@ async def test_configured_certificate(create_pc: CreatePC) -> None: pc.add_transceiver(webrtc.MediaType.audio) offer = await pc.create_offer() assert fingerprint.value.upper() in offer.sdp - assert pc.get_configuration().certificates[0].get_fingerprints() == [fingerprint] + certificates = pc.get_configuration().certificates + assert certificates is not None + assert certificates[0].get_fingerprints() == [fingerprint] @pytest.mark.asyncio @@ -282,8 +287,12 @@ def test_peer_reflexive_candidate_hides_its_address() -> None: 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', + 'sdp_m_line_index': 0, + 'username_fragment': None, + 'url': None, + 'relay_protocol': None, }) - assert not candidate.candidate + assert candidate.candidate == '' assert candidate.type == webrtc.RTCIceCandidateType.prflx assert candidate.address is None assert candidate.port == 62341 @@ -298,7 +307,7 @@ def test_rtc_error() -> None: 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') + webrtc.RTCErrorInit(mistyped('nonsense')) def test_rtc_error_init_from_json() -> None: @@ -335,6 +344,8 @@ async def test_created_descriptions(pc: webrtc.RTCPeerConnection) -> None: ) await pc.set_local_description(offer) # not compared by identity: gathered candidates change the description (and its object) between reads + assert pc.pending_local_description is not None + assert pc.local_description is not None assert pc.pending_local_description.type == pc.local_description.type == offer.type assert pc.current_local_description is None @@ -349,10 +360,12 @@ async def test_provisional_answers_without_sdp( await callee.set_local_description(webrtc.RTCSessionDescriptionInit('pranswer')) assert callee.signaling_state == webrtc.RTCSignalingState.have_local_pranswer + assert callee.pending_local_description is not None assert callee.pending_local_description.type == webrtc.RTCSdpType.pranswer # without a type, the final answer await callee.set_local_description() assert callee.signaling_state == webrtc.RTCSignalingState.stable + assert callee.current_local_description is not None assert callee.current_local_description.type == webrtc.RTCSdpType.answer diff --git a/tests/test_data_channel.py b/tests/test_data_channel.py index d698ad1..43b775b 100644 --- a/tests/test_data_channel.py +++ b/tests/test_data_channel.py @@ -12,13 +12,31 @@ import asyncio import pytest +from typing_extensions import TypedDict, Unpack import webrtc -from tests.helpers import connect, wait_for_event, wait_until +from tests.helpers import connect, mistyped, wait_for_event, wait_until + + +class ChannelOptions(TypedDict, total=False, closed=True): + """Options of RTCDataChannelInit the tests set.""" + + max_packet_life_time: int + max_retransmits: int + protocol: str + negotiated: bool + id: int + + +class Negotiation(TypedDict, total=False, closed=True): + """Whether a channel is negotiated, and its id.""" + + negotiated: bool + id: int async def open_pair( - caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection, **options: object + caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection, **options: Unpack[ChannelOptions] ) -> tuple[webrtc.RTCDataChannel, webrtc.RTCDataChannel]: """Opens a channel of the caller and returns it with its remote end.""" init = webrtc.RTCDataChannelInit(**options) @@ -33,7 +51,9 @@ async def open_pair( opened = wait_for_event(channel, 'open') announced = wait_for_event(callee, 'datachannel') await connect(caller, callee) - remote = (await announced).channel + event = await announced + assert isinstance(event, webrtc.RTCDataChannelEvent) + remote = event.channel await opened return channel, remote @@ -41,7 +61,7 @@ async def open_pair( @pytest.mark.asyncio @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] + caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection, options: Negotiation ) -> 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) @@ -51,7 +71,7 @@ async def test_messages_both_ways( assert channel.ready_state == remote.ready_state == webrtc.RTCDataChannelState.open assert channel.id == remote.id - received = [] + received: list[str | bytes | webrtc.Blob] = [] got_all = asyncio.get_running_loop().create_future() @remote.on('message') @@ -85,14 +105,14 @@ async def test_buffered_amount_and_low_event( 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 = [] + remote_events: list[tuple[str, webrtc.RTCDataChannelState]] = [] def on_remote_closing(_event: webrtc.Event) -> None: remote_events.append(('closing', remote.ready_state)) remote.on('closing', on_remote_closing) remote_closed = wait_for_event(remote, 'close') - local_closing = [] + local_closing: list[webrtc.Event] = [] channel.on('closing', local_closing.append) channel.close() @@ -115,7 +135,7 @@ def on_remote_closing(_event: webrtc.Event) -> None: ], ids=['both limits', 'negotiated without id', 'id out of range'], ) -def test_invalid_data_channel_init(pc: webrtc.RTCPeerConnection, init: dict[str, object], error: str) -> None: +def test_invalid_data_channel_init(pc: webrtc.RTCPeerConnection, init: ChannelOptions, 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', webrtc.RTCDataChannelInit(**init)) @@ -142,7 +162,7 @@ def test_data_channel_options(pc: webrtc.RTCPeerConnection) -> None: assert channel.ordered is False assert channel.ready_state == webrtc.RTCDataChannelState.connecting with pytest.raises(TypeError): - channel.send(42) + channel.send(mistyped(42)) def test_create_data_channel_on_closed_connection(pc: webrtc.RTCPeerConnection) -> None: @@ -157,6 +177,7 @@ async def test_max_message_size_before_an_answer(pc: webrtc.RTCPeerConnection) - """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 is not None assert pc.sctp.max_message_size == 65536 @@ -166,7 +187,9 @@ async def test_send_larger_than_max_message_size( ) -> None: """A message larger than the negotiated max message size isn't sent, and the channel stays open.""" channel, _ = await open_pair(caller, callee) + assert caller.sctp is not None size = caller.sctp.max_message_size + assert size is not None assert size > 65536 with pytest.raises(ValueError, match='larger than the maxMessageSize'): @@ -185,6 +208,7 @@ async def test_stats_are_current(caller: webrtc.RTCPeerConnection, callee: webrt await received # libwebrtc reuses a report for 50 ms: the stats right after a message count it after = (await callee.get_stats()).of_type('data-channel')[0].bytes_received + assert isinstance(before, int) assert after == before + 5 @@ -193,10 +217,18 @@ async def test_max_channels_once_connected(caller: webrtc.RTCPeerConnection, cal """The max number of channels is known once SCTP is connected.""" caller.create_data_channel('channels') await caller.set_local_description() + assert caller.sctp is not None assert caller.sctp.max_channels is None await connect(caller, callee) - await wait_until(lambda: callee.sctp.state == webrtc.SctpTransportState.connected, 'SCTP to connect') + + def connected() -> bool: + assert callee.sctp is not None + return callee.sctp.state == webrtc.SctpTransportState.connected + + await wait_until(connected, 'SCTP to connect') + assert callee.sctp is not None + assert callee.sctp.max_channels is not None assert callee.sctp.max_channels > 0 @@ -208,23 +240,29 @@ async def test_binary_type(caller: webrtc.RTCPeerConnection, callee: webrtc.RTCP first = wait_for_event(remote, 'message') channel.send(b'\x01\x02') - assert (await first).data == b'\x01\x02' + message = await first + assert isinstance(message, webrtc.MessageEvent) + assert message.data == b'\x01\x02' remote.binary_type = 'blob' assert remote.binaryType == webrtc.BinaryType.blob second = wait_for_event(remote, 'message') channel.send(webrtc.Blob([b'\x03', 'a', webrtc.Blob([b'\x04'])])) - blob = (await second).data + message = await second + assert isinstance(message, webrtc.MessageEvent) + blob = message.data assert isinstance(blob, webrtc.Blob) 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' + message = await text + assert isinstance(message, webrtc.MessageEvent) + assert message.data == 'text' with pytest.raises(ValueError, match='not a valid BinaryType'): - remote.binary_type = 'buffer' + remote.binary_type = mistyped('buffer') assert remote.binary_type == webrtc.BinaryType.blob @@ -237,4 +275,4 @@ async def test_blob() -> None: 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 not webrtc.Blob(type='é').type + assert webrtc.Blob(type='é').type == '' diff --git a/tests/test_e2e.py b/tests/test_e2e.py index bc41294..d9e1674 100644 --- a/tests/test_e2e.py +++ b/tests/test_e2e.py @@ -19,10 +19,11 @@ async def set_local_and_gather( pc: webrtc.RTCPeerConnection, description: webrtc.RTCSessionDescriptionInit -) -> webrtc.RTCSessionDescription | None: +) -> webrtc.RTCSessionDescription: """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) + assert pc.local_description is not None return pc.local_description @@ -43,6 +44,7 @@ async def test_peers_connect_and_send_audio(caller: webrtc.RTCPeerConnection, ca await callee.set_remote_description(offer) # libwebrtc drops a=end-of-candidates when it serializes a remote description + assert callee.remote_description is not None assert callee.remote_description.sdp == offer.sdp.replace('a=end-of-candidates\r\n', '') answer = await set_local_and_gather(callee, await callee.create_answer()) await caller.set_remote_description(answer) diff --git a/tests/test_enums.py b/tests/test_enums.py index 70b1e9d..78a9d7e 100644 --- a/tests/test_enums.py +++ b/tests/test_enums.py @@ -10,11 +10,13 @@ from __future__ import annotations import enum +import typing import pytest import webrtc import webrtc.enums +from tests.helpers import mistyped def test_members_are_their_values() -> None: @@ -31,13 +33,22 @@ def test_enums_are_exported() -> None: for name, value in vars(webrtc.enums).items() if isinstance(value, type) and issubclass(value, enum.Enum) and value.__module__ == 'webrtc.enums' ] - assert enums + assert len(enums) > 0 for name, value in enums: if not name.startswith('_'): assert getattr(webrtc, name) is value assert name in webrtc.__all__ +def test_value_aliases_match_enums() -> None: + """Each Literal alias of the values a parameter takes has exactly the values of its enum.""" + aliases = [name for name in webrtc.__all__ if name.endswith('Value')] + assert len(aliases) > 0 + for name in aliases: + cls = getattr(webrtc, name.removesuffix('Value')) + assert typing.get_args(getattr(webrtc, name)) == tuple(member.value for member in cls), name + + def test_native_getters_return_members(pc: webrtc.RTCPeerConnection) -> None: """The native API returns members.""" assert pc.signaling_state is webrtc.RTCSignalingState.stable @@ -59,13 +70,13 @@ 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' + transceiver.direction = mistyped('nonsense') with pytest.raises(TypeError): - transceiver.direction = 1 + transceiver.direction = mistyped(1) with pytest.raises(TypeError): pc.add_transceiver('data') with pytest.raises(TypeError): - webrtc.RTCPeerConnection(webrtc.RTCConfiguration(bundle_policy='nonsense')) + webrtc.RTCPeerConnection(webrtc.RTCConfiguration(bundle_policy=mistyped('nonsense'))) def test_data_channel_priority(pc: webrtc.RTCPeerConnection) -> None: diff --git a/tests/test_events.py b/tests/test_events.py index 8992d8b..7b4265f 100644 --- a/tests/test_events.py +++ b/tests/test_events.py @@ -20,7 +20,7 @@ @pytest.mark.asyncio 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 = [] + calls: list[tuple[object, ...]] = [] @pc.on('negotiationneeded') def decorated(event: webrtc.Event) -> None: @@ -74,7 +74,7 @@ def test_handlers_need_a_running_loop(pc: webrtc.RTCPeerConnection) -> None: @pytest.mark.asyncio 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 = [] + states: list[webrtc.RTCSignalingState] = [] def on_change(_event: webrtc.Event) -> None: states.append(pc.signaling_state) @@ -93,7 +93,7 @@ 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 = [] + events: list[str] = [] def on_event(event: webrtc.Event) -> None: events.append(event.type) @@ -115,21 +115,26 @@ def on_event(event: webrtc.Event) -> None: @pytest.mark.asyncio 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 = [] + candidates: list[webrtc.RTCIceCandidate | None] = [] def on_candidate(event: webrtc.RTCPeerConnectionIceEvent) -> None: candidates.append(event.candidate) + def end_of_candidates(event: webrtc.Event) -> bool: + assert isinstance(event, webrtc.RTCPeerConnectionIceEvent) + return event.candidate is None + pc.on('icecandidate', on_candidate) - gathered = wait_for_event(pc, 'icecandidate', predicate=lambda event: event.candidate is None) + gathered = wait_for_event(pc, 'icecandidate', predicate=end_of_candidates) pc.add_transceiver(webrtc.MediaType.audio) await pc.set_local_description() await gathered - host = [c for c in candidates if c is not None and c.candidate] - assert host, candidates + host = [c for c in candidates if c is not None and c.candidate != ''] + assert len(host) > 0, 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 any(c is not None and c.candidate == '' for c in candidates), candidates + assert pc.local_description is not None assert 'a=end-of-candidates' in pc.local_description.sdp, pc.local_description.sdp @@ -139,7 +144,7 @@ async def test_descriptions_change_with_signaling_events( ) -> 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 = [] + seen: list[dict[str, object]] = [] def on_change(_event: webrtc.Event) -> None: seen.append({ @@ -155,22 +160,32 @@ def on_change(_event: webrtc.Event) -> None: assert rolled_back == {'state': webrtc.RTCSignalingState.stable, 'local': None, 'remote': None} assert received['state'] == webrtc.RTCSignalingState.have_remote_offer assert received['local'] is None - assert received['remote'].type == webrtc.RTCSdpType.offer + remote = received['remote'] + assert isinstance(remote, webrtc.RTCSessionDescription) + assert remote.type == webrtc.RTCSdpType.offer @pytest.mark.asyncio async def test_restart_ice_before_negotiation_needs_nothing(pc: webrtc.RTCPeerConnection) -> None: """restart_ice before the first negotiation doesn't fire negotiationneeded.""" - events = [] + events: list[webrtc.Event] = [] pc.on('negotiationneeded', events.append) pc.restart_ice() await asyncio.sleep(QUIET_PERIOD) assert events == [] +class _LoopObjects: + """The objects of the first loop, used from the second.""" + + caller: webrtc.RTCPeerConnection + callee: webrtc.RTCPeerConnection + channel: webrtc.RTCDataChannel + + 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 = {} + objects = _LoopObjects() async def first() -> None: caller, callee = webrtc.RTCPeerConnection(), webrtc.RTCPeerConnection() @@ -178,15 +193,15 @@ async def first() -> None: opened = wait_for_event(channel, 'open') await connect(caller, callee) await opened - objects.update(caller=caller, callee=callee, channel=channel) + objects.caller, objects.callee, objects.channel = caller, callee, channel async def second() -> None: - channel = objects['channel'] + channel = objects.channel closed = wait_for_event(channel, 'close') - objects['callee'].close() + objects.callee.close() await closed assert channel.ready_state == webrtc.RTCDataChannelState.closed - objects['caller'].close() + objects.caller.close() asyncio.run(first()) asyncio.run(second()) diff --git a/tests/test_ice_transport.py b/tests/test_ice_transport.py index c7c177f..2cf6406 100644 --- a/tests/test_ice_transport.py +++ b/tests/test_ice_transport.py @@ -20,6 +20,12 @@ from tests.helpers import connect, wait_for_event, wait_until +def local_parameters(transport: webrtc.RTCIceTransport) -> webrtc.RTCIceParameters: + parameters = transport.get_local_parameters() + assert parameters is not None + return parameters + + @pytest.mark.asyncio async def test_candidates_parameters_and_role( caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPeerConnection @@ -27,19 +33,22 @@ async def test_candidates_parameters_and_role( """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() + assert caller.sctp is not None ice = caller.sctp.transport.ice_transport assert ice.role == webrtc.RTCIceRole.unknown assert ice.get_remote_parameters() is None await connect(caller, callee) await wait_until(lambda: ice.role == webrtc.RTCIceRole.controlling, 'the controlling role') + assert callee.sctp is not None remote_ice = callee.sctp.transport.ice_transport local, remote = ice.get_local_parameters(), ice.get_remote_parameters() 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() + assert local.username_fragment != '' + assert local.password != '' + assert remote is not None + assert remote.username_fragment == local_parameters(remote_ice).username_fragment + assert len(ice.get_local_candidates()) > 0 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() @@ -51,6 +60,7 @@ async def test_component(caller: webrtc.RTCPeerConnection, callee: webrtc.RTCPee """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 is not None assert transceiver.sender.transport.ice_transport.component == webrtc.RTCIceComponent.rtp @@ -59,6 +69,7 @@ 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() + assert transceiver.sender.transport is not None ice = transceiver.sender.transport.ice_transport state = ice.gathering_state pc.close() @@ -72,13 +83,13 @@ async def test_two_transports_connect() -> None: local, remote = webrtc.RTCIceTransport(), webrtc.RTCIceTransport() assert local.role is None assert local.state == webrtc.RTCIceTransportState.new - assert local.get_local_parameters() + assert local.get_local_parameters() is not None assert local.get_remote_parameters() is None for transport, other in ((local, remote), (remote, local)): def on_candidate(event: webrtc.RTCPeerConnectionIceEvent, other: webrtc.RTCIceTransport = other) -> None: - if event.candidate: + if event.candidate is not None: other.add_remote_candidate(event.candidate) transport.on('icecandidate', on_candidate) @@ -87,13 +98,14 @@ def on_candidate(event: webrtc.RTCPeerConnectionIceEvent, other: webrtc.RTCIceTr remote.gather() assert local.gathering_state == webrtc.CricketIceGatheringState.gathering # both take the controlling role: one of them switches - local.start(remote.get_local_parameters(), 'controlling') - remote.start(local.get_local_parameters(), 'controlling') + local.start(local_parameters(remote), 'controlling') + remote.start(local_parameters(local), 'controlling') await asyncio.gather(*connected) assert local.state == remote.state == webrtc.RTCIceTransportState.connected assert {local.role, remote.role} == {webrtc.RTCIceRole.controlling, webrtc.RTCIceRole.controlled} pair = local.get_selected_candidate_pair() + assert pair is not None assert pair.local.candidate in [c.candidate for c in local.get_local_candidates()] assert pair.remote.candidate in [c.candidate for c in local.get_remote_candidates()] diff --git a/tests/test_lifetime.py b/tests/test_lifetime.py index fa93b1e..18863c6 100644 --- a/tests/test_lifetime.py +++ b/tests/test_lifetime.py @@ -28,7 +28,7 @@ from tests.helpers import QUIET_PERIOD, connect, exchange_offer_answer, wait_for_event if TYPE_CHECKING: - from collections.abc import Callable, Iterator + from collections.abc import Awaitable, Callable, Iterator def collect() -> None: @@ -36,6 +36,12 @@ def collect() -> None: gc.collect() +async def received_track_event(event: Awaitable[webrtc.Event]) -> webrtc.RTCTrackEvent: + received = await event + assert isinstance(received, webrtc.RTCTrackEvent) + return received + + 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) @@ -93,7 +99,7 @@ def test_factories_return_to_baseline_when_everything_is_gone() -> None: pc.add_track(generator) # one factory for everything - assert alive_factories() == (baseline or 1) + assert alive_factories() == (baseline if baseline != 0 else 1) pc.close() del pc, stream, generator @@ -153,7 +159,10 @@ def test_destroyed_track_wrapper_is_not_notified() -> None: collect() for i in range(10): - pc.get_senders()[0].track.enabled = bool(i % 2) + track = pc.get_senders()[0].track + assert track is not None + track.enabled = bool(i % 2) + del track collect() pc.close() @@ -249,7 +258,10 @@ async def test_drop_and_refetch_transports( collect() assert caller.get_senders()[0].transport == caller.get_transceivers()[0].receiver.transport - assert caller.get_senders()[0].transport.ice_transport == caller.get_receivers()[0].transport.ice_transport + sender_transport, receiver_transport = caller.get_senders()[0].transport, caller.get_receivers()[0].transport + assert sender_transport is not None + assert receiver_transport is not None + assert sender_transport.ice_transport == receiver_transport.ice_transport @pytest.mark.asyncio @@ -261,6 +273,7 @@ async def test_transports_outlive_closed_connection( await exchange_offer_answer(caller, callee) transport = caller.get_senders()[0].transport + assert transport is not None ice_transport = transport.ice_transport caller.close() callee.close() @@ -343,7 +356,7 @@ 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 = [] + received: list[str] = [] def on_close(_event: webrtc.Event) -> None: received.append('close') @@ -430,7 +443,7 @@ async def test_stream_tracks_read_while_they_change( sender = caller.add_track(audio, audio_stream) track_event = wait_for_event(callee, 'track') await exchange_offer_answer(caller, callee) - remote = (await track_event).streams[0] + remote = (await received_track_event(track_event)).streams[0] stop = threading.Event() @@ -595,7 +608,13 @@ def create() -> weakref.ref[webrtc.RTCRtpSender | webrtc.RTCRtpReceiver]: 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 _: owner.track) + track = owner.track + assert track is not None + + def on_ended(_event: webrtc.Event) -> webrtc.MediaStreamTrack | None: + return owner.track + + track.on('ended', on_ended) pc.close() return weakref.ref(owner) @@ -637,7 +656,7 @@ async def session() -> None: channel = caller.create_data_channel('session') received = wait_for_event(callee, 'track') await connect(caller, callee) - remote = (await received).track + remote = (await received_track_event(received)).track reader = webrtc.MediaStreamTrackProcessor(webrtc.MediaStreamTrackProcessorInit(remote)).readable.get_reader() writer = generator.writable.get_writer() await writer.write( @@ -646,7 +665,9 @@ async def session() -> None: webrtc.VideoFrameBufferInit(format='RGBA', coded_width=64, coded_height=48, timestamp=0), ) ) - (await asyncio.wait_for(reader.read(), 5)).value.close() + frame = (await asyncio.wait_for(reader.read(), 5)).value + assert frame is not None + frame.close() await caller.get_stats() channel.send('bye') for track in stream.get_tracks(): diff --git a/tests/test_media_e2e.py b/tests/test_media_e2e.py index 8a28e1d..07e15ab 100644 --- a/tests/test_media_e2e.py +++ b/tests/test_media_e2e.py @@ -58,13 +58,18 @@ def dominant_frequency(samples: list[float], rate: int) -> float: return (len(crossings) - 1) * rate / (crossings[-1] - crossings[0]) -async def read_frames(reader: webrtc.ReadableStreamDefaultReader, count: int) -> tuple[list[int], bytearray]: +async def read_frames( + reader: webrtc.ReadableStreamDefaultReader[webrtc.VideoFrame | webrtc.AudioData], count: int +) -> tuple[list[int], bytearray]: """Reads frames of the size, returns their timestamps and the last one in RGBA.""" - timestamps = [] + timestamps: list[int] = [] for _ in range(count): frame = (await asyncio.wait_for(reader.read(), TIMEOUT)).value + assert isinstance(frame, webrtc.VideoFrame) assert (frame.coded_width, frame.coded_height) == (WIDTH, HEIGHT) - assert frame.metadata().rtp_timestamp > 0 + rtp_timestamp = frame.metadata().rtp_timestamp + assert rtp_timestamp is not None + assert rtp_timestamp > 0 timestamps.append(frame.timestamp) rgba = bytearray(frame.allocation_size(webrtc.VideoFrameCopyToOptions(format='RGBA'))) await frame.copy_to(rgba, webrtc.VideoFrameCopyToOptions(format='RGBA')) @@ -105,9 +110,10 @@ async def test_audio_through_a_connection(caller: webrtc.RTCPeerConnection, call reader = webrtc.MediaStreamTrackProcessor( webrtc.MediaStreamTrackProcessorInit(remote, max_buffer_size=100) ).readable.get_reader() - samples = [] + samples: list[float] = [] for chunk in range(150): audio = (await asyncio.wait_for(reader.read(), TIMEOUT)).value + assert isinstance(audio, webrtc.AudioData) assert audio.sample_rate == 48000 plane = array.array('f', [0.0] * audio.number_of_frames) audio.copy_to( @@ -133,7 +139,9 @@ async def test_remote_track_end_closes_the_processor( 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(webrtc.MediaStreamTrackProcessorInit(remote)).readable.get_reader() - (await asyncio.wait_for(reader.read(), TIMEOUT)).value.close() + frame = (await asyncio.wait_for(reader.read(), TIMEOUT)).value + assert isinstance(frame, webrtc.VideoFrame) + frame.close() callee.close() await asyncio.wait_for(reader.closed, TIMEOUT) generator.track.stop() diff --git a/tests/test_media_stream_track_processor.py b/tests/test_media_stream_track_processor.py index ddd323d..43f2d58 100644 --- a/tests/test_media_stream_track_processor.py +++ b/tests/test_media_stream_track_processor.py @@ -11,6 +11,7 @@ import array import asyncio +from typing import TypeVar import pytest @@ -32,10 +33,23 @@ def video_frame(timestamp: int, width: int = 4, height: int = 2) -> webrtc.Video ) -async def read(reader: webrtc.ReadableStreamDefaultReader) -> webrtc.ReadableStreamReadResult: +_T = TypeVar('_T') + + +async def read(reader: webrtc.ReadableStreamDefaultReader[_T]) -> webrtc.ReadableStreamReadResult[_T]: return await asyncio.wait_for(reader.read(), TIMEOUT) +def as_video_frame(value: object) -> webrtc.VideoFrame: + assert isinstance(value, webrtc.VideoFrame) + return value + + +def as_audio_data(value: object) -> webrtc.AudioData: + assert isinstance(value, webrtc.AudioData) + return value + + @pytest.mark.asyncio 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.""" @@ -95,7 +109,7 @@ async def test_full_buffer_drops_the_oldest_frames(video_stream: webrtc.MediaStr # the total first: frames keep arriving, and the discarded ones are the total less the 2 queued at any time total = processor.total_frames assert processor.discarded_frames >= total - 2 - first, second = (await read(reader)).value, (await read(reader)).value + first, second = as_video_frame((await read(reader)).value), as_video_frame((await read(reader)).value) assert first.timestamp < second.timestamp first.close() second.close() @@ -106,7 +120,7 @@ async def test_cancel_stops_reading(video_stream: webrtc.MediaStream) -> None: """Canceling the stream detaches the processor from the track.""" processor = webrtc.MediaStreamTrackProcessor(webrtc.MediaStreamTrackProcessorInit(video_stream.get_tracks()[0])) reader = processor.readable.get_reader() - (await read(reader)).value.close() + as_video_frame((await read(reader)).value).close() await reader.cancel() total = processor.total_frames await asyncio.sleep(QUIET_PERIOD) @@ -139,7 +153,7 @@ async def test_generator_forwards_frames_with_their_timestamps() -> None: frame = video_frame(timestamp * 1000) await writer.write(frame) assert frame.format is None, 'written frames are closed' - frames = [r.value for r in await asyncio.wait_for(asyncio.gather(*reads), TIMEOUT)] + frames = [as_video_frame(r.value) for r in await asyncio.wait_for(asyncio.gather(*reads), TIMEOUT)] assert [f.timestamp for f in frames] == [0, 1000, 2000, 3000] out = bytearray(12) await frames[0].copy_to(out) @@ -195,7 +209,7 @@ async def test_muted_generator_drops_frames() -> None: generator.muted = False await unmuted await writer.write(video_frame(2)) - frame = (await read(reader)).value + frame = as_video_frame((await read(reader)).value) assert frame.timestamp == 2 frame.close() track.stop() @@ -221,7 +235,7 @@ async def test_audio_generator_sends_10_ms_frames() -> None: ) await writer.write(data) assert data.format is None, 'written data is closed' - received = [(await read(reader)).value for _ in range(2)] + received = [as_audio_data((await read(reader)).value) for _ in range(2)] for audio in received: assert (audio.number_of_frames, audio.sample_rate) == (480, 48000) out = array.array('h', [0] * 480) @@ -265,7 +279,7 @@ def transform(frame: webrtc.VideoFrame, controller: webrtc.TransformStreamDefaul webrtc.MediaStreamTrackProcessorInit(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 + frame = as_video_frame((await read(reader)).value) assert frame.timestamp == 42 assert frame.coded_width == 640 frame.close() @@ -284,6 +298,6 @@ async def test_frames_wait_in_the_native_queue_only(video_stream: webrtc.MediaSt reader = processor.readable.get_reader() frames = await asyncio.wait_for(asyncio.gather(*(reader.read() for _ in range(4))), TIMEOUT) for result in frames: - result.value.close() + as_video_frame(result.value).close() await wait_until(lambda: processor.discarded_frames > 0, 'frames to be dropped') - assert not processor.readable._controller._queue + assert len(processor.readable._controller._queue) == 0 diff --git a/tests/test_media_stress.py b/tests/test_media_stress.py index 8b010a5..7454b9b 100644 --- a/tests/test_media_stress.py +++ b/tests/test_media_stress.py @@ -30,6 +30,12 @@ def frame(timestamp: int = 0, width: int = 16, height: int = 16) -> webrtc.Video ) +def close(media: object) -> None: + """Closes media read from a processor.""" + assert isinstance(media, (webrtc.VideoFrame, webrtc.AudioData)) + media.close() + + def collected(refs: list[weakref.ref[object]]) -> bool: gc.collect() return all(ref() is None for ref in refs) @@ -46,7 +52,7 @@ async def test_cancel_during_pending_reads(video_stream: webrtc.MediaStream) -> await reader.cancel() for result in await asyncio.wait_for(asyncio.gather(*reads), TIMEOUT): if not result.done: - result.value.close() + close(result.value) @pytest.mark.asyncio @@ -54,12 +60,12 @@ async def test_stop_track_while_reading(video_stream: webrtc.MediaStream, audio_ """Tracks stopped while tasks read them close their streams.""" tracks = [*video_stream.get_tracks(), *audio_stream.get_tracks()] - read = [] + read: list[webrtc.MediaStreamTrack] = [] async def read_all(track: webrtc.MediaStreamTrack) -> int: count = 0 async for media in webrtc.MediaStreamTrackProcessor(webrtc.MediaStreamTrackProcessorInit(track)).readable: - media.close() + close(media) count += 1 if count == 1: read.append(track) @@ -82,7 +88,7 @@ async def test_many_processors_of_one_track(video_stream: webrtc.MediaStream) -> for _ in range(20) ] for result in await asyncio.wait_for(asyncio.gather(*(r.read() for r in readers)), TIMEOUT): - result.value.close() + close(result.value) track.stop() await asyncio.wait_for(asyncio.gather(*(r.closed for r in readers)), TIMEOUT) @@ -91,7 +97,7 @@ async def test_many_processors_of_one_track(video_stream: webrtc.MediaStream) -> 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 = [] + refs: list[weakref.ref[object]] = [] for _ in range(20): processor = webrtc.MediaStreamTrackProcessor(webrtc.MediaStreamTrackProcessorInit(track)) processor.readable.get_reader().read() @@ -99,7 +105,7 @@ async def test_garbage_collected_with_pending_reads(video_stream: webrtc.MediaSt del processor await wait_until(lambda: collected(refs), 'the processors to be collected') reader = webrtc.MediaStreamTrackProcessor(webrtc.MediaStreamTrackProcessorInit(track)).readable.get_reader() - (await asyncio.wait_for(reader.read(), TIMEOUT)).value.close() + close((await asyncio.wait_for(reader.read(), TIMEOUT)).value) @pytest.mark.asyncio @@ -114,14 +120,14 @@ async def test_close_connection_while_reading( webrtc.MediaStreamTrackProcessor(webrtc.MediaStreamTrackProcessorInit(remote)).readable.get_reader() for _ in range(5) ] - (await asyncio.wait_for(readers[0].read(), TIMEOUT)).value.close() + close((await asyncio.wait_for(readers[0].read(), TIMEOUT)).value) pending = [r.read() for r in readers] callee.close() caller.close() await asyncio.wait_for(asyncio.gather(*(r.closed for r in readers)), TIMEOUT) for result in await asyncio.gather(*pending): if not result.done: - result.value.close() + close(result.value) generator.track.stop() @@ -160,7 +166,7 @@ async def cycle() -> tuple[weakref.ref[object], weakref.ref[object]]: ) reader = processor.readable.get_reader() await generator.writable.get_writer().write(frame()) - (await reader.read()).value.close() + close((await reader.read()).value) await reader.cancel() generator.track.stop() return weakref.ref(processor), weakref.ref(generator) @@ -205,7 +211,8 @@ async def test_reader_that_never_yields_queues_nothing() -> None: for timestamp in range(2000): # both are done right away, so the loop never runs in between await writer.write(frame(timestamp)) - (await reader.read()).value.close() + close((await reader.read()).value) assert len(TaskQueue.of(loop)._items) <= 2 - assert len(loop._ready) < 100 + # the callbacks ready to run, private to the loop + assert len(vars(loop)['_ready']) < 100 generator.track.stop() diff --git a/tests/test_native_calls.py b/tests/test_native_calls.py index 867ef1b..37b67bd 100644 --- a/tests/test_native_calls.py +++ b/tests/test_native_calls.py @@ -16,15 +16,17 @@ import pytest +import wrtc +from tests.helpers import mistyped from webrtc.utils.native_calls import call_native OnSuccess = Callable[[object], None] -OnFailure = Callable[[SimpleNamespace], None] +OnFailure = Callable[[wrtc.RTCCallbackException], None] -def _error(error: Exception) -> SimpleNamespace: +def _error(error: Exception) -> wrtc.RTCCallbackException: """A stand-in of the native exception passed to on_failure.""" - return SimpleNamespace(toPython=lambda: error) + return mistyped(SimpleNamespace(toPython=lambda: error)) def _later(callback: Callable[..., None], *args: object, delay: float = 0.0) -> None: @@ -41,7 +43,10 @@ def method(on_success: Callable[[int], None], _on_failure: OnFailure, *numbers: @pytest.mark.asyncio async def test_no_result() -> None: - assert await call_native(lambda on_success, _: _later(on_success)) is None + def method(on_success: Callable[[], None], _on_failure: OnFailure) -> None: + _later(on_success) + + assert await call_native(method) is None @pytest.mark.asyncio @@ -57,7 +62,7 @@ def method(_on_success: OnSuccess, on_failure: OnFailure) -> None: 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 = [] + errors: list[dict[str, object]] = [] loop.set_exception_handler(lambda _, context: errors.append(context)) settled = threading.Event() diff --git a/tests/test_robustness_chaos.py b/tests/test_robustness_chaos.py index 34e322a..1653b22 100644 --- a/tests/test_robustness_chaos.py +++ b/tests/test_robustness_chaos.py @@ -25,7 +25,7 @@ def run_chaos(seed: int, steps: int, timeout: float) -> None: try: 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:]}') + pytest.fail(f'seed {seed} is stuck after:\n{(e.stdout if e.stdout is not None else b"")[-3000:]}') output = result.stdout + result.stderr assert result.returncode == 0, f'seed {seed}, exit code {result.returncode}:\n{output[-5000:]}' assert 'done' in result.stdout, output[-3000:] diff --git a/tests/test_rtp_sender_receiver.py b/tests/test_rtp_sender_receiver.py index 05b086b..c0a8e63 100644 --- a/tests/test_rtp_sender_receiver.py +++ b/tests/test_rtp_sender_receiver.py @@ -11,19 +11,26 @@ import asyncio import time +from typing import TYPE_CHECKING import pytest import webrtc -from tests.helpers import connect, exchange_offer_answer, next_task, wait_for_event, wait_until +from tests.helpers import connect, exchange_offer_answer, mistyped, next_task, wait_for_event, wait_until + +if TYPE_CHECKING: + from collections.abc import Callable 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 audio is not None assert any(codec.mime_type == 'audio/opus' for codec in audio.codecs) - assert audio.header_extensions - assert webrtc.RTCRtpReceiver.get_capabilities('video').codecs + assert len(audio.header_extensions) > 0 + video = webrtc.RTCRtpReceiver.get_capabilities('video') + assert video is not None + assert len(video.codecs) > 0 assert webrtc.RTCRtpSender.get_capabilities('data') is None @@ -89,18 +96,18 @@ async def test_parameters_expire_with_their_task(pc: webrtc.RTCPeerConnection) - @pytest.mark.parametrize( - ('encodings', 'error'), + ('rids', 'error'), [ - ([{'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'), + (['a', 'a'], 'needs a distinct rid'), + (['a', None], 'needs a distinct rid'), + (['no-dash'], 'not a valid rid'), + ([''], 'not a valid rid'), ], ids=['duplicate rid', 'missing rid', 'invalid rid', 'empty rid'], ) -def test_invalid_send_encodings(pc: webrtc.RTCPeerConnection, encodings: list[dict[str, str]], error: str) -> None: +def test_invalid_send_encodings(pc: webrtc.RTCPeerConnection, rids: list[str | None], error: str) -> None: """The rids of send encodings are unique, present when there are several, and alphanumeric.""" - init = webrtc.RTCRtpTransceiverInit(send_encodings=[webrtc.RTCRtpEncodingParameters(**e) for e in encodings]) + init = webrtc.RTCRtpTransceiverInit(send_encodings=[webrtc.RTCRtpEncodingParameters(rid=rid) for rid in rids]) with pytest.raises(ValueError, match=error): pc.add_transceiver(webrtc.MediaType.video, init) @@ -118,8 +125,8 @@ async def test_negotiated_codecs(caller: webrtc.RTCPeerConnection, callee: webrt transceiver = caller.add_transceiver(webrtc.MediaType.audio) await exchange_offer_answer(caller, callee) - assert transceiver.sender.get_parameters().codecs - assert callee.get_transceivers()[0].receiver.get_parameters().codecs + assert len(transceiver.sender.get_parameters().codecs) > 0 + assert len(callee.get_transceivers()[0].receiver.get_parameters().codecs) > 0 @pytest.mark.asyncio @@ -140,7 +147,9 @@ async def test_replace_track( 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'] + capabilities = webrtc.RTCRtpReceiver.get_capabilities('audio') + assert capabilities is not None + opus = [c for c in capabilities.codecs if c.mime_type == 'audio/opus'] transceiver.set_codec_preferences(opus) transceiver.set_codec_preferences([]) with pytest.raises(webrtc.InvalidModificationError): @@ -159,9 +168,11 @@ async def test_sender_codecs_leave_out_unknown_remote_codecs( """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() + assert caller.local_description is not None await callee.set_remote_description(caller.local_description) await callee.set_local_description() # the answer lists a codec this side doesn't know first + assert callee.local_description is not None sdp = callee.local_description.sdp m_line = next(line for line in sdp.split('\r\n') if line.startswith('m=audio')) sdp = sdp.replace(m_line, m_line + ' 125').replace( @@ -170,7 +181,7 @@ async def test_sender_codecs_leave_out_unknown_remote_codecs( await caller.set_remote_description(webrtc.RTCSessionDescriptionInit('answer', sdp)) parameters = sender.get_parameters() - assert parameters.codecs + assert len(parameters.codecs) > 0 assert all('flarglblurp' not in codec.mime_type for codec in parameters.codecs) await sender.set_parameters(parameters) @@ -214,8 +225,8 @@ async def test_simulcast_receiver_parameters( caller.add_transceiver(webrtc.MediaType.video, webrtc.RTCRtpTransceiverInit(send_encodings=encodings)) await exchange_offer_answer(caller, callee) parameters = callee.get_transceivers()[0].receiver.get_parameters() - assert parameters.codecs - assert parameters.header_extensions + assert len(parameters.codecs) > 0 + assert len(parameters.header_extensions) > 0 @pytest.mark.asyncio @@ -229,9 +240,17 @@ async def test_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: not event.tone) + tones: list[str] = [] + + def on_tone_change(event: webrtc.RTCDTMFToneChangeEvent) -> None: + tones.append(event.tone) + + def is_last(event: webrtc.Event) -> bool: + assert isinstance(event, webrtc.RTCDTMFToneChangeEvent) + return event.tone == '' + + dtmf.on('tonechange', on_tone_change) + done = wait_for_event(dtmf, 'tonechange', timeout=5, predicate=is_last) with pytest.raises(webrtc.InvalidCharacterError): dtmf.insert_dtmf('12X') @@ -249,7 +268,9 @@ async def test_synchronization_sources( caller.add_track(video_stream.get_tracks()[0], video_stream) remote_track = wait_for_event(callee, 'track') await connect(caller, callee) - receiver = (await remote_track).receiver + track_event = await remote_track + assert isinstance(track_event, webrtc.RTCTrackEvent) + receiver = track_event.receiver # sources are known once media is decoded (audio is only played out by a real audio device) await wait_until(receiver.get_synchronization_sources, 'a synchronization source', timeout=5) @@ -266,19 +287,21 @@ async def test_synchronization_sources( @pytest.mark.parametrize( 'encoding', [ - {'max_bitrate': -1}, - {'max_bitrate': 2**32}, - {'max_bitrate': 1.5}, - {'max_framerate': float('inf')}, - {'scale_resolution_down_by': float('nan')}, + lambda: webrtc.RTCRtpEncodingParameters(max_bitrate=-1), + lambda: webrtc.RTCRtpEncodingParameters(max_bitrate=2**32), + lambda: webrtc.RTCRtpEncodingParameters(max_bitrate=mistyped(1.5)), + lambda: webrtc.RTCRtpEncodingParameters(max_framerate=float('inf')), + lambda: webrtc.RTCRtpEncodingParameters(scale_resolution_down_by=float('nan')), ], ) -def test_encodings_have_their_webidl_types(pc: webrtc.RTCPeerConnection, encoding: dict[str, float]) -> None: +def test_encodings_have_their_webidl_types( + pc: webrtc.RTCPeerConnection, encoding: Callable[[], webrtc.RTCRtpEncodingParameters] +) -> 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, - webrtc.RTCRtpTransceiverInit(send_encodings=[webrtc.RTCRtpEncodingParameters(**encoding)]), + webrtc.RTCRtpTransceiverInit(send_encodings=[encoding()]), ) diff --git a/tests/test_session_description.py b/tests/test_session_description.py index b1735fd..f1e5bbe 100644 --- a/tests/test_session_description.py +++ b/tests/test_session_description.py @@ -9,21 +9,28 @@ from __future__ import annotations +from typing import TYPE_CHECKING + import pytest import webrtc +from tests.helpers import mistyped + +if TYPE_CHECKING: + from collections.abc import Callable def test_type_is_required() -> None: """The init of a description needs its type, as its WebIDL dictionary does.""" + no_arguments: Callable[[], object] = mistyped(webrtc.RTCSessionDescription) with pytest.raises(TypeError): - webrtc.RTCSessionDescription() + no_arguments() with pytest.raises(TypeError): webrtc.RTCSessionDescriptionInit.from_json({'sdp': ''}) with pytest.raises(TypeError): - webrtc.RTCSessionDescription('offer', None) + webrtc.RTCSessionDescription('offer', mistyped(None)) with pytest.raises(ValueError, match='not a valid RTCSdpType'): - webrtc.RTCSessionDescription({'type': 'offer'}) + webrtc.RTCSessionDescription(mistyped({'type': 'offer'})) @pytest.mark.parametrize( @@ -38,5 +45,5 @@ def test_sdp_is_empty_by_default(init: webrtc.RTCSessionDescriptionInit | webrtc """A description from the JSON form, an init or a type has an empty SDP.""" description = webrtc.RTCSessionDescription(init) assert description.type == webrtc.RTCSdpType.rollback - assert not description.sdp + assert description.sdp == '' assert description.to_json() == {'type': 'rollback', 'sdp': ''} diff --git a/tests/test_stats.py b/tests/test_stats.py index 1092d77..12aee6e 100644 --- a/tests/test_stats.py +++ b/tests/test_stats.py @@ -24,7 +24,9 @@ async def send_audio( caller.add_track(stream.get_tracks()[0], stream) track_event = wait_for_event(callee, 'track') await connect(caller, callee) - remote = (await track_event).track + event = await track_event + assert isinstance(event, webrtc.RTCTrackEvent) + remote = event.track await wait_until_unmuted(remote) return remote @@ -38,7 +40,7 @@ async def test_connection_stats( report = await caller.get_stats() assert isinstance(report, webrtc.RTCStatsReport) - assert report.of_type('peer-connection') + assert len(report.of_type('peer-connection')) > 0 outbound = report.of_type('outbound-rtp')[0] assert outbound['kind'] == outbound.kind == 'audio' assert abs(outbound.timestamp - time.time() * 1000) < 60_000 @@ -52,8 +54,8 @@ async def test_sender_stats( await send_audio(caller, callee, audio_stream) sender_report = await caller.get_senders()[0].get_stats() - assert sender_report.of_type('outbound-rtp') - assert not sender_report.of_type('inbound-rtp') + assert len(sender_report.of_type('outbound-rtp')) > 0 + assert len(sender_report.of_type('inbound-rtp')) == 0 assert len(await caller.get_stats(audio_stream.get_tracks()[0])) == len(sender_report) @@ -88,7 +90,7 @@ async def test_closed_connection_has_stats( """A closed connection still has stats.""" await send_audio(caller, callee, audio_stream) caller.close() - assert (await caller.get_stats()).of_type('peer-connection') + assert len((await caller.get_stats()).of_type('peer-connection')) > 0 @pytest.mark.asyncio @@ -101,6 +103,10 @@ async def test_remote_audio_is_played_out( async def decoded() -> bool: inbound = (await receiver.get_stats()).of_type('inbound-rtp') - return bool(inbound) and inbound[0].get('totalSamplesReceived', 0) > 0 + if len(inbound) == 0: + return False + received = inbound[0].get('totalSamplesReceived', 0) + assert isinstance(received, int) + return received > 0 await wait_until(decoded, 'decoded remote audio') diff --git a/tests/test_streams.py b/tests/test_streams.py index 270893d..53c8308 100644 --- a/tests/test_streams.py +++ b/tests/test_streams.py @@ -15,6 +15,7 @@ from typing import TYPE_CHECKING, NoReturn import pytest +from typing_extensions import override import webrtc from tests.helpers import wait_until @@ -33,7 +34,7 @@ def __init__(self, chunks: Iterable[object]) -> None: def pull(self, controller: webrtc.ReadableStreamDefaultController) -> None: self.pulls += 1 - if self.chunks: + if len(self.chunks) > 0: controller.enqueue(self.chunks.pop(0)) else: controller.close() @@ -45,6 +46,8 @@ def cancel(self, reason: object) -> None: class Controlled: """An underlying source keeping its controller, for the test to enqueue or error.""" + controller: webrtc.ReadableStreamDefaultController + def start(self, controller: webrtc.ReadableStreamDefaultController) -> None: self.controller = controller @@ -77,7 +80,7 @@ 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 = [] + seen: list[object] = [] async for chunk in stream: seen.append(chunk) if chunk == 2: @@ -91,8 +94,13 @@ async def test_async_iteration_and_cancel() -> None: @pytest.mark.asyncio async def test_cancel_reaches_source() -> None: """Canceling a stream calls the source and settles pending reads as done.""" - source = Chunks([]) - source.pull = lambda _: None + + class Idle(Chunks): + @override + def pull(self, controller: webrtc.ReadableStreamDefaultController) -> None: + pass + + source = Idle([]) stream = webrtc.ReadableStream(source, high_water_mark=0) reader = stream.get_reader() read = reader.read() @@ -132,7 +140,7 @@ async def test_locked_stream() -> None: @pytest.mark.asyncio async def test_writer_backpressure_and_order() -> None: """Writes reach the sink in order, one at a time, and ready follows the queue.""" - written = [] + written: list[int] = [] class Sink: @staticmethod @@ -173,7 +181,7 @@ def write(_chunk: int, _controller: webrtc.WritableStreamDefaultController) -> N @pytest.mark.asyncio async def test_abort_drops_queued_writes() -> None: """Aborting fails the writes not done yet and tells the sink.""" - reasons = [] + reasons: list[str] = [] class Sink: @staticmethod @@ -196,7 +204,7 @@ def abort(reason: str) -> None: @pytest.mark.asyncio async def test_pipe_through_transform() -> None: """A readable stream piped through a transform stream into a writable one.""" - written = [] + written: list[int] = [] class Double: @staticmethod @@ -216,7 +224,7 @@ def write(chunk: int, _controller: webrtc.WritableStreamDefaultController) -> No @pytest.mark.asyncio async def test_pipe_to_aborts_on_error() -> None: """An error of the source aborts the destination.""" - aborted = [] + aborted: list[Exception] = [] class Source: @staticmethod @@ -245,6 +253,7 @@ class Woken: def __init__(self, waiting: weakref.WeakSet[asyncio.Future[None]]) -> None: self.waiting = waiting self.next = 0 + self.woken: asyncio.Future[None] | None = None def pull(self, controller: webrtc.ReadableStreamDefaultController) -> asyncio.Future[None]: # kept by the source, as a processor keeps its pending read @@ -272,12 +281,15 @@ 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[asyncio.Future[None]] = weakref.WeakSet() - written = [] + written: list[int] = [] + + def write(chunk: int, _controller: webrtc.WritableStreamDefaultController) -> None: + written.append(chunk) def start() -> asyncio.Future[None]: # only the last pipe is referenced source = webrtc.ReadableStream(Woken(waiting), high_water_mark=0) - sink = webrtc.WritableStream({'write': lambda chunk, _: written.append(chunk)}) + sink = webrtc.WritableStream({'write': write}) return source.pipe_through(webrtc.TransformStream()).pipe_to(sink) done = start() @@ -293,11 +305,21 @@ def start() -> asyncio.Future[None]: @pytest.mark.asyncio 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, _: written.append(chunk)}) - pipe = source.pipe_through(transform).pipe_to(sink) + written: list[int] = [] + + def pull(controller: webrtc.ReadableStreamDefaultController) -> None: + controller.enqueue(2) + + def transform(chunk: int, controller: webrtc.TransformStreamDefaultController) -> None: + controller.enqueue(chunk * 10) + + def write(chunk: int, _controller: webrtc.WritableStreamDefaultController) -> None: + written.append(chunk) + + source = webrtc.ReadableStream({'pull': pull}) + transform_stream = webrtc.TransformStream({'transform': transform}) + sink = webrtc.WritableStream({'write': write}) + pipe = source.pipe_through(transform_stream).pipe_to(sink) await wait_until(lambda: len(written) >= 3, 'chunks written') pipe.cancel() assert written[:3] == [20, 20, 20] diff --git a/tests/test_task_queue.py b/tests/test_task_queue.py index 0076b9f..b0cf836 100644 --- a/tests/test_task_queue.py +++ b/tests/test_task_queue.py @@ -24,7 +24,7 @@ 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 = [] + order: list[int | str] = [] def post_many() -> None: for i in range(10): @@ -49,7 +49,7 @@ 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 = [] + order: list[str] = [] done = loop.create_future() def first() -> None: @@ -71,7 +71,7 @@ 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 = [] + order: list[str] = [] resumed = asyncio.Event() async def awaiting() -> None: @@ -102,9 +102,9 @@ 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 + loop: asyncio.AbstractEventLoop - refs = [] + refs: list[tuple[weakref.ref[asyncio.AbstractEventLoop], weakref.ref[Held]]] = [] for _ in range(5): loop = asyncio.new_event_loop() held = Held() diff --git a/tests/test_track_settings.py b/tests/test_track_settings.py index a0b102b..e0374fa 100644 --- a/tests/test_track_settings.py +++ b/tests/test_track_settings.py @@ -10,6 +10,7 @@ from __future__ import annotations import pytest +from typing_extensions import TypedDict import webrtc from tests.helpers import capture_mode, connect_track, run_isolated, wait_until @@ -17,6 +18,13 @@ C = webrtc.MediaTrackConstraints +class VideoConstraints(TypedDict, total=False, closed=True): + """Constraints taken by both MediaTrackConstraints and get_user_media.""" + + width: int | webrtc.ConstrainULongRange + frame_rate: float | webrtc.ConstrainDoubleRange + + @pytest.mark.asyncio 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.""" @@ -26,12 +34,15 @@ async def test_camera_settings_and_capabilities() -> None: await wait_until(lambda: track.get_settings().frame_rate is not None, 'the frame rate') settings = track.get_settings() assert (settings.width, settings.height, settings.aspect_ratio) == (320, 240, 320 / 240) + assert settings.frame_rate is not None assert abs(settings.frame_rate - 30) < 5 assert settings.device_id == 'synthetic-camera' assert settings.resize_mode == 'none' capabilities = track.get_capabilities() assert capabilities.width == webrtc.ULongRange(1, 4096) + assert capabilities.frame_rate is not None + assert capabilities.frame_rate.max is not None assert capabilities.frame_rate.max >= 60 assert track.get_constraints() == webrtc.MediaTrackConstraints(width=320, height=240, frame_rate=30) track.stop() @@ -58,7 +69,12 @@ async def test_apply_constraints_to_the_camera(video_stream: webrtc.MediaStream) C(width=160, height=webrtc.ConstrainULongRange(exact=120), frame_rate=webrtc.ConstrainDoubleRange(max=10)) ) await wait_until(lambda: track.get_settings().width == 160, 'the new size') - await wait_until(lambda: (track.get_settings().frame_rate or 0) < 12, 'the new frame rate') + + def slowed() -> bool: + frame_rate = track.get_settings().frame_rate + return frame_rate is None or frame_rate < 12 + + await wait_until(slowed, 'the new frame rate') assert track.get_settings().height == 120 assert track.getConstraints().height == webrtc.ConstrainULongRange(exact=120) @@ -104,15 +120,15 @@ async def test_remote_track_settings( 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 not video.content_hint - assert not audio.contentHint + assert video.content_hint == '' + assert 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 not video.content_hint + assert video.content_hint == '' @pytest.mark.asyncio @@ -129,7 +145,7 @@ async def test_constraints_of_an_ended_track(video_stream: webrtc.MediaStream) - [{'frame_rate': float('nan')}, {'frame_rate': float('inf')}, {'width': '1'}], ) async def test_constraints_have_their_webidl_types( - video_stream: webrtc.MediaStream, constraints: dict[str, object] + video_stream: webrtc.MediaStream, constraints: VideoConstraints ) -> None: """Unsigned longs and restricted doubles: other values are a TypeError.""" with pytest.raises(TypeError): @@ -148,7 +164,7 @@ async def test_constraints_have_their_webidl_types( ], ) async def test_camera_stays_within_its_capabilities( - video_stream: webrtc.MediaStream, constraints: dict[str, object], expected: tuple[int, int, float] + video_stream: webrtc.MediaStream, constraints: VideoConstraints, expected: tuple[int, int, float] ) -> None: """Ideal values beyond the capabilities select the nearest ones.""" track = video_stream.get_tracks()[0] diff --git a/tests/test_tracks.py b/tests/test_tracks.py index 2c378d9..b125337 100644 --- a/tests/test_tracks.py +++ b/tests/test_tracks.py @@ -26,7 +26,8 @@ async def test_remote_tracks_have_their_own_id_and_label( caller.add_transceiver(webrtc.MediaType.video) await caller.set_local_description() - tracks = [] + tracks: list[list[webrtc.MediaStreamTrack]] = [] + assert caller.local_description is not None for pc in (callee, callee2): await pc.set_remote_description(caller.local_description) tracks.append([t.receiver.track for t in pc.get_transceivers()]) @@ -77,7 +78,7 @@ 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 = [] + events: list[webrtc.RTCTrackEvent] = [] callee.on('track', events.append) transceiver = caller.add_transceiver(webrtc.MediaType.audio) await exchange_offer_answer(caller, callee) @@ -119,7 +120,7 @@ async def test_remote_track_mute_and_stream_events( stream = webrtc.MediaStream([audio, video]) caller.add_track(audio, stream) transceiver = caller.add_transceiver(video, webrtc.RTCRtpTransceiverInit(streams=[stream])) - events = [] + events: list[webrtc.RTCTrackEvent] = [] callee.on('track', events.append) await connect(caller, callee) @@ -134,6 +135,8 @@ async def test_remote_track_mute_and_stream_events( transceiver.direction = webrtc.TransceiverDirection.inactive await exchange_offer(caller, callee) - assert (await removed).track == remote_video + removed_event = await removed + assert isinstance(removed_event, webrtc.MediaStreamTrackEvent) + assert removed_event.track == remote_video await muted assert remote_video.muted diff --git a/tests/test_video.py b/tests/test_video.py index e3b96eb..b8e7af7 100644 --- a/tests/test_video.py +++ b/tests/test_video.py @@ -9,6 +9,8 @@ from __future__ import annotations +import functools + import pytest import webrtc @@ -31,18 +33,26 @@ def test_get_user_media_needs_audio_or_video() -> None: @pytest.mark.parametrize( - ('constraints', 'error'), + ('get_user_media', 'error'), [ - ({'width': webrtc.ConstrainULongRange(exact=0)}, webrtc.OverconstrainedError), - ({'frame_rate': webrtc.ConstrainDoubleRange(max=0)}, webrtc.OverconstrainedError), - ({'width': webrtc.ConstrainULongRange(min=0, max=-1)}, TypeError), + ( + functools.partial(webrtc.get_user_media, width=webrtc.ConstrainULongRange(exact=0)), + webrtc.OverconstrainedError, + ), + ( + functools.partial(webrtc.get_user_media, frame_rate=webrtc.ConstrainDoubleRange(max=0)), + webrtc.OverconstrainedError, + ), + (functools.partial(webrtc.get_user_media, width=webrtc.ConstrainULongRange(min=0, max=-1)), TypeError), ], ids=['exact', 'max', 'negative'], ) -def test_get_user_media_constraint_beyond_the_camera(constraints: dict[str, object], error: type[Exception]) -> None: +def test_get_user_media_constraint_beyond_the_camera( + get_user_media: functools.partial[webrtc.MediaStream], 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) + get_user_media(audio=False, video=True) def test_get_user_media_ideal_beyond_the_camera() -> None: @@ -85,6 +95,7 @@ async def test_remote_video_track( track_event = wait_for_event(callee, 'track') await connect(caller, callee) event = await track_event + assert isinstance(event, webrtc.RTCTrackEvent) assert event.track.kind == webrtc.MediaType.video assert [s.id for s in event.streams] == [video_stream.id] diff --git a/tests/test_video_frame.py b/tests/test_video_frame.py index a491935..1ce7ff0 100644 --- a/tests/test_video_frame.py +++ b/tests/test_video_frame.py @@ -14,18 +14,28 @@ import struct import pytest +from typing_extensions import TypedDict, Unpack import webrtc +from tests.helpers import mistyped from webrtc import PlaneLayout, VideoPixelFormat # a 4x2 I420 frame: 8 samples of Y, 2 of U, 2 of V I420_DATA = bytes(range(1, 13)) -def i420_4x2(data: bytes = I420_DATA, **init: object) -> webrtc.VideoFrame: +class I420Init(TypedDict, total=False, closed=True): + duration: int | None + layout: list[PlaneLayout] | None + visible_rect: webrtc.DOMRectInit | None + display_width: int | None + display_height: int | None + + +def i420_4x2(data: bytes | bytearray = I420_DATA, *, timestamp: int = 0, **init: Unpack[I420Init]) -> webrtc.VideoFrame: return webrtc.VideoFrame( data, - webrtc.VideoFrameBufferInit(**{'format': 'I420', 'coded_width': 4, 'coded_height': 2, 'timestamp': 0, **init}), + webrtc.VideoFrameBufferInit(format='I420', coded_width=4, coded_height=2, timestamp=timestamp, **init), ) @@ -69,8 +79,10 @@ def test_init_from_json() -> None: def test_buffer_needs_an_init() -> None: """A frame of a buffer has no defaults for its format and size.""" + # a buffer, where only a frame may come without an init + buffer: webrtc.VideoFrame = mistyped(I420_DATA) with pytest.raises(TypeError, match='needs a VideoFrameBufferInit'): - webrtc.VideoFrame(I420_DATA) + webrtc.VideoFrame(buffer) @pytest.mark.parametrize( @@ -161,7 +173,7 @@ async def test_copy_to_errors() -> None: @pytest.mark.asyncio @pytest.mark.parametrize('format', ['RGBA', 'RGBX', 'BGRA', 'BGRX']) -async def test_convert_i420_to_rgb(format: str) -> None: +async def test_convert_i420_to_rgb(format: webrtc.VideoPixelFormatValue) -> 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) @@ -193,7 +205,7 @@ async def test_rgb_formats_swap_and_alpha() -> None: bytes([1, 2, 3, 4] * 4), webrtc.VideoFrameBufferInit(format='RGBA', coded_width=2, coded_height=2, timestamp=0) ) assert frame.color_space.matrix == 'rgb' - assert frame.color_space.full_range + assert frame.color_space.full_range is True out = bytearray(16) await frame.copy_to(out, webrtc.VideoFrameCopyToOptions(format='BGRA')) assert list(out[:4]) == [3, 2, 1, 4] @@ -214,7 +226,7 @@ async def test_rgb_formats_swap_and_alpha() -> None: ('NV12', 12), ], ) -async def test_other_formats_round_trip(format: str, size: int) -> None: +async def test_other_formats_round_trip(format: webrtc.VideoPixelFormatValue, 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( @@ -244,6 +256,7 @@ 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, webrtc.VideoFrameInit(visible_rect=webrtc.DOMRectInit(x=2, y=0, width=2, height=2))) + assert crop.visible_rect is not None assert (crop.coded_width, crop.visible_rect.x, crop.visible_rect.width) == (4, 2, 2) assert (crop.display_width, crop.display_height) == (4, 2) assert crop.timestamp == 1234 @@ -309,7 +322,8 @@ async def test_close_and_clone() -> None: frame.clone() with pytest.raises(webrtc.InvalidStateError): webrtc.VideoFrame(frame) - assert clone.format == VideoPixelFormat.I420 + clone_format = clone.format + assert clone_format == VideoPixelFormat.I420 with clone: assert clone.allocation_size() == 12 assert clone.format is None diff --git a/tests/wpt/__main__.py b/tests/wpt/__main__.py index 9b11ca8..8d4dc5b 100644 --- a/tests/wpt/__main__.py +++ b/tests/wpt/__main__.py @@ -30,13 +30,13 @@ def _print_result(case: str, result: runner.CaseResult) -> None: harness = result['harness'] - message = f' ({harness["message"]})' if harness['message'] else '' + message = f' ({harness["message"]})' if harness['message'] not in {None, ''} else '' logger.info('%s: harness %s%s', case, harness['status'], message) for test in result['tests']: logger.info(' %-8s %s', test['status'], test['name']) - if test['status'] != 'PASS' and test['message']: + if test['status'] != 'PASS' and test['message'] not in {None, ''}: logger.info(' %s', test['message']) - if result['unsupported']: + if len(result['unsupported']) > 0: logger.info(' unsupported: %s', ', '.join(result['unsupported'])) @@ -47,7 +47,7 @@ def run(args: argparse.Namespace) -> None: def update(args: argparse.Namespace) -> None: expectations = Expectations.load() - cases = [c for c in args.cases or discover() if not expectations.skip_reason(c)] + cases = [c for c in args.cases or discover() if expectations.skip_reason(c) in {None, ''}] tests: Counter[str] = Counter() harness: Counter[str] = Counter() @@ -70,7 +70,7 @@ def run_repeatedly(case: str) -> list[runner.CaseResult]: logger.info('\nfiles: %s', dict(harness)) logger.info('tests: %s', dict(tests)) - if unsupported: + if len(unsupported) > 0: logger.info('unsupported members used by tests:') for name, count in unsupported.most_common(): logger.info(' %4d %s', count, name) diff --git a/tests/wpt/bridge.py b/tests/wpt/bridge.py index 3151da6..c34a89e 100644 --- a/tests/wpt/bridge.py +++ b/tests/wpt/bridge.py @@ -15,17 +15,33 @@ import asyncio import dataclasses import enum +import inspect import sys import time -from typing import TYPE_CHECKING, Callable, Union +from typing import TYPE_CHECKING, Callable, TypeVar, Union, cast import pythonmonkey as pm +from typing_extensions import TypedDict import webrtc import webrtc.enums if TYPE_CHECKING: - from collections.abc import Coroutine + from collections.abc import Awaitable, Coroutine + + from _typeshed import DataclassInstance + + from webrtc.models.media_track_constraints import ConstrainDouble, ConstrainULong + + class _UserMediaOptions(TypedDict, total=False, closed=True): + audio: bool + video: bool + width: ConstrainULong | None + height: ConstrainULong | None + frame_rate: ConstrainDouble | None + + +_T = TypeVar('_T') Result = dict[str, object] Buffer = Union[bytes, bytearray, memoryview] @@ -60,33 +76,74 @@ def _event_to_js(event: webrtc.Event) -> dict[str, object]: return {'__event': type(event).__name__, 'type': event.type, 'init': init} -def _dictionary_to_js(value: object) -> dict[str, object]: +def _dictionary_to_js(value: DataclassInstance) -> 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]]] = [ +def _enum_to_js(value: enum.Enum) -> 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}), + return value.value + + +def _rect_to_js(value: webrtc.DOMRectReadOnly) -> object: + return {'__rect': [value.x, value.y, value.width, value.height]} + + +def _blob_to_js(value: webrtc.Blob) -> object: + return {'__blob': bytearray(bytes(value)), 'type': value.type} + + +def _description_to_js(value: webrtc.RTCSessionDescriptionInit) -> object: # a dictionary in WebIDL, but a WebRTCObject (it holds a native one) here, not a dataclass - (webrtc.RTCSessionDescriptionInit, lambda value: value.to_json()), + return value.to_json() + + +def _wrapper_to_js(value: webrtc.WebRTCObject[object]) -> object: # 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()]}, - ), + return {'__type': type(value).__name__, '__id': hash(value), '__obj': value} + + +def _plain_to_js(value: object) -> object: + return {'__type': type(value).__name__, '__id': id(value), '__obj': value} + + +def _error_to_js(value: BaseException) -> object: + return {'__error': _error(value)['error']} + + +def _stats_to_js(value: webrtc.RTCStatsReport) -> object: + return {'__statsReport': [[stats_id, dict(stats)] for stats_id, stats in value.items()]} + + +def _bytes_to_js(value: bytes) -> object: # PythonMonkey shares a bytearray as a Uint8Array, the shim copies it - (bytes, lambda value: {'__bytes': bytearray(value)}), + return {'__bytes': bytearray(value)} + + +def _dict_to_js(value: dict[object, object]) -> object: + return {k: to_js(v) for k, v in value.items()} + + +def _sequence_to_js(value: list[object] | tuple[object, ...]) -> object: + return [to_js(v) for v in value] + + +# in order: the first type that matches converts the value +_CONVERTERS: list[tuple[type | tuple[type, ...], Callable[..., object]]] = [ + (enum.Enum, _enum_to_js), + (webrtc.DOMRectReadOnly, _rect_to_js), + (webrtc.Blob, _blob_to_js), + (webrtc.RTCSessionDescriptionInit, _description_to_js), + (webrtc.WebRTCObject, _wrapper_to_js), + (_PLAIN_INTERFACES, _plain_to_js), + (BaseException, _error_to_js), + (webrtc.RTCStatsReport, _stats_to_js), + (bytes, _bytes_to_js), (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]), + (dict, _dict_to_js), + ((list, tuple), _sequence_to_js), ] @@ -106,7 +163,7 @@ def _to_enum(value: dict[str, object]) -> object: try: return enum_cls(value['value']) except ValueError: - if value.get('strict', True): + if bool(value.get('strict', True)): 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 @@ -173,7 +230,7 @@ def _guard(func: Callable[[], object]) -> Result: return _error(e) -async def _guard_async(awaitable: Callable[[], Coroutine[object, object, object]]) -> Result: +async def _guard_async(awaitable: Callable[[], Awaitable[object]]) -> Result: try: return {'ok': to_js(await awaitable())} except BRIDGED_ERRORS as e: @@ -196,10 +253,33 @@ def set_attr(obj: object, name: str, value: object) -> Result: return _guard(lambda: setattr(obj, name, from_js(value))) +def _from_js_as(kind: type[_T], value: object) -> _T: + """A value from the shim that converts to a type, like a dictionary the library has a model of.""" + result = from_js(value) + if not isinstance(result, kind): + msg = f'expected {kind.__name__} from JS, got {type(result).__name__}' + raise TypeError(msg) + return result + + 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 {}))) + if not isinstance(args, list): + msg = f'args from JS is not an array: {args!r}' + raise TypeError(msg) + if kwargs is not None and not isinstance(kwargs, dict): + msg = f'kwargs from JS is not an object: {kwargs!r}' + raise TypeError(msg) + keywords: dict[object, object] = dict(kwargs) if kwargs is not None else {} + return getattr(obj, name)(*from_js(list(args)), **from_js(keywords)) + + +def _awaitable(value: object) -> Awaitable[object]: + if not inspect.isawaitable(value): + msg = f'{type(value).__name__} is not awaitable' + raise TypeError(msg) + return value def call_method(obj: object, name: str, arguments: dict[str, object]) -> Result: @@ -207,7 +287,7 @@ def call_method(obj: object, name: str, arguments: dict[str, object]) -> Result: def call_async_method(obj: object, name: str, arguments: dict[str, object]) -> asyncio.Future[Result]: - return _start(_guard_async(lambda: _call(obj, name, arguments))) + return _start(_guard_async(lambda: _awaitable(_call(obj, name, arguments)))) def await_attr(obj: object, name: str) -> asyncio.Future[Result]: @@ -224,7 +304,7 @@ def video_frame_copy_to(frame: webrtc.VideoFrame, destination: Buffer, options: async def copy() -> dict[str, object]: data = bytearray(destination) - layout = await frame.copy_to(data, from_js(options)) + layout = await frame.copy_to(data, _from_js_as(webrtc.VideoFrameCopyToOptions, options)) return {'layout': layout, 'data': bytes(data)} return _start(_guard_async(copy)) @@ -235,7 +315,7 @@ def audio_data_copy_to(audio: webrtc.AudioData, destination: Buffer, options: ob def copy() -> bytes: data = bytearray(destination) - audio.copy_to(data, from_js(options)) + audio.copy_to(data, _from_js_as(webrtc.AudioDataCopyToOptions, options)) return bytes(data) return _guard(copy) @@ -246,7 +326,8 @@ def construct(name: str, kwargs: dict[str, object]) -> Result: def get_user_media(kwargs: dict[str, object]) -> Result: - return _guard(lambda: webrtc.get_user_media(**from_js(dict(kwargs)))) + # the shim converts the constraints as WebIDL does, the library validates them + return _guard(lambda: webrtc.get_user_media(**cast('_UserMediaOptions', from_js(dict(kwargs))))) def call_static(class_name: str, name: str, args: list[object]) -> Result: diff --git a/tests/wpt/child.py b/tests/wpt/child.py index b681be5..fb2659e 100644 --- a/tests/wpt/child.py +++ b/tests/wpt/child.py @@ -15,6 +15,7 @@ import sys import pythonmonkey as pm +from typing_extensions import TypedDict from tests.wpt import bridge from tests.wpt.loader import build_scripts, load, split_case @@ -31,6 +32,23 @@ logger = logging.getLogger(__name__) +# what testharness reports, as JS objects: statuses are numbers, messages may be null or undefined +class _JsHarness(TypedDict): + status: float + message: object + + +class _JsTest(TypedDict): + name: str + status: float + message: object + + +class _JsResult(TypedDict): + harness: _JsHarness + tests: list[_JsTest] + + def _text(value: object) -> str | None: # JS null and undefined arrive as PythonMonkey objects return value if isinstance(value, str) else None @@ -53,10 +71,10 @@ async def run_in_process(case: str) -> CaseResult: test_file = load(path) loop = asyncio.get_running_loop() - completed = loop.create_future() + completed: asyncio.Future[_JsResult] = loop.create_future() unsupported: set[str] = set() - def complete(result: dict) -> None: + def complete(result: _JsResult) -> None: if not completed.done(): completed.set_result(result) diff --git a/tests/wpt/expectations.py b/tests/wpt/expectations.py index a7315d5..4be9674 100644 --- a/tests/wpt/expectations.py +++ b/tests/wpt/expectations.py @@ -23,7 +23,9 @@ import json from dataclasses import dataclass, field from pathlib import Path -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, Union + +from typing_extensions import TypedDict if TYPE_CHECKING: from tests.wpt.runner import CaseResult @@ -32,6 +34,14 @@ HARNESS_KEY = '[harness]' +# one status, or the statuses a flaky test may have +Expected = Union[str, list[str]] + + +class _File(TypedDict, total=False): + skip: dict[str, str] + results: dict[str, dict[str, Expected]] + def _allowed(expected: str | list[str]) -> list[str]: return expected if isinstance(expected, list) else [expected] @@ -50,7 +60,7 @@ class Expectations: def load(cls) -> Expectations: if not PATH.exists(): return cls() - data = json.loads(PATH.read_text()) + data: _File = json.loads(PATH.read_text()) return cls(skip=data.get('skip', {}), results=data.get('results', {})) def save(self) -> None: @@ -73,7 +83,7 @@ def record(self, case: str, results: list[CaseResult]) -> None: seen.setdefault(test['name'], set()).add(test['status']) previous = self.results.get(case, {}) - entry = {} + entry: dict[str, Expected] = {} for name, statuses in seen.items(): kept = previous.get(name) if isinstance(kept, list) and statuses <= set(kept): @@ -83,7 +93,7 @@ def record(self, case: str, results: list[CaseResult]) -> None: elif statuses.isdisjoint({'PASS', 'OK'}): entry[name] = statuses.pop() - if entry: + if len(entry) > 0: self.results[case] = dict(sorted(entry.items())) else: self.results.pop(case, None) @@ -92,7 +102,7 @@ def mismatches(self, case: str, result: CaseResult) -> list[str]: expected = dict(self.results.get(case, {})) expected_harness = expected.pop(HARNESS_KEY, 'OK') - problems = [] + problems: list[str] = [] harness = result['harness'] if harness['status'] not in _allowed(expected_harness): problems.append( diff --git a/tests/wpt/loader.py b/tests/wpt/loader.py index 79c2d23..0cc5e96 100644 --- a/tests/wpt/loader.py +++ b/tests/wpt/loader.py @@ -15,6 +15,8 @@ from pathlib import Path from urllib.parse import urljoin, urlsplit +from typing_extensions import override + WPT_ROOT = Path(__file__).resolve().parents[2] / 'wpt' # Platform globals the shell lacks, then the WebRTC API POLYFILLS = Path(__file__).with_name('polyfills.js') @@ -77,28 +79,33 @@ def __init__(self, test_file: TestFile) -> None: self._inline: list[str] | None = None self._in_title = False + @override 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'])) + src = attrs.get('src') + if src is not None and src != '': + self.test_file.scripts.append(('src', src)) else: self._inline = [] elif tag == 'meta' and attrs.get('name') == 'variant': - self.test_file.variants.append(attrs.get('content') or '') + content = attrs.get('content') + self.test_file.variants.append(content if content is not None else '') elif tag == 'meta' and attrs.get('name') == 'timeout': self.test_file.long_timeout = attrs.get('content') == 'long' elif tag == 'title': self._in_title = True + @override 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 + @override 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))) @@ -118,7 +125,7 @@ def _load_js(path: Path) -> TestFile: test_file = TestFile(path, scripts=[('src', HARNESS)]) for line in path.read_text(encoding='utf-8').splitlines(): match = _META.match(line.strip()) - if not match: + if match is None: continue key, value = match.group(1), match.group(2).strip() if key == 'script': @@ -135,14 +142,16 @@ def _load_js(path: Path) -> TestFile: def load(path: Path) -> TestFile: test_file = _load_js(path) if path.name.endswith('.js') else _load_html(path) - test_file.variants = test_file.variants or [''] - test_file.title = test_file.title.strip() or path.name + if len(test_file.variants) == 0: + test_file.variants = [''] + title = test_file.title.strip() + test_file.title = title if title != '' else path.name return test_file def _is_test(path: Path) -> bool: parts = path.relative_to(WPT_ROOT).parts - if HELPER_DIRS.intersection(parts[:-1]): + if len(HELPER_DIRS.intersection(parts[:-1])) > 0: return False name = path.name if '-manual.' in name or name.endswith(('-ref.html', '-notref.html')): @@ -153,7 +162,7 @@ def _is_test(path: Path) -> bool: def discover() -> list[str]: """Returns every test case as a path relative to the WPT root, followed by its variant, if any.""" - cases = [] + cases: list[str] = [] for test_dir in TEST_DIRS: paths = WPT_ROOT.glob(test_dir) if '*' in test_dir else (WPT_ROOT / test_dir).rglob('*') for path in sorted(paths): @@ -168,7 +177,7 @@ def case_id(path: Path, variant: str = '') -> str: def split_case(case: str) -> tuple[Path, str]: path, _, variant = case.partition('?') - return WPT_ROOT / path, '?' + variant if variant else '' + return WPT_ROOT / path, '?' + variant if variant != '' else '' def _resolve(src: str, test_path: Path) -> Path: diff --git a/tests/wpt/runner.py b/tests/wpt/runner.py index 20a743d..95ecb10 100644 --- a/tests/wpt/runner.py +++ b/tests/wpt/runner.py @@ -79,5 +79,6 @@ def run(case: str) -> CaseResult: for line in proc.stdout.splitlines(): if line.startswith(RESULT_PREFIX): - return json.loads(line[len(RESULT_PREFIX) :]) + result: CaseResult = json.loads(line[len(RESULT_PREFIX) :]) + return result return harness_result('CRASH', f'exit code {proc.returncode}\n{proc.stderr[-2000:]}') diff --git a/tests/wpt/test_wpt.py b/tests/wpt/test_wpt.py index b7864e1..ae440a0 100644 --- a/tests/wpt/test_wpt.py +++ b/tests/wpt/test_wpt.py @@ -36,7 +36,7 @@ def _cases() -> list[str | ParameterSet]: return [] return [ pytest.param(case, marks=pytest.mark.skip(reason=reason)) - if (reason := expectations.skip_reason(case)) + if (reason := expectations.skip_reason(case)) is not None and reason != '' else case for case in discover() ] @@ -45,5 +45,5 @@ def _cases() -> list[str | ParameterSet]: @pytest.mark.parametrize('case', _cases()) def test_wpt(case: str) -> None: problems = expectations.mismatches(case, runner.run(case)) - if problems: + if len(problems) > 0: pytest.fail('\n'.join(problems), pytrace=False) diff --git a/uv.lock b/uv.lock index df0f224..ec5de1c 100644 --- a/uv.lock +++ b/uv.lock @@ -1512,6 +1512,9 @@ wheels = [ name = "wrtc" version = "0.0.0.dev10" source = { editable = "." } +dependencies = [ + { name = "typing-extensions" }, +] [package.dev-dependencies] dev = [ @@ -1539,6 +1542,7 @@ wpt = [ ] [package.metadata] +requires-dist = [{ name = "typing-extensions", specifier = ">=4.10" }] [package.metadata.requires-dev] dev = [