diff --git a/README.md b/README.md index 934f9d4..3f59a76 100644 --- a/README.md +++ b/README.md @@ -190,5 +190,25 @@ the desktop itself continues to control only the Daemon. Applications that intentionally implement their own media protocol can use the advanced `ApplicationChannels` Device channel and consume raw WSPK frames. +## Live speaker PCM + +Firmware advertising `audio.stream.live.v1` accepts an unlimited-duration, +backpressured PCM stream. Each `write()` blocks when the device's fixed-size +playback queue has no credit, so the SDK never grows an unbounded host buffer: + +```python +with app.robot.audio.open_stream( + sample_rate_hz=24000, + channels=1, + sample_width_bytes=2, +) as speaker: + speaker.write(pcm_chunk) +``` + +Normal context exit drains queued PCM and sends the WSPK `LAST` marker; +exceptional exit calls `abort()` and immediately stops playback. The existing +`play_file()` and `play_pcm()` APIs remain bounded file transfers with size and +SHA-256 validation. + See [examples](examples/README.md), [Runtime contract](docs/contracts/runtime-profile-index.md), the [microphone contract](docs/microphone-audio.md), and [troubleshooting](docs/troubleshooting.md). diff --git a/docs/microphone-audio.md b/docs/microphone-audio.md index 54239e1..e3a647f 100644 --- a/docs/microphone-audio.md +++ b/docs/microphone-audio.md @@ -30,6 +30,11 @@ recording.save("recording.wav") ## Responsibility boundary +The opposite direction uses `robot.audio.open_stream()` for 24 kHz mono +PCM16. Robot-microphone upload and robot-speaker playback are real-time but +half-duplex: opening either direction first closes the other. Camera preview +does not participate in this arbitration. + The Runtime/Daemon owns the physical device WebSocket, WSPK framing, connection lifecycle, and source-aware routing. It must keep device media payloads opaque when forwarding them to an Application; it does not decode, diff --git a/src/watcherobot/__init__.py b/src/watcherobot/__init__.py index adbf7f8..23923a1 100644 --- a/src/watcherobot/__init__.py +++ b/src/watcherobot/__init__.py @@ -6,7 +6,7 @@ JobFailedError, WatcheRobotError, ) -from .audio import AudioPlayback, PCMAudio +from .audio import AudioLiveStream, AudioPlayback, PCMAudio from .job import Job, JobState from .inputs import BackTouchEvent, InputDomain, InputEvent, RollerEvent, ScreenTouchEvent from .media import AudioFormat, AudioFrame, AudioRecording, ImageFrame, MicrophoneSession @@ -39,6 +39,7 @@ __all__ = [ "AudioFormat", "AudioFrame", + "AudioLiveStream", "AudioPlayback", "AudioRecording", "AuthenticationError", diff --git a/src/watcherobot/application/transport.py b/src/watcherobot/application/transport.py index 8b95cad..7e6d76f 100644 --- a/src/watcherobot/application/transport.py +++ b/src/watcherobot/application/transport.py @@ -62,6 +62,7 @@ def __init__(self, *, command_timeout: float = 5.0) -> None: self._audio_credits = 0 self._audio_slots_per_packet = 1 self._audio_flow_error: str | None = None + self._audio_live_first_frame = True def set_callbacks( self, @@ -143,6 +144,35 @@ def send_audio_stream( self._send_audio_stream(bytes(pcm), stream_id, chunk_bytes) ) + def begin_live_audio_stream(self, *, stream_id: int, chunk_bytes: int = 960) -> None: + self._submit(self._begin_live_audio_stream(stream_id, chunk_bytes)).result( + timeout=self.command_timeout + 1.0 + ) + + def write_live_audio_stream( + self, + pcm: bytes, + *, + stream_id: int, + sequence: int, + chunk_bytes: int = 960, + ) -> int: + return self._submit( + self._write_live_audio_stream( + bytes(pcm), stream_id, sequence, chunk_bytes + ) + ).result(timeout=self.command_timeout + 1.0) + + def end_live_audio_stream(self, *, stream_id: int, sequence: int) -> None: + self._submit(self._end_live_audio_stream(stream_id, sequence)).result( + timeout=self.command_timeout + 1.0 + ) + + def cancel_live_audio_stream(self, *, stream_id: int) -> None: + self._submit(self._cancel_live_audio_stream(stream_id)).result( + timeout=self.command_timeout + 1.0 + ) + def send_desktop(self, frame: str | bytes) -> Future[None]: return self._submit(self._send(ApplicationChannel.DESKTOP, frame)) @@ -207,6 +237,7 @@ async def _run(self) -> None: self._disconnect_callback() communicator_task.cancel() stop_task.cancel() + await self._fail_audio_flow("disconnected") await asyncio.gather( communicator_task, stop_task, @@ -214,6 +245,14 @@ async def _run(self) -> None: ) self._communicators = None + async def _fail_audio_flow(self, reason: str) -> None: + condition = self._audio_credit_condition + if condition is None: + return + async with condition: + self._audio_flow_error = reason + condition.notify_all() + async def _send( self, channel: ApplicationChannel, @@ -296,6 +335,89 @@ async def _send_audio_stream( self._audio_credits = 0 self._audio_slots_per_packet = 1 + async def _begin_live_audio_stream( + self, + stream_id: int, + chunk_bytes: int = 960, + ) -> None: + if chunk_bytes <= 0 or chunk_bytes > 4096 or chunk_bytes % 2 != 0: + raise ValueError("chunk_bytes must be an even value between 2 and 4096") + if self._audio_credit_condition is None: + self._audio_credit_condition = asyncio.Condition() + async with self._audio_credit_condition: + self._audio_flow_stream_id = stream_id + self._audio_slots_per_packet = max( + 1, + (chunk_bytes + AUDIO_DEVICE_SLOT_BYTES - 1) + // AUDIO_DEVICE_SLOT_BYTES, + ) + self._audio_credits = 4 + self._audio_flow_error = None + self._audio_live_first_frame = True + self._audio_credit_condition.notify_all() + + async def _write_live_audio_stream( + self, + pcm: bytes, + stream_id: int, + sequence: int, + chunk_bytes: int = 960, + ) -> int: + if not pcm: + return sequence + if len(pcm) % 2 != 0: + raise ValueError("PCM chunk ends with a partial sample") + for offset in range(0, len(pcm), chunk_bytes): + await self._take_audio_credit(stream_id) + payload = pcm[offset : offset + chunk_bytes] + await self._send( + ApplicationChannel.DEVICE, + build_wspk( + FRAME_AUDIO, + FLAG_FIRST if self._audio_live_first_frame else 0, + stream_id, + sequence, + payload, + ), + ) + self._audio_live_first_frame = False + sequence = (sequence + 1) & 0xFFFFFFFF + return sequence + + async def _end_live_audio_stream(self, stream_id: int, sequence: int) -> None: + condition = self._audio_credit_condition + if condition is None: + raise WatcheRobotError("audio flow control is not initialized") + try: + await self._send( + ApplicationChannel.DEVICE, + build_wspk( + FRAME_AUDIO, + FLAG_LAST, + stream_id, + sequence, + b"", + ), + ) + finally: + async with condition: + if self._audio_flow_stream_id == stream_id: + self._audio_flow_stream_id = 0 + self._audio_credits = 0 + self._audio_slots_per_packet = 1 + condition.notify_all() + + async def _cancel_live_audio_stream(self, stream_id: int) -> None: + condition = self._audio_credit_condition + if condition is None: + return + async with condition: + if self._audio_flow_stream_id == stream_id: + self._audio_flow_stream_id = 0 + self._audio_credits = 0 + self._audio_flow_error = "cancelled" + condition.notify_all() + async def _take_audio_credit(self, stream_id: int) -> None: condition = self._audio_credit_condition if condition is None: @@ -308,12 +430,12 @@ async def wait_for_credit() -> None: or self._audio_credits > 0 or self._audio_flow_error is not None ) - if self._audio_flow_stream_id != stream_id: - raise WatcheRobotError("audio stream was replaced") if self._audio_flow_error is not None: raise WatcheRobotError( f"audio stream failed: {self._audio_flow_error}" ) + if self._audio_flow_stream_id != stream_id: + raise WatcheRobotError("audio stream was replaced") self._audio_credits -= 1 try: diff --git a/src/watcherobot/audio.py b/src/watcherobot/audio.py index e4b8cb3..4ed3284 100644 --- a/src/watcherobot/audio.py +++ b/src/watcherobot/audio.py @@ -1,11 +1,14 @@ from __future__ import annotations import hashlib +import threading import wave from dataclasses import dataclass from pathlib import Path +from types import TracebackType from typing import Callable +from .errors import WatcheRobotError from .job import CommandTransport, Job, JobState from .media import AudioFormat @@ -58,6 +61,88 @@ def cancel(self) -> None: self._cancel_callback(self) +class AudioLiveStream: + """Synchronous, backpressured PCM stream from the host to the robot.""" + + def __init__( + self, + stream_id: int, + write_callback: Callable[[AudioLiveStream, bytes, int], int], + close_callback: Callable[[AudioLiveStream, int], None], + abort_callback: Callable[[AudioLiveStream], None], + ) -> None: + self.stream_id = stream_id + self._write_callback = write_callback + self._close_callback = close_callback + self._abort_callback = abort_callback + self._sequence = 0 + self._closed = False + self._lock = threading.RLock() + self._write_lock = threading.Lock() + + @property + def closed(self) -> bool: + with self._lock: + return self._closed + + def write(self, pcm_chunk: bytes) -> None: + payload = bytes(pcm_chunk) + if len(payload) % OUTPUT_AUDIO_FORMAT.sample_width_bytes != 0: + raise ValueError("PCM chunk ends with a partial sample") + with self._write_lock: + with self._lock: + if self._closed: + raise WatcheRobotError("audio live stream is closed") + if not payload: + return + sequence = self._sequence + next_sequence = self._write_callback(self, payload, sequence) + with self._lock: + if not self._closed: + self._sequence = next_sequence + + def close(self) -> None: + with self._write_lock: + with self._lock: + if self._closed: + return + self._closed = True + sequence = self._sequence + try: + self._close_callback(self, sequence) + except Exception: + try: + self._abort_callback(self) + except Exception: + pass + raise + + def abort(self) -> None: + with self._lock: + if self._closed: + return + self._closed = True + self._abort_callback(self) + + def _mark_closed(self) -> None: + with self._lock: + self._closed = True + + def __enter__(self) -> AudioLiveStream: + return self + + def __exit__( + self, + exc_type: type[BaseException] | None, + exc_value: BaseException | None, + traceback: TracebackType | None, + ) -> None: + if exc_type is None: + self.close() + else: + self.abort() + + def load_pcm_wave(path: str | Path) -> PCMAudio: """Read a WAV file in the single playback format supported by protocol v1.""" source = Path(path) diff --git a/src/watcherobot/robot.py b/src/watcherobot/robot.py index 3bcb9a1..edae67d 100644 --- a/src/watcherobot/robot.py +++ b/src/watcherobot/robot.py @@ -12,7 +12,7 @@ from ._internal.audio_status import AudioStatusKind, classify_audio_status from .errors import CommandError, WatcheRobotError -from .audio import AudioPlayback, PCMAudio, load_pcm_wave +from .audio import AudioLiveStream, AudioPlayback, PCMAudio, load_pcm_wave from .job import Job, JobState from .inputs import InputDomain, parse_input_event from .media import AudioFormat, AudioRecording, ImageFrame, MicrophoneSession @@ -140,6 +140,28 @@ def play_pcm( ) ) + def open_stream( + self, + *, + sample_rate_hz: int = 24000, + channels: int = 1, + sample_width_bytes: int = 2, + ) -> AudioLiveStream: + audio_format = AudioFormat( + sample_rate_hz=sample_rate_hz, + channels=channels, + sample_width_bytes=sample_width_bytes, + encoding="pcm_s16le", + ) + if audio_format != AudioFormat( + sample_rate_hz=24000, + channels=1, + sample_width_bytes=2, + encoding="pcm_s16le", + ): + raise ValueError("live playback requires PCM S16LE, 24000 Hz, mono") + return self._robot._open_live_audio_stream() + def stop(self) -> None: self._robot._stop_audio_playback() @@ -302,6 +324,7 @@ def __init__( self._audio_playback_lock = threading.Lock() self._audio_api_lock = threading.Lock() self._audio_playback: AudioPlayback | None = None + self._live_audio_stream: AudioLiveStream | None = None self._audio_send_future: Any | None = None self._audio_cleanup_future: Any | None = None self._audio_cleanup_required = False @@ -353,6 +376,13 @@ def close(self) -> None: microphone.close() except Exception: pass + with self._audio_playback_lock: + live_stream = self._live_audio_stream + if live_stream is not None and not live_stream.closed: + try: + live_stream.abort() + except Exception: + pass with self._face_tracking_lock: preview = self._face_tracking_preview if preview is not None and not preview.closed: @@ -425,6 +455,8 @@ def _cancel_audio_sender(self) -> None: send_future.cancel() def _start_local_audio(self, sound_id: str) -> Job: + self._close_microphone_for_speaker() + self._abort_active_live_audio() with self._audio_api_lock: self._begin_audio_transition() try: @@ -439,6 +471,8 @@ def _start_local_audio(self, sound_id: str) -> Job: def _start_audio_playback(self, audio: PCMAudio) -> AudioPlayback: if "audio.stream" not in self.capabilities: raise WatcheRobotError("robot firmware does not advertise audio.stream") + self._close_microphone_for_speaker() + self._abort_active_live_audio() with self._audio_api_lock: self._begin_audio_transition() with self._audio_playback_lock: @@ -491,6 +525,121 @@ def finish_send(future: Future[None]) -> None: send_future.add_done_callback(finish_send) return playback + def _allocate_audio_stream_id(self) -> int: + with self._audio_playback_lock: + stream_id = self._next_audio_stream_id + self._next_audio_stream_id = 1 if stream_id >= 0xFFFF else stream_id + 1 + return stream_id + + def _close_microphone_for_speaker(self) -> None: + with self._media_lock: + microphone = self._microphone + if microphone is not None and not microphone.closed: + microphone.close() + + def _abort_active_live_audio(self) -> bool: + with self._audio_playback_lock: + stream = self._live_audio_stream + if stream is not None and not stream.closed: + stream.abort() + return True + return False + + def _open_live_audio_stream(self) -> AudioLiveStream: + if "audio.stream.live.v1" not in self.capabilities: + raise WatcheRobotError( + "robot firmware does not advertise audio.stream.live.v1; upgrade firmware" + ) + self._close_microphone_for_speaker() + self._abort_active_live_audio() + stream_id = self._allocate_audio_stream_id() + with self._audio_api_lock: + self._begin_audio_transition() + device_started = False + try: + self._command( + "ctrl.audio.stream.begin", + { + "mode": "live", + "stream_id": stream_id, + "sample_rate_hz": 24000, + "channels": 1, + "sample_width_bytes": 2, + }, + ) + device_started = True + self._transport.begin_live_audio_stream( + stream_id=stream_id, + chunk_bytes=960, + ) + except Exception: + if device_started: + try: + self._transport.cancel_live_audio_stream(stream_id=stream_id) + except Exception: + pass + try: + self._command("ctrl.audio.stop", {}) + except Exception: + pass + self._end_audio_transition(command_succeeded=False) + raise + self._replace_audio_playback() + stream = AudioLiveStream( + stream_id, + self._write_live_audio_stream, + self._close_live_audio_stream, + self._abort_live_audio_stream, + ) + with self._audio_playback_lock: + self._live_audio_stream = stream + self._end_audio_transition(command_succeeded=True) + return stream + + def _write_live_audio_stream( + self, + stream: AudioLiveStream, + payload: bytes, + sequence: int, + ) -> int: + with self._audio_playback_lock: + if self._live_audio_stream is not stream: + raise WatcheRobotError("audio live stream was replaced") + return self._transport.write_live_audio_stream( + payload, + stream_id=stream.stream_id, + sequence=sequence, + chunk_bytes=960, + ) + + def _close_live_audio_stream(self, stream: AudioLiveStream, sequence: int) -> None: + with self._audio_api_lock: + with self._audio_playback_lock: + if self._live_audio_stream is not stream: + return + self._transport.end_live_audio_stream( + stream_id=stream.stream_id, + sequence=sequence, + ) + with self._audio_playback_lock: + if self._live_audio_stream is stream: + self._live_audio_stream = None + + def _abort_live_audio_stream(self, stream: AudioLiveStream) -> None: + with self._audio_api_lock: + with self._audio_playback_lock: + if self._live_audio_stream is not stream: + return + try: + try: + self._transport.cancel_live_audio_stream(stream_id=stream.stream_id) + finally: + self._command("ctrl.audio.stop", {}) + finally: + with self._audio_playback_lock: + if self._live_audio_stream is stream: + self._live_audio_stream = None + def _begin_audio_transition(self) -> None: while True: with self._audio_playback_lock: @@ -566,6 +715,8 @@ def _cancel_audio_playback(self, playback: AudioPlayback) -> None: self._end_audio_transition(command_succeeded=True) def _stop_audio_playback(self) -> None: + if self._abort_active_live_audio(): + return with self._audio_api_lock: self._begin_audio_transition() try: @@ -579,6 +730,7 @@ def _stop_audio_playback(self) -> None: def _open_microphone(self, *, queue_size: int, decode_opus: bool) -> MicrophoneSession: if queue_size <= 0: raise ValueError("queue_size must be positive") + self._abort_active_live_audio() with self._media_lock: if self._closed or self._closing: raise WatcheRobotError("robot connection is closed") @@ -943,6 +1095,11 @@ def _on_disconnect(self) -> None: self.inputs._close("disconnected") if microphone is not None: microphone._mark_remote_closed() + with self._audio_playback_lock: + live_stream = self._live_audio_stream + self._live_audio_stream = None + if live_stream is not None: + live_stream._mark_closed() with self._face_tracking_lock: preview = self._face_tracking_preview self._face_tracking_preview = None diff --git a/src/watcherobot/runtime/daemon/connections/registry.py b/src/watcherobot/runtime/daemon/connections/registry.py index 038881d..149cd34 100644 --- a/src/watcherobot/runtime/daemon/connections/registry.py +++ b/src/watcherobot/runtime/daemon/connections/registry.py @@ -15,6 +15,7 @@ class ExternalClientRole(str, Enum): UNKNOWN = "unknown" DESKTOP = "desktop" DEVICE = "hardware" + MEDIA = "media" class ConnectionRegistryError(RuntimeError): @@ -89,6 +90,51 @@ def declare_role( connection.metadata = {} return declared_role + def find_device_control( + self, + *, + pair_request_id: str, + daemon_instance_id: str, + session_token: str, + peer_ip: str, + ) -> ExternalConnection | None: + for connection in self.connections_for(ExternalClientRole.DEVICE): + metadata = connection.metadata + if ( + metadata.get("pair_request_id") == pair_request_id + and metadata.get("daemon_instance_id") == daemon_instance_id + and metadata.get("session_token") == session_token + and metadata.get("peer_ip") == peer_ip + ): + return connection + return None + + def bind_media( + self, + connection: ExternalConnection, + control: ExternalConnection, + ) -> None: + if connection.role is not ExternalClientRole.UNKNOWN: + raise ClientRoleLockedError("media connection role is already locked") + connection.role = ExternalClientRole.MEDIA + connection.metadata = {"control_websocket": control.websocket} + + async def close_media_for_control( + self, + control: ExternalConnection, + *, + code: int, + reason: str, + ) -> int: + sidecars = [ + connection + for connection in self.connections_for(ExternalClientRole.MEDIA) + if connection.metadata.get("control_websocket") is control.websocket + ] + for connection in sidecars: + await connection.websocket.close(code=code, reason=reason) + return len(sidecars) + def connections_for( self, role: ExternalClientRole, diff --git a/src/watcherobot/runtime/daemon/connections/websocket_server.py b/src/watcherobot/runtime/daemon/connections/websocket_server.py index 7dd70a5..cb2b64a 100644 --- a/src/watcherobot/runtime/daemon/connections/websocket_server.py +++ b/src/watcherobot/runtime/daemon/connections/websocket_server.py @@ -25,8 +25,10 @@ PairingProtocolError, build_hardware_hello_ack, build_hardware_hello_nack, + build_media_hello_ack, parse_device_session_end, parse_hardware_hello, + parse_media_hello, ) from watcherobot.runtime.daemon.pairing.session import PairingSessionError from watcherobot.runtime.daemon.routing.raw import RawFrameRouter @@ -42,7 +44,7 @@ ] BusinessFrameListener = Callable[ [ExternalConnection, str | bytes], - Awaitable[None] | None, + Awaitable[bool | None] | bool | None, ] ExternalDisconnectListener = Callable[ [ExternalConnection], @@ -61,6 +63,8 @@ class ExternalWebSocketServer: CLOSE_DEVICE_SLOT_OCCUPIED = 4412 CLOSE_PAIRING_PROTOCOL_MISMATCH = 4413 CLOSE_DESKTOP_LOOPBACK_REQUIRED = 4414 + CLOSE_MEDIA_CONTROL_REQUIRED = 4415 + CLOSE_INVALID_MEDIA_FRAME = 4416 def __init__( self, @@ -158,7 +162,9 @@ async def _handle_connection(self, websocket: ServerConnection) -> None: return async for frame in websocket: - if self._is_control_message(frame, "sys.client.hello"): + if self._is_control_message(frame, "sys.client.hello") or self._is_control_message( + frame, "sys.media.hello" + ): await websocket.close( code=self.CLOSE_ROLE_LOCKED, reason="client role is already locked", @@ -176,6 +182,15 @@ async def _handle_connection(self, websocket: ServerConnection) -> None: ) ) continue + if connection.role is ExternalClientRole.MEDIA: + if not self._is_video_frame(frame): + await websocket.close( + code=self.CLOSE_INVALID_MEDIA_FRAME, + reason="video sidecar accepts WSPK video only", + ) + return + await self.router.route_external(connection, frame) + continue if ( connection.role is ExternalClientRole.DEVICE and self._is_control_message( @@ -192,12 +207,20 @@ async def _handle_connection(self, websocket: ServerConnection) -> None: frame, ) if inspect.isawaitable(result): - await result + result = await result + if result is True: + continue await self.router.route_external(connection, frame) except ConnectionClosed: pass finally: removed = self.registry.remove(websocket) + if removed is not None and removed.role is ExternalClientRole.DEVICE: + await self.registry.close_media_for_control( + removed, + code=1001, + reason="control connection closed", + ) if ( removed is not None and self._external_disconnect_listener is not None @@ -298,6 +321,9 @@ async def _handle_hello( connection: ExternalConnection, frame: str | bytes, ) -> bool: + if self._is_control_message(frame, "sys.media.hello"): + return await self._handle_media_hello(connection, frame) + payload = self._parse_control_payload(frame, "sys.client.hello") if payload is None: await connection.websocket.close( @@ -395,8 +421,29 @@ async def _handle_hello( ) return False + if role is ExternalClientRole.DESKTOP: + capabilities = payload.get("capabilities", []) + connection.metadata = { + "client_name": str(payload.get("client_name", "")), + "capabilities": ( + list(capabilities) + if isinstance(capabilities, list) + and all(isinstance(item, str) for item in capabilities) + else [] + ), + } + if role is ExternalClientRole.DEVICE: assert hardware_ack is not None + assert hello is not None + connection.metadata.update( + { + "pair_request_id": hello.pair_request_id, + "daemon_instance_id": hello.daemon_instance_id, + "session_token": hello.session_token, + "peer_ip": self._peer_ip(connection.websocket), + } + ) acknowledgement = hardware_ack else: acknowledgement = { @@ -415,6 +462,37 @@ async def _handle_hello( ) return True + async def _handle_media_hello( + self, + connection: ExternalConnection, + frame: str | bytes, + ) -> bool: + try: + hello = parse_media_hello(frame) + except PairingProtocolError: + await connection.websocket.close( + code=self.CLOSE_MEDIA_CONTROL_REQUIRED, + reason="invalid sys.media.hello", + ) + return False + control = self.registry.find_device_control( + pair_request_id=hello.pair_request_id, + daemon_instance_id=hello.daemon_instance_id, + session_token=hello.session_token, + peer_ip=self._peer_ip(connection.websocket), + ) + if control is None: + await connection.websocket.close( + code=self.CLOSE_MEDIA_CONTROL_REQUIRED, + reason="matching control connection required", + ) + return False + self.registry.bind_media(connection, control) + await connection.websocket.send( + json.dumps(build_media_hello_ack(), separators=(",", ":")) + ) + return True + async def _reject_hardware_hello( self, connection: ExternalConnection, @@ -480,3 +558,12 @@ def _parse_control_payload( ): return None return dict(message["data"]) + + @staticmethod + def _is_video_frame(frame: str | bytes) -> bool: + return ( + isinstance(frame, bytes) + and len(frame) >= 14 + and frame[:4] == b"WSPK" + and frame[4] == 2 + ) diff --git a/src/watcherobot/runtime/daemon/pairing/protocol.py b/src/watcherobot/runtime/daemon/pairing/protocol.py index 8ad8fa4..4e9a75d 100644 --- a/src/watcherobot/runtime/daemon/pairing/protocol.py +++ b/src/watcherobot/runtime/daemon/pairing/protocol.py @@ -50,6 +50,18 @@ "mode", } ) +_MEDIA_HELLO_DATA_FIELDS = frozenset( + { + "pairing_protocol", + "pairing_version", + "pair_request_id", + "daemon_instance_id", + "session_token", + "mode", + "channel", + "version", + } +) _SESSION_END_FIELDS = frozenset({"type", "code", "data"}) _SESSION_END_DATA_FIELDS = frozenset({"pair_request_id", "reason"}) _SENSITIVE_FIELDS = frozenset({"pairing_code", "session_token"}) @@ -98,6 +110,16 @@ class HardwareHello: mode: str +@dataclass(frozen=True) +class MediaHello: + pair_request_id: str + daemon_instance_id: str + session_token: str + mode: str + channel: str + version: int + + @dataclass(frozen=True) class DeviceSessionEnd: pair_request_id: str @@ -289,6 +311,35 @@ def parse_hardware_hello(payload: RawPayload) -> HardwareHello: ) +def parse_media_hello(payload: RawPayload) -> MediaHello: + """Parse a video sidecar hello bound to an online control session.""" + + message = _decode_object(payload, max_bytes=MAX_HELLO_PAYLOAD_BYTES) + _require_exact_fields(message, _HELLO_FIELDS) + _require_literal(message, "type", "sys.media.hello") + if type(message.get("code")) is not int or message["code"] != 0: + raise PairingProtocolError("code must be integer zero") + data = message.get("data") + if not isinstance(data, Mapping): + raise PairingProtocolError("data must be a JSON object") + data = dict(data) + _require_exact_fields(data, _MEDIA_HELLO_DATA_FIELDS) + _require_literal(data, "pairing_protocol", LAN_PAIRING_PROTOCOL) + _require_literal(data, "pairing_version", LAN_PAIRING_VERSION) + channel = _require_literal(data, "channel", "video") + version = data.get("version") + if type(version) is not int or version != 1: + raise PairingProtocolError("version must be integer one") + return MediaHello( + pair_request_id=_require_lower_hex(data, "pair_request_id", _LOWER_HEX_32), + daemon_instance_id=_require_lower_hex(data, "daemon_instance_id", _LOWER_HEX_32), + session_token=_require_lower_hex(data, "session_token", _LOWER_HEX_64), + mode=_require_literal(data, "mode", LAN_PAIRING_TARGET_MODE), + channel=channel, + version=version, + ) + + def parse_device_session_end(payload: RawPayload) -> DeviceSessionEnd: """Parse the hardware's explicit normal session termination.""" @@ -390,11 +441,24 @@ def build_hardware_hello_ack() -> dict[str, object]: "frame_duration_ms": 60, "packetization": "one_opus_packet_per_wspk", "version": 1, - } + }, + "video_uplink": { + "transport": "websocket_sidecar", + "hello_type": "sys.media.hello", + "version": 1, + }, }, }, } + +def build_media_hello_ack() -> dict[str, object]: + return { + "type": "sys.ack", + "code": 0, + "data": {"type": "sys.media.hello", "channel": "video", "version": 1}, + } + def build_hardware_hello_nack( *, code: int, diff --git a/src/watcherobot/runtime/daemon/pairing/session.py b/src/watcherobot/runtime/daemon/pairing/session.py index bc49e42..52c259c 100644 --- a/src/watcherobot/runtime/daemon/pairing/session.py +++ b/src/watcherobot/runtime/daemon/pairing/session.py @@ -19,7 +19,7 @@ ) -DEFAULT_DISCOVERY_TIMEOUT_SECONDS = 10.0 +DEFAULT_DISCOVERY_TIMEOUT_SECONDS = 30.0 DEFAULT_CONNECT_TIMEOUT_SECONDS = 10.0 DEFAULT_RECONNECT_TIMEOUT_SECONDS = 30.0 diff --git a/src/watcherobot/runtime/daemon/preview/delivery.py b/src/watcherobot/runtime/daemon/preview/delivery.py new file mode 100644 index 0000000..a911479 --- /dev/null +++ b/src/watcherobot/runtime/daemon/preview/delivery.py @@ -0,0 +1,214 @@ +"""Per-browser latest-frame delivery with explicit consumer credit.""" + +from __future__ import annotations + +import asyncio +import json +import time +from dataclasses import asdict, dataclass +from typing import Any, Callable + +from websockets.exceptions import ConnectionClosed + +from ..connections.registry import ( + ExternalClientRole, + ExternalConnectionRegistry, +) + + +PREVIEW_CREDIT_CAPABILITY = "face_tracking.preview.credit.v1" + + +@dataclass(frozen=True) +class PreviewRelayFrame: + stream_id: int + sequence: int + telemetry: str + image: bytes + completed_at: float + + +@dataclass +class PreviewDeliveryStats: + offered_frames: int = 0 + sent_frames: int = 0 + acknowledged_frames: int = 0 + pending_overwrites: int = 0 + unexpected_acknowledgements: int = 0 + send_errors: int = 0 + + +@dataclass +class _ConsumerState: + websocket: Any + credit_required: bool + in_flight_sequence: int | None = None + pending: PreviewRelayFrame | None = None + pending_overwrites: int = 0 + sender: asyncio.Task[None] | None = None + + +class LatestPreviewDelivery: + """Keep at most one in-flight and one replaceable frame per browser.""" + + def __init__( + self, + registry: ExternalConnectionRegistry, + *, + clock: Callable[[], float] = time.monotonic, + ) -> None: + self._registry = registry + self._clock = clock + self._consumers: dict[object, _ConsumerState] = {} + self.stats = PreviewDeliveryStats() + + def offer(self, frame: PreviewRelayFrame) -> int: + self.stats.offered_frames += 1 + delivered = 0 + for connection in self._registry.connections_for(ExternalClientRole.DESKTOP): + websocket = connection.websocket + state = self._consumers.get(websocket) + if state is None: + capabilities = connection.metadata.get("capabilities", []) + state = _ConsumerState( + websocket=websocket, + credit_required=( + isinstance(capabilities, list) + and PREVIEW_CREDIT_CAPABILITY in capabilities + ), + ) + self._consumers[websocket] = state + self._offer_to_consumer(state, frame) + delivered += 1 + return delivered + + def acknowledge(self, websocket: object, *, sequence: int) -> bool: + state = self._consumers.get(websocket) + if state is None or state.in_flight_sequence != sequence: + self.stats.unexpected_acknowledgements += 1 + return False + state.in_flight_sequence = None + self.stats.acknowledged_frames += 1 + if state.sender is None: + self._send_pending(state) + return True + + def connection_lost(self, websocket: object) -> None: + state = self._consumers.pop(websocket, None) + if state is not None and state.sender is not None: + state.sender.cancel() + + def discard_pending(self) -> None: + """Forget unsent frames while preserving at most one in-flight frame.""" + + for state in self._consumers.values(): + state.pending = None + state.pending_overwrites = 0 + + async def stop(self) -> None: + tasks = [ + state.sender + for state in self._consumers.values() + if state.sender is not None + ] + self._consumers.clear() + for task in tasks: + task.cancel() + if tasks: + await asyncio.gather(*tasks, return_exceptions=True) + + def snapshot(self) -> dict[str, object]: + return { + **asdict(self.stats), + "active_consumers": len(self._consumers), + "credit_consumers": sum( + 1 for state in self._consumers.values() if state.credit_required + ), + } + + def _offer_to_consumer( + self, + state: _ConsumerState, + frame: PreviewRelayFrame, + ) -> None: + if state.in_flight_sequence is None and state.sender is None: + self._start_send(state, frame, skipped_frames=0) + return + if state.pending is not None: + state.pending_overwrites += 1 + self.stats.pending_overwrites += 1 + state.pending = frame + + def _send_pending(self, state: _ConsumerState) -> None: + frame = state.pending + if frame is None or state.sender is not None: + return + skipped_frames = state.pending_overwrites + state.pending = None + state.pending_overwrites = 0 + self._start_send(state, frame, skipped_frames=skipped_frames) + + def _start_send( + self, + state: _ConsumerState, + frame: PreviewRelayFrame, + *, + skipped_frames: int, + ) -> None: + state.in_flight_sequence = frame.sequence + state.sender = asyncio.create_task( + self._send(state, frame, skipped_frames), + name=f"preview-delivery-{frame.sequence}", + ) + + async def _send( + self, + state: _ConsumerState, + frame: PreviewRelayFrame, + skipped_frames: int, + ) -> None: + failed = False + try: + telemetry = self._telemetry_for_delivery(frame, skipped_frames) + await state.websocket.send(telemetry) + await state.websocket.send(frame.image) + self.stats.sent_frames += 1 + except ( + ConnectionClosed, + OSError, + RuntimeError, + TypeError, + ValueError, + ): + failed = True + self.stats.send_errors += 1 + finally: + state.sender = None + if self._consumers.get(state.websocket) is not state: + state.pending = None + state.pending_overwrites = 0 + return + if failed: + state.in_flight_sequence = None + state.pending = None + state.pending_overwrites = 0 + self._consumers.pop(state.websocket, None) + return + if not state.credit_required: + state.in_flight_sequence = None + if state.in_flight_sequence is None: + self._send_pending(state) + + def _telemetry_for_delivery( + self, + frame: PreviewRelayFrame, + skipped_frames: int, + ) -> str: + payload = json.loads(frame.telemetry) + queue_ms = round(max(0.0, self._clock() - frame.completed_at) * 1000, 1) + payload["relay"] = [ + int(round(frame.completed_at * 1000)), + queue_ms, + skipped_frames, + ] + return json.dumps(payload, separators=(",", ":")) diff --git a/src/watcherobot/runtime/daemon/preview/face_tracking.py b/src/watcherobot/runtime/daemon/preview/face_tracking.py index 878614a..646b05f 100644 --- a/src/watcherobot/runtime/daemon/preview/face_tracking.py +++ b/src/watcherobot/runtime/daemon/preview/face_tracking.py @@ -3,6 +3,7 @@ from __future__ import annotations import json +from collections.abc import Awaitable, Callable from itertools import count from ..application.session import ApplicationChannel @@ -18,12 +19,18 @@ class FaceTrackingPreviewBroker: START_TYPE = "ctrl.face_tracking.preview.start" STOP_TYPE = "ctrl.face_tracking.preview.stop" + CONSUMED_TYPE = "daemon.face_tracking.preview.consumed" def __init__( self, registry: ExternalConnectionRegistry | None = None, + *, + on_preview_start: Callable[[], Awaitable[None]] | None = None, + on_preview_ack: Callable[[ExternalConnection, int], bool] | None = None, ) -> None: self._registry = registry + self._on_preview_start = on_preview_start + self._on_preview_ack = on_preview_ack self._owners: set[object] = set() self._application_owner = object() self._command_sequence = count(1) @@ -37,10 +44,16 @@ async def observe_frame( self, source: ExternalConnection, frame: str | bytes, - ) -> None: + ) -> bool: if source.role is not ExternalClientRole.DESKTOP: - return - self._observe_owner(source.websocket, frame) + return False + acknowledged_sequence = self._consumed_sequence(frame) + if acknowledged_sequence is not None: + if self._on_preview_ack is not None: + self._on_preview_ack(source, acknowledged_sequence) + return True + await self._observe_owner(source.websocket, frame) + return False async def observe_application_frame( self, @@ -48,11 +61,13 @@ async def observe_application_frame( frame: str | bytes, ) -> None: if source is ApplicationChannel.DEVICE: - self._observe_owner(self._application_owner, frame) + await self._observe_owner(self._application_owner, frame) - def _observe_owner(self, owner: object, frame: str | bytes) -> None: + async def _observe_owner(self, owner: object, frame: str | bytes) -> None: message_type = self._message_type(frame) if message_type == self.START_TYPE: + if self._on_preview_start is not None: + await self._on_preview_start() self._owners.add(owner) elif message_type == self.STOP_TYPE: self._owners.clear() @@ -111,3 +126,25 @@ def _message_type(frame: str | bytes) -> str | None: if not isinstance(command_id, str) or not command_id: return None return message_type + + @classmethod + def _consumed_sequence(cls, frame: str | bytes) -> int | None: + if not isinstance(frame, str): + return None + try: + message = json.loads(frame) + except json.JSONDecodeError: + return None + if not isinstance(message, dict): + return None + if message.get("type") != cls.CONSUMED_TYPE: + return None + data = message.get("data") + if not isinstance(data, dict): + return None + sequence = data.get("sequence") + if not isinstance(sequence, int) or isinstance(sequence, bool): + return None + if sequence < 0 or sequence > 0xFFFFFFFF: + return None + return sequence diff --git a/src/watcherobot/runtime/daemon/preview/udp_service.py b/src/watcherobot/runtime/daemon/preview/udp_service.py index d45d5e7..214ef4a 100644 --- a/src/watcherobot/runtime/daemon/preview/udp_service.py +++ b/src/watcherobot/runtime/daemon/preview/udp_service.py @@ -12,6 +12,7 @@ from ..connections.registry import ExternalClientRole, ExternalConnectionRegistry from ..pairing.session import DevicePairingSession from .udp_feedback import encode_preview_ack +from .delivery import PreviewRelayFrame from .udp_protocol import ( CompletedPreviewFrame, FaceTrackingUdpProtocolError, @@ -50,6 +51,7 @@ def __init__( session: DevicePairingSession, registry: ExternalConnectionRegistry, publisher: Callable[[str | bytes], Awaitable[int]] | None = None, + bundle_publisher: Callable[[PreviewRelayFrame], Awaitable[int]] | None = None, host: str = "0.0.0.0", port: int = 37022, clock: Callable[[], float] = time.monotonic, @@ -57,12 +59,14 @@ def __init__( self._session = session self._registry = registry self._publisher_callback = publisher + self._bundle_publisher_callback = bundle_publisher self._host = host self._port = port self._clock = clock self._transport: asyncio.DatagramTransport | None = None self._publisher: asyncio.Task[None] | None = None self._publish_event = asyncio.Event() + self._listener_lock = asyncio.Lock() self._pending: CompletedPreviewFrame | None = None self._pending_completed_at = 0.0 self._reassembler: FaceTrackingUdpReassembler | None = None @@ -115,6 +119,13 @@ async def stop(self) -> None: pass self._reset_session() + async def refresh_listener(self) -> None: + """Rebind the UDP endpoint before a new preview stream starts.""" + + async with self._listener_lock: + await self.stop() + await self.start() + def handle_datagram(self, data: bytes, address: tuple[str, int]) -> None: self.stats.datagrams_received += 1 credentials = self._session.preview_transport_credentials() @@ -187,17 +198,29 @@ async def _publish_loop(self) -> None: raise FaceTrackingUdpProtocolError( "preview telemetry is not an object" ) + telemetry = json.dumps(payload, separators=(",", ":")) + except (FaceTrackingUdpProtocolError, json.JSONDecodeError): + self.stats.invalid_datagrams += 1 + continue + bundle_publisher = self._bundle_publisher_callback + if bundle_publisher is not None: + await bundle_publisher( + PreviewRelayFrame( + stream_id=frame.stream_id, + sequence=frame.sequence, + telemetry=telemetry, + image=image, + completed_at=completed_at, + ) + ) + else: published_at = self._clock() payload["relay"] = [ int(round(completed_at * 1000)), round(max(0.0, published_at - completed_at) * 1000, 1), ] - telemetry = json.dumps(payload, separators=(",", ":")) - except (FaceTrackingUdpProtocolError, json.JSONDecodeError): - self.stats.invalid_datagrams += 1 - continue - await self._publish(telemetry) - await self._publish(image) + await self._publish(json.dumps(payload, separators=(",", ":"))) + await self._publish(image) self.stats.published_frames += 1 async def _publish(self, frame: str | bytes) -> int: diff --git a/src/watcherobot/runtime/daemon/routing/raw.py b/src/watcherobot/runtime/daemon/routing/raw.py index 8ea1143..8766bfe 100644 --- a/src/watcherobot/runtime/daemon/routing/raw.py +++ b/src/watcherobot/runtime/daemon/routing/raw.py @@ -61,7 +61,7 @@ async def route_external( ExternalClientRole.DEVICE, frame, ) - if source.role is ExternalClientRole.DEVICE: + if source.role in (ExternalClientRole.DEVICE, ExternalClientRole.MEDIA): return await self._registry.send_to_role( ExternalClientRole.DESKTOP, frame, @@ -88,6 +88,6 @@ def _application_channel_for_external( ) -> ApplicationChannel | None: if role is ExternalClientRole.DESKTOP: return ApplicationChannel.DESKTOP - if role is ExternalClientRole.DEVICE: + if role in (ExternalClientRole.DEVICE, ExternalClientRole.MEDIA): return ApplicationChannel.DEVICE return None diff --git a/src/watcherobot/runtime/daemon/runtime.py b/src/watcherobot/runtime/daemon/runtime.py index 90e2471..86c9b8b 100644 --- a/src/watcherobot/runtime/daemon/runtime.py +++ b/src/watcherobot/runtime/daemon/runtime.py @@ -26,6 +26,7 @@ ) from watcherobot.runtime.daemon.connections.registry import ( ExternalClientRole, + ExternalConnection, ExternalConnectionRegistry, ) from watcherobot.runtime.daemon.connections.websocket_server import ( @@ -46,6 +47,10 @@ from watcherobot.runtime.daemon.preview.face_tracking import ( FaceTrackingPreviewBroker, ) +from watcherobot.runtime.daemon.preview.delivery import ( + LatestPreviewDelivery, + PreviewRelayFrame, +) from watcherobot.runtime.daemon.preview.udp_service import ( FaceTrackingUdpPreviewService, ) @@ -114,7 +119,15 @@ def __init__( ) self._clock = clock self.connection_registry = connection_registry - self.face_tracking_preview = FaceTrackingPreviewBroker(connection_registry) + self.preview_delivery = LatestPreviewDelivery( + connection_registry, + clock=clock, + ) + self.face_tracking_preview = FaceTrackingPreviewBroker( + connection_registry, + on_preview_start=self._refresh_preview_udp_listener, + on_preview_ack=self._acknowledge_preview_frame, + ) self.application.bridge.set_frame_callback( self._route_application_frame ) @@ -130,7 +143,7 @@ def __init__( device_disconnect_listener=self._device_disconnected, device_session_end_listener=self._device_session_ended, business_frame_listener=self.face_tracking_preview.observe_frame, - external_disconnect_listener=(self.face_tracking_preview.connection_lost), + external_disconnect_listener=self._external_connection_lost, ) self.pairing_udp = PairingUdpService( session=self.device_pairing, @@ -142,8 +155,9 @@ def __init__( self.preview_udp = FaceTrackingUdpPreviewService( session=self.device_pairing, registry=connection_registry, - publisher=self._publish_preview_frame, + bundle_publisher=self._publish_preview_bundle, port=preview_udp_port, + clock=clock, ) self.control_server = DaemonControlServer( controller=self, @@ -155,6 +169,14 @@ def __init__( self.auto_start_error: str | None = None self._shutdown_event = asyncio.Event() + async def _refresh_preview_udp_listener(self) -> None: + self.preview_delivery.discard_pending() + await self.preview_udp.refresh_listener() + self.logs.record( + "Preview UDP listener refreshed " + f"(port={self.preview_udp.bound_port})" + ) + async def start(self) -> None: self.logs.record("Daemon Runtime starting") await self.external_server.start() @@ -195,6 +217,7 @@ async def stop(self) -> None: await self.control_server.stop() await self.application.stop() await self.preview_udp.stop() + await self.preview_delivery.stop() await self.pairing_udp.stop() await self.external_server.stop() @@ -356,7 +379,10 @@ def device_status(self) -> dict[str, object]: if device["online"] and peer_ip is not None else None ) - device["preview_transport"] = self.preview_udp.snapshot() + device["preview_transport"] = { + **self.preview_udp.snapshot(), + "delivery": self.preview_delivery.snapshot(), + } return {"device": device} async def _authorize_hardware_hello( @@ -422,20 +448,38 @@ async def _route_application_frame( ) return delivered - async def _publish_preview_frame(self, frame: str | bytes) -> int: + def _acknowledge_preview_frame( + self, + connection: ExternalConnection, + sequence: int, + ) -> bool: + return self.preview_delivery.acknowledge( + connection.websocket, + sequence=sequence, + ) + + async def _external_connection_lost( + self, + connection: ExternalConnection, + ) -> None: + self.preview_delivery.connection_lost(connection.websocket) + await self.face_tracking_preview.connection_lost(connection) + + async def _publish_preview_bundle(self, frame: PreviewRelayFrame) -> int: if self.application.registry.active_run is not None: try: await self.application.bridge.send_to_application( ApplicationChannel.DEVICE, - frame, + frame.telemetry, + ) + await self.application.bridge.send_to_application( + ApplicationChannel.DEVICE, + frame.image, ) except ChannelNotConnectedError: return 0 return 1 - return await self.connection_registry.send_to_role( - ExternalClientRole.DESKTOP, - frame, - ) + return self.preview_delivery.offer(frame) def _catalog_entry_payload(entry: CatalogEntry) -> dict[str, object]: diff --git a/tests/application/test_transport.py b/tests/application/test_transport.py index d45e3d7..98279be 100644 --- a/tests/application/test_transport.py +++ b/tests/application/test_transport.py @@ -4,8 +4,11 @@ import json import struct +import pytest + from watcherobot.application.transport import DaemonApplicationTransport -from watcherobot.protocol import FRAME_VIDEO +from watcherobot.errors import WatcheRobotError +from watcherobot.protocol import FLAG_FIRST, FLAG_LAST, FRAME_AUDIO, FRAME_VIDEO, parse_wspk from watcherobot.runtime.daemon.application.session import ApplicationChannel @@ -58,6 +61,128 @@ async def capture( asyncio.run(scenario()) +def test_live_audio_stream_uses_credit_and_finishes_with_last_frame() -> None: + async def scenario() -> None: + transport = DaemonApplicationTransport(command_timeout=1.0) + sent: list[bytes] = [] + + async def capture(channel: ApplicationChannel, frame: str | bytes) -> None: + assert channel is ApplicationChannel.DEVICE + assert isinstance(frame, bytes) + sent.append(frame) + + transport._send = capture # type: ignore[method-assign] + await transport._begin_live_audio_stream(stream_id=9, chunk_bytes=960) + sequence = await transport._write_live_audio_stream( + b"\x01\x00" * 600, + stream_id=9, + sequence=0, + chunk_bytes=960, + ) + await transport._end_live_audio_stream(stream_id=9, sequence=sequence) + + frames = [parse_wspk(packet) for packet in sent] + assert [frame.sequence for frame in frames] == [0, 1, 2] + assert frames[0].flags == FLAG_FIRST + assert frames[0].payload == b"\x01\x00" * 480 + assert frames[1].flags == 0 + assert frames[-1].flags == FLAG_LAST + assert frames[-1].frame_type == FRAME_AUDIO + assert frames[-1].payload == b"" + + asyncio.run(scenario()) + + +def test_live_audio_stream_does_not_repeat_first_flag_after_sequence_wrap() -> None: + async def scenario() -> None: + transport = DaemonApplicationTransport(command_timeout=1.0) + sent: list[bytes] = [] + + async def capture(channel: ApplicationChannel, frame: str | bytes) -> None: + assert channel is ApplicationChannel.DEVICE + assert isinstance(frame, bytes) + sent.append(frame) + + transport._send = capture # type: ignore[method-assign] + await transport._begin_live_audio_stream(stream_id=10, chunk_bytes=960) + sequence = await transport._write_live_audio_stream( + b"\x01\x00" * 960, + stream_id=10, + sequence=0xFFFFFFFF, + chunk_bytes=960, + ) + + frames = [parse_wspk(packet) for packet in sent] + assert [frame.sequence for frame in frames] == [0xFFFFFFFF, 0] + assert [frame.flags for frame in frames] == [FLAG_FIRST, 0] + assert sequence == 1 + + asyncio.run(scenario()) + + +def test_live_audio_stream_blocks_on_device_credit_and_wakes_on_status() -> None: + async def scenario() -> None: + transport = DaemonApplicationTransport(command_timeout=1.0) + sent: list[bytes] = [] + + async def capture(channel: ApplicationChannel, frame: str | bytes) -> None: + assert channel is ApplicationChannel.DEVICE + assert isinstance(frame, bytes) + sent.append(frame) + + transport._send = capture # type: ignore[method-assign] + await transport._begin_live_audio_stream(stream_id=12, chunk_bytes=960) + write = asyncio.create_task( + transport._write_live_audio_stream( + b"\x01\x00" * (480 * 5), + stream_id=12, + sequence=0, + chunk_bytes=960, + ) + ) + for _ in range(100): + if len(sent) == 4: + break + await asyncio.sleep(0) + assert len(sent) == 4 + assert not write.done() + + await transport._update_audio_flow( + {"stream_id": 12, "reason": "playback", "pending_frames": 0, "queue_depth": 16} + ) + assert await write == 5 + assert len(sent) == 5 + + asyncio.run(scenario()) + + +def test_live_audio_stream_device_failure_wakes_blocked_writer() -> None: + async def scenario() -> None: + transport = DaemonApplicationTransport(command_timeout=1.0) + + async def capture(_channel: ApplicationChannel, _frame: str | bytes) -> None: + return None + + transport._send = capture # type: ignore[method-assign] + await transport._begin_live_audio_stream(stream_id=13, chunk_bytes=960) + write = asyncio.create_task( + transport._write_live_audio_stream( + b"\x01\x00" * (480 * 5), + stream_id=13, + sequence=0, + chunk_bytes=960, + ) + ) + await asyncio.sleep(0) + await transport._update_audio_flow( + {"stream_id": 13, "reason": "playback_write_failed"} + ) + with pytest.raises(WatcheRobotError, match="playback_write_failed"): + await write + + asyncio.run(scenario()) + + def test_transport_dispatches_face_preview_packet_as_video_frame() -> None: transport = DaemonApplicationTransport() received = [] diff --git a/tests/fixtures/contracts/watcher_lan_pairing_v1.json b/tests/fixtures/contracts/watcher_lan_pairing_v1.json index 0bcd2dc..d244454 100644 --- a/tests/fixtures/contracts/watcher_lan_pairing_v1.json +++ b/tests/fixtures/contracts/watcher_lan_pairing_v1.json @@ -71,10 +71,38 @@ "frame_duration_ms": 60, "packetization": "one_opus_packet_per_wspk", "version": 1 + }, + "video_uplink": { + "transport": "websocket_sidecar", + "hello_type": "sys.media.hello", + "version": 1 } } } }, + "media_hello": { + "type": "sys.media.hello", + "code": 0, + "data": { + "pairing_protocol": "watcher-lan-pairing", + "pairing_version": "1.0", + "pair_request_id": "21a9dbf05ea3443480e62076f79a3b12", + "daemon_instance_id": "f730f29e670c49f7a3320c4314eb9805", + "session_token": "f84a1e16ce6f35f14d167f227a93ea93d1a9c4d9eb5517112030f2839d57ae4b", + "mode": "desktop_link", + "channel": "video", + "version": 1 + } + }, + "media_hello_ack": { + "type": "sys.ack", + "code": 0, + "data": { + "type": "sys.media.hello", + "channel": "video", + "version": 1 + } + }, "session_end": { "type": "sys.device.session.end", "code": 0, diff --git a/tests/runtime/test_daemon_runtime_routing.py b/tests/runtime/test_daemon_runtime_routing.py index 3e2b72f..2c7f701 100644 --- a/tests/runtime/test_daemon_runtime_routing.py +++ b/tests/runtime/test_daemon_runtime_routing.py @@ -8,6 +8,7 @@ from websockets.asyncio.client import connect from watcherobot.runtime.daemon.runtime import DaemonRuntime +from watcherobot.runtime.daemon.preview.delivery import PreviewRelayFrame from watcherobot.runtime.daemon.application.session import ApplicationChannel from tests.runtime.pairing_helpers import connect_runtime_hardware @@ -101,8 +102,18 @@ async def capture( runtime.application.bridge.send_to_application = capture # type: ignore[method-assign] - assert await runtime._publish_preview_frame(b"FTW1") == 1 - assert delivered == [(ApplicationChannel.DEVICE, b"FTW1")] + frame = PreviewRelayFrame( + stream_id=1, + sequence=2, + telemetry='{"v":1,"kind":"frame","seq":2}', + image=b"FTW1", + completed_at=1.0, + ) + assert await runtime._publish_preview_bundle(frame) == 1 + assert delivered == [ + (ApplicationChannel.DEVICE, frame.telemetry), + (ApplicationChannel.DEVICE, frame.image), + ] asyncio.run(scenario()) diff --git a/tests/runtime/test_external_websocket_routing.py b/tests/runtime/test_external_websocket_routing.py index b73c9b6..d896b5b 100644 --- a/tests/runtime/test_external_websocket_routing.py +++ b/tests/runtime/test_external_websocket_routing.py @@ -41,6 +41,29 @@ def _hello(role: str, **metadata) -> str: ) +def _media_hello(**metadata) -> str: + return json.dumps( + { + "type": "sys.media.hello", + "code": 0, + "data": { + "pairing_protocol": "watcher-lan-pairing", + "pairing_version": "1.0", + "pair_request_id": "21a9dbf05ea3443480e62076f79a3b12", + "daemon_instance_id": "f730f29e670c49f7a3320c4314eb9805", + "session_token": ( + "f84a1e16ce6f35f14d167f227a93ea93" + "d1a9c4d9eb5517112030f2839d57ae4b" + ), + "mode": "desktop_link", + "channel": "video", + "version": 1, + **metadata, + }, + } + ) + + async def _connect_as(server: ExternalWebSocketServer, role: str): websocket = await connect(server.url, max_size=None) await websocket.send(_hello(role)) @@ -84,6 +107,34 @@ async def scenario() -> None: asyncio.run(scenario()) +def test_desktop_hello_preserves_preview_delivery_capability() -> None: + async def scenario() -> None: + server = ExternalWebSocketServer(host="127.0.0.1", port=0) + await server.start() + desktop = await connect(server.url, max_size=None) + await desktop.send( + _hello( + "desktop", + client_name="media-debugger", + capabilities=["face_tracking.preview.credit.v1"], + ) + ) + await asyncio.wait_for(desktop.recv(), timeout=1) + try: + [connection] = server.registry.connections_for( + ExternalClientRole.DESKTOP + ) + assert connection.metadata == { + "client_name": "media-debugger", + "capabilities": ["face_tracking.preview.credit.v1"], + } + finally: + await desktop.close() + await server.stop() + + asyncio.run(scenario()) + + def test_desktop_role_is_restricted_to_loopback_peers() -> None: assert ExternalWebSocketServer.is_loopback_address("127.0.0.1") assert ExternalWebSocketServer.is_loopback_address("::1") @@ -271,3 +322,65 @@ async def handle_session_end(message, peer_ip: str) -> None: await server.stop() asyncio.run(scenario()) + + +def test_video_sidecar_requires_matching_online_control_and_routes_video() -> None: + async def scenario() -> None: + server = ExternalWebSocketServer( + host="127.0.0.1", + port=0, + hardware_hello_authorizer=_allow_hardware, + ) + await server.start() + desktop = await _connect_as(server, "desktop") + orphan = await connect(server.url, max_size=None) + await orphan.send(_media_hello()) + await asyncio.wait_for(orphan.wait_closed(), timeout=1) + assert orphan.close_code == server.CLOSE_MEDIA_CONTROL_REQUIRED + + device = await _connect_as(server, "hardware") + media = await connect(server.url, max_size=None) + await media.send(_media_hello()) + ack = json.loads(await asyncio.wait_for(media.recv(), timeout=1)) + assert ack["data"] == { + "type": "sys.media.hello", + "channel": "video", + "version": 1, + } + video = b"WSPK\x02\x00" + b"\x01\x00\x00\x00\x00\x00\x00\x00" + await media.send(video) + assert await asyncio.wait_for(desktop.recv(), timeout=1) == video + + await device.close() + await asyncio.wait_for(media.wait_closed(), timeout=1) + assert media.close_code == 1001 + await desktop.close() + await orphan.close() + await server.stop() + + asyncio.run(scenario()) + + +def test_video_sidecar_failure_does_not_close_control() -> None: + async def scenario() -> None: + server = ExternalWebSocketServer( + host="127.0.0.1", + port=0, + hardware_hello_authorizer=_allow_hardware, + ) + await server.start() + device = await _connect_as(server, "hardware") + media = await connect(server.url, max_size=None) + await media.send(_media_hello()) + await asyncio.wait_for(media.recv(), timeout=1) + await media.send(b"WSPK\x01\x00") + await asyncio.wait_for(media.wait_closed(), timeout=1) + assert media.close_code == server.CLOSE_INVALID_MEDIA_FRAME + + await device.send(json.dumps({"type": "sys.ping", "data": {}})) + pong = json.loads(await asyncio.wait_for(device.recv(), timeout=1)) + assert pong["type"] == "sys.pong" + await device.close() + await server.stop() + + asyncio.run(scenario()) diff --git a/tests/runtime/test_face_tracking_preview_broker.py b/tests/runtime/test_face_tracking_preview_broker.py index 2660463..7339457 100644 --- a/tests/runtime/test_face_tracking_preview_broker.py +++ b/tests/runtime/test_face_tracking_preview_broker.py @@ -91,6 +91,123 @@ async def scenario() -> None: asyncio.run(scenario()) +def test_preview_start_refreshes_udp_listener_before_forwarding() -> None: + async def scenario() -> None: + listener_refreshed = asyncio.Event() + + async def refresh_listener() -> None: + listener_refreshed.set() + + broker = FaceTrackingPreviewBroker( + on_preview_start=refresh_listener, + ) + server = ExternalWebSocketServer( + host="127.0.0.1", + port=0, + hardware_hello_authorizer=_allow_hardware, + business_frame_listener=broker.observe_frame, + ) + broker.bind_registry(server.registry) + await server.start() + desktop = await _connect_as(server, "desktop") + device = await _connect_as(server, "hardware") + try: + start = json.dumps( + { + "type": "ctrl.face_tracking.preview.start", + "data": {"command_id": "preview-start-refresh"}, + } + ) + await desktop.send(start) + assert await asyncio.wait_for(device.recv(), timeout=1) == start + assert listener_refreshed.is_set() + finally: + await desktop.close() + await device.close() + await server.stop() + + asyncio.run(scenario()) + + +def test_preview_consumed_ack_is_handled_by_daemon_without_reaching_device() -> None: + async def scenario() -> None: + acknowledgements: list[tuple[object, int]] = [] + acknowledged = asyncio.Event() + + def acknowledge(connection, sequence: int) -> bool: + acknowledgements.append((connection.websocket, sequence)) + acknowledged.set() + return True + + broker = FaceTrackingPreviewBroker(on_preview_ack=acknowledge) + server = ExternalWebSocketServer( + host="127.0.0.1", + port=0, + hardware_hello_authorizer=_allow_hardware, + business_frame_listener=broker.observe_frame, + ) + broker.bind_registry(server.registry) + await server.start() + desktop = await _connect_as(server, "desktop") + device = await _connect_as(server, "hardware") + try: + await desktop.send( + json.dumps( + { + "type": "daemon.face_tracking.preview.consumed", + "code": 0, + "data": {"sequence": 31}, + } + ) + ) + await asyncio.wait_for(acknowledged.wait(), timeout=1) + assert len(acknowledgements) == 1 + assert acknowledgements[0][1] == 31 + try: + await asyncio.wait_for(device.recv(), timeout=0.05) + except asyncio.TimeoutError: + pass + else: + raise AssertionError("preview consumption ACK leaked to device") + finally: + await desktop.close() + await device.close() + await server.stop() + + asyncio.run(scenario()) + + +def test_malformed_preview_consumed_ack_is_forwarded_without_crashing() -> None: + async def scenario() -> None: + broker = FaceTrackingPreviewBroker() + server = ExternalWebSocketServer( + host="127.0.0.1", + port=0, + hardware_hello_authorizer=_allow_hardware, + business_frame_listener=broker.observe_frame, + ) + broker.bind_registry(server.registry) + await server.start() + desktop = await _connect_as(server, "desktop") + device = await _connect_as(server, "hardware") + malformed = json.dumps( + { + "type": "daemon.face_tracking.preview.consumed", + "code": 0, + "data": [], + } + ) + try: + await desktop.send(malformed) + assert await asyncio.wait_for(device.recv(), timeout=1) == malformed + finally: + await desktop.close() + await device.close() + await server.stop() + + asyncio.run(scenario()) + + def test_normal_preview_stop_disarms_disconnect_cleanup() -> None: async def scenario() -> None: broker = FaceTrackingPreviewBroker() diff --git a/tests/runtime/test_face_tracking_preview_delivery.py b/tests/runtime/test_face_tracking_preview_delivery.py new file mode 100644 index 0000000..2a208f2 --- /dev/null +++ b/tests/runtime/test_face_tracking_preview_delivery.py @@ -0,0 +1,193 @@ +from __future__ import annotations + +import asyncio +import json +from dataclasses import dataclass, field + +from watcherobot.runtime.daemon.connections.registry import ExternalClientRole +from watcherobot.runtime.daemon.preview.delivery import ( + LatestPreviewDelivery, + PreviewRelayFrame, +) + + +class FakeClock: + def __init__(self, now: float = 10.0) -> None: + self.now = now + + def __call__(self) -> float: + return self.now + + +class FakeWebSocket: + def __init__(self, blocked: asyncio.Event | None = None) -> None: + self.frames: list[str | bytes] = [] + self._blocked = blocked + + async def send(self, frame: str | bytes) -> None: + if self._blocked is not None: + await self._blocked.wait() + self.frames.append(frame) + + +class FailingWebSocket: + def __init__(self) -> None: + self.send_attempts = 0 + + async def send(self, _frame: str | bytes) -> None: + self.send_attempts += 1 + raise RuntimeError("connection is closing") + + +@dataclass +class FakeConnection: + websocket: FakeWebSocket + metadata: dict[str, object] = field( + default_factory=lambda: { + "capabilities": ["face_tracking.preview.credit.v1"] + } + ) + + +class FakeRegistry: + def __init__(self, *connections: FakeConnection) -> None: + self.connections = list(connections) + + def connections_for(self, role: ExternalClientRole) -> list[FakeConnection]: + assert role is ExternalClientRole.DESKTOP + return list(self.connections) + + +def relay_frame(sequence: int, *, completed_at: float = 10.0) -> PreviewRelayFrame: + return PreviewRelayFrame( + stream_id=7, + sequence=sequence, + telemetry=json.dumps( + {"v": 1, "kind": "frame", "seq": sequence}, + separators=(",", ":"), + ), + image=b"FTW1" + sequence.to_bytes(4, "little"), + completed_at=completed_at, + ) + + +async def settle() -> None: + await asyncio.sleep(0) + await asyncio.sleep(0) + + +def test_credit_delivery_keeps_one_frame_in_flight_and_skips_to_latest() -> None: + async def scenario() -> None: + clock = FakeClock() + websocket = FakeWebSocket() + delivery = LatestPreviewDelivery( + FakeRegistry(FakeConnection(websocket)), + clock=clock, + ) + + assert delivery.offer(relay_frame(1)) == 1 + await settle() + assert len(websocket.frames) == 2 + + delivery.offer(relay_frame(2)) + delivery.offer(relay_frame(3)) + await settle() + assert len(websocket.frames) == 2 + + clock.now = 10.125 + assert delivery.acknowledge(websocket, sequence=1) + await settle() + + assert len(websocket.frames) == 4 + telemetry = json.loads(websocket.frames[2]) + assert telemetry["seq"] == 3 + assert telemetry["relay"][1] == 125.0 + assert telemetry["relay"][2] == 1 + assert websocket.frames[3] == relay_frame(3).image + assert delivery.snapshot()["pending_overwrites"] == 1 + await delivery.stop() + + asyncio.run(scenario()) + + +def test_slow_preview_consumer_does_not_block_another_browser() -> None: + async def scenario() -> None: + release_slow = asyncio.Event() + slow = FakeWebSocket(release_slow) + fast = FakeWebSocket() + delivery = LatestPreviewDelivery( + FakeRegistry(FakeConnection(slow), FakeConnection(fast)) + ) + + delivery.offer(relay_frame(11)) + await settle() + + assert slow.frames == [] + assert len(fast.frames) == 2 + release_slow.set() + await settle() + assert len(slow.frames) == 2 + await delivery.stop() + + asyncio.run(scenario()) + + +def test_unknown_or_repeated_ack_never_releases_newer_frame() -> None: + async def scenario() -> None: + websocket = FakeWebSocket() + delivery = LatestPreviewDelivery(FakeRegistry(FakeConnection(websocket))) + delivery.offer(relay_frame(21)) + await settle() + + assert not delivery.acknowledge(websocket, sequence=20) + delivery.offer(relay_frame(22)) + await settle() + assert len(websocket.frames) == 2 + + assert delivery.acknowledge(websocket, sequence=21) + await settle() + assert json.loads(websocket.frames[2])["seq"] == 22 + assert not delivery.acknowledge(websocket, sequence=21) + await delivery.stop() + + asyncio.run(scenario()) + + +def test_send_failure_discards_pending_frame_instead_of_retrying_backlog() -> None: + async def scenario() -> None: + websocket = FailingWebSocket() + connection = FakeConnection(websocket) # type: ignore[arg-type] + delivery = LatestPreviewDelivery(FakeRegistry(connection)) + + delivery.offer(relay_frame(31)) + delivery.offer(relay_frame(32)) + await settle() + + assert websocket.send_attempts == 1 + assert delivery.snapshot()["send_errors"] == 1 + assert delivery.snapshot()["active_consumers"] == 0 + await delivery.stop() + + asyncio.run(scenario()) + + +def test_connection_loss_does_not_flush_pending_legacy_frame() -> None: + async def scenario() -> None: + release = asyncio.Event() + websocket = FakeWebSocket(release) + connection = FakeConnection(websocket, metadata={}) + delivery = LatestPreviewDelivery(FakeRegistry(connection)) + + delivery.offer(relay_frame(41)) + delivery.offer(relay_frame(42)) + await settle() + delivery.connection_lost(websocket) + await settle() + release.set() + await settle() + + assert websocket.frames == [] + assert delivery.snapshot()["active_consumers"] == 0 + await delivery.stop() + + asyncio.run(scenario()) diff --git a/tests/runtime/test_face_tracking_udp_service.py b/tests/runtime/test_face_tracking_udp_service.py index bedbb12..068ffa5 100644 --- a/tests/runtime/test_face_tracking_udp_service.py +++ b/tests/runtime/test_face_tracking_udp_service.py @@ -12,6 +12,7 @@ build_preview_bundle, encode_preview_datagrams, ) +from watcherobot.runtime.daemon.preview.delivery import PreviewRelayFrame from watcherobot.runtime.daemon.preview.udp_service import FaceTrackingUdpPreviewService @@ -162,3 +163,57 @@ async def publish(frame: str | bytes) -> int: await service.stop() asyncio.run(scenario()) + + +def test_service_can_publish_an_atomic_preview_pair_to_credit_delivery() -> None: + async def scenario() -> None: + published: list[PreviewRelayFrame] = [] + + async def publish(frame: PreviewRelayFrame) -> int: + published.append(frame) + return 1 + + service = FaceTrackingUdpPreviewService( + session=connected_session(), + registry=FakeRegistry(), + bundle_publisher=publish, + port=0, + ) + await service.start() + for packet in packets(sequence=19): + service.handle_datagram(packet, (PEER_IP, 50000)) + await asyncio.sleep(0) + await asyncio.sleep(0) + + assert len(published) == 1 + assert published[0].stream_id == 123 + assert published[0].sequence == 19 + assert json.loads(published[0].telemetry)["seq"] == 19 + assert published[0].image == b"FTW1" + bytes(20) + b"jpeg" + await service.stop() + + asyncio.run(scenario()) + + +def test_service_refresh_rebinds_listener_and_accepts_new_preview() -> None: + async def scenario() -> None: + registry = FakeRegistry() + service = FaceTrackingUdpPreviewService( + session=connected_session(), registry=registry, port=0 + ) + await service.start() + previous_transport = service._transport + + await service.refresh_listener() + + assert service.bound_port > 0 + assert service._transport is not previous_transport + for packet in packets(sequence=23): + service.handle_datagram(packet, (PEER_IP, 50000)) + await asyncio.sleep(0) + await asyncio.sleep(0) + assert service.stats.published_frames == 1 + assert json.loads(registry.frames[0][1])["seq"] == 23 + await service.stop() + + asyncio.run(scenario()) diff --git a/tests/runtime/test_lan_pairing_protocol.py b/tests/runtime/test_lan_pairing_protocol.py index 286bd85..89ba407 100644 --- a/tests/runtime/test_lan_pairing_protocol.py +++ b/tests/runtime/test_lan_pairing_protocol.py @@ -10,6 +10,7 @@ LAN_PAIRING_VERSION, DeviceSessionEnd, HardwareHello, + MediaHello, PairAccept, PairBusy, PairCancel, @@ -18,8 +19,10 @@ build_device_state_event, build_hardware_hello_ack, build_hardware_hello_nack, + build_media_hello_ack, encode_udp_message, parse_hardware_hello, + parse_media_hello, parse_device_session_end, parse_udp_message, redact_sensitive_fields, @@ -49,6 +52,7 @@ def test_protocol_vectors_fix_name_version_and_message_shapes(vectors) -> None: assert isinstance(parse_udp_message(valid["pair_busy"]), PairBusy) assert isinstance(parse_udp_message(valid["pair_cancel"]), PairCancel) assert isinstance(parse_hardware_hello(valid["hardware_hello"]), HardwareHello) + assert isinstance(parse_media_hello(valid["media_hello"]), MediaHello) assert isinstance(parse_device_session_end(valid["session_end"]), DeviceSessionEnd) hello_data = valid["hardware_hello"]["data"] @@ -57,6 +61,11 @@ def test_protocol_vectors_fix_name_version_and_message_shapes(vectors) -> None: assert "capabilities" not in hello_data assert "pairing_code" not in hello_data assert valid["hello_ack"]["data"]["negotiated"]["audio_uplink"]["codec"] == "opus" + assert valid["hello_ack"]["data"]["negotiated"]["video_uplink"] == { + "transport": "websocket_sidecar", + "hello_type": "sys.media.hello", + "version": 1, + } @pytest.mark.parametrize("pairing_code", ["000000", "999999"]) @@ -174,3 +183,16 @@ def test_hardware_ack_nack_and_state_event_have_stable_envelopes(vectors) -> Non assert build_device_state_event( vectors["valid"]["device_state"]["data"], ) == vectors["valid"]["device_state"] + assert build_media_hello_ack() == vectors["valid"]["media_hello_ack"] + + +def test_media_hello_rejects_wrong_channel_version_and_unknown_fields(vectors) -> None: + valid = vectors["valid"]["media_hello"] + for field, value in (("channel", "audio"), ("version", 2)): + payload = {**valid, "data": {**valid["data"], field: value}} + with pytest.raises(PairingProtocolError, match=field): + parse_media_hello(payload) + + payload = {**valid, "data": {**valid["data"], "device_id": "legacy"}} + with pytest.raises(PairingProtocolError, match="unknown fields"): + parse_media_hello(payload) diff --git a/tests/runtime/test_pairing_session.py b/tests/runtime/test_pairing_session.py index 8c5af63..b411298 100644 --- a/tests/runtime/test_pairing_session.py +++ b/tests/runtime/test_pairing_session.py @@ -128,6 +128,20 @@ def test_accept_moves_to_connecting_and_never_exposes_secrets() -> None: assert SESSION_TOKEN not in snapshot_text +def test_default_discovery_window_accepts_a_slow_local_device_response() -> None: + session = make_session() + session.start_pairing( + pairing_code="123456", + target_mode="desktop_link", + websocket_port=8765, + now=10.0, + ) + + assert session.expire(now=25.0) is False + session.accept_device(make_accept(), peer_ip="192.168.3.25", now=25.0) + assert session.state is DevicePairingState.CONNECTING + + def test_hardware_hello_is_the_only_transition_to_online() -> None: session = make_session() session.start_pairing( @@ -203,7 +217,7 @@ def test_abnormal_disconnect_reserves_slot_for_same_session_reconnect() -> None: @pytest.mark.parametrize( ("advance", "expected_error"), [ - (("discovering", 20.0), "pairing_not_found"), + (("discovering", 40.0), "pairing_not_found"), (("connecting", 23.0), "device_connect_timeout"), (("reconnecting", 51.0), "reconnect_timeout"), ], diff --git a/tests/runtime/test_pairing_udp.py b/tests/runtime/test_pairing_udp.py index 5701671..c0fb723 100644 --- a/tests/runtime/test_pairing_udp.py +++ b/tests/runtime/test_pairing_udp.py @@ -251,7 +251,7 @@ async def scenario() -> None: states: list[dict[str, object]] = [] service = PairingUdpService( session=session, - clock=lambda: 20.0, + clock=lambda: 40.0, interface_provider=lambda: (), channel_factory=FakeChannelFactory(), state_listener=lambda snapshot: states.append(dict(snapshot)), diff --git a/tests/test_api.py b/tests/test_api.py index c248ffe..c5719c9 100644 --- a/tests/test_api.py +++ b/tests/test_api.py @@ -5,6 +5,7 @@ import pytest from watcherobot import Job +from watcherobot.audio import AudioLiveStream from watcherobot.errors import CommandError, WatcheRobotError from watcherobot.robot import WatcheRobot from watcherobot.protocol import ( @@ -29,6 +30,7 @@ def __init__(self): "motion", "audio", "audio.stream", + "audio.stream.live.v1", "light", "microphone", "camera.capture", @@ -38,6 +40,7 @@ def __init__(self): self.next_session_id = 100 self.closed = False self.audio_streams = [] + self.live_audio_frames = [] def set_callbacks(self, message_callback, binary_callback, disconnect_callback): self.message_callback = message_callback @@ -70,6 +73,23 @@ def send_command_nowait(self, message_type, data): future.set_result({"type": "sys.ack", "code": 0, "data": {}}) return future + def begin_live_audio_stream(self, *, stream_id, chunk_bytes=960): + self.live_audio_frames.append(("begin", stream_id, chunk_bytes)) + + def write_live_audio_stream(self, pcm, *, stream_id, sequence, chunk_bytes=960): + payload = bytes(pcm) + for offset in range(0, len(payload), chunk_bytes): + chunk = payload[offset : offset + chunk_bytes] + self.live_audio_frames.append(("data", stream_id, sequence, chunk)) + sequence = (sequence + 1) & 0xFFFFFFFF + return sequence + + def end_live_audio_stream(self, *, stream_id, sequence): + self.live_audio_frames.append(("end", stream_id, sequence)) + + def cancel_live_audio_stream(self, *, stream_id): + self.live_audio_frames.append(("cancel", stream_id)) + class FakeOpusDecoder: def decode(self, packet): @@ -131,6 +151,222 @@ def test_robot_supports_negotiated_capabilities(): robot.supports("") +def test_live_audio_stream_writes_incremental_pcm_without_a_total_size(): + transport = FakeTransport() + robot = WatcheRobot._from_transport(transport) + + stream = robot.audio.open_stream() + stream.write(b"\x01\x00" * 600) + stream.write(b"\x02\x00") + stream.close() + + assert stream.closed + assert transport.commands == [ + ( + "ctrl.audio.stream.begin", + { + "mode": "live", + "stream_id": stream.stream_id, + "sample_rate_hz": 24000, + "channels": 1, + "sample_width_bytes": 2, + }, + ) + ] + assert transport.live_audio_frames == [ + ("begin", stream.stream_id, 960), + ("data", stream.stream_id, 0, b"\x01\x00" * 480), + ("data", stream.stream_id, 1, b"\x01\x00" * 120), + ("data", stream.stream_id, 2, b"\x02\x00"), + ("end", stream.stream_id, 3), + ] + + +def test_live_audio_stream_context_aborts_on_exception(): + transport = FakeTransport() + robot = WatcheRobot._from_transport(transport) + + with pytest.raises(RuntimeError, match="boom"): + with robot.audio.open_stream() as stream: + stream.write(b"\x01\x00") + raise RuntimeError("boom") + + assert stream.closed + assert transport.live_audio_frames[-1] == ("cancel", stream.stream_id) + assert transport.commands[-1] == ("ctrl.audio.stop", {}) + + +def test_live_audio_stream_requires_negotiated_capability(): + transport = FakeTransport() + transport.capabilities = tuple( + capability + for capability in transport.capabilities + if capability != "audio.stream.live.v1" + ) + robot = WatcheRobot._from_transport(transport) + + with pytest.raises(WatcheRobotError, match="audio.stream.live.v1"): + robot.audio.open_stream() + + +def test_live_audio_stream_setup_failure_stops_the_authorized_device_stream(): + class BeginFailingTransport(FakeTransport): + def begin_live_audio_stream(self, *, stream_id, chunk_bytes=960): + raise WatcheRobotError("local media channel is disconnected") + + transport = BeginFailingTransport() + robot = WatcheRobot._from_transport(transport) + + with pytest.raises(WatcheRobotError, match="media channel"): + robot.audio.open_stream() + + assert transport.commands[-2][0] == "ctrl.audio.stream.begin" + assert transport.commands[-1] == ("ctrl.audio.stop", {}) + + +def test_live_audio_stream_close_and_abort_are_idempotent(): + transport = FakeTransport() + robot = WatcheRobot._from_transport(transport) + closed = robot.audio.open_stream() + closed.close() + closed.close() + closed.abort() + assert [frame[0] for frame in transport.live_audio_frames].count("end") == 1 + + aborted = robot.audio.open_stream() + aborted.abort() + aborted.abort() + aborted.close() + assert [frame[0] for frame in transport.live_audio_frames].count("cancel") == 1 + + +def test_live_audio_stream_abort_wakes_a_blocked_write_immediately(): + write_started = threading.Event() + write_released = threading.Event() + aborted = threading.Event() + + def write_callback(_stream, _payload, sequence): + write_started.set() + assert write_released.wait(1) + raise WatcheRobotError("audio stream failed: cancelled") + + def abort_callback(_stream): + aborted.set() + write_released.set() + + stream = AudioLiveStream( + 7, + write_callback, + lambda _stream, _sequence: None, + abort_callback, + ) + errors = [] + + def write_audio() -> None: + try: + stream.write(b"\x01\x00") + except Exception as error: + errors.append(error) + + writer = threading.Thread(target=write_audio) + writer.start() + assert write_started.wait(1) + + stream.abort() + writer.join(1) + + assert aborted.is_set() + assert stream.closed + assert not writer.is_alive() + assert len(errors) == 1 + assert isinstance(errors[0], WatcheRobotError) + + +def test_live_audio_stream_close_failure_falls_back_to_device_stop(): + class EndFailingTransport(FakeTransport): + def end_live_audio_stream(self, *, stream_id, sequence): + raise WatcheRobotError("connection lost while sending LAST") + + transport = EndFailingTransport() + robot = WatcheRobot._from_transport(transport) + stream = robot.audio.open_stream() + + with pytest.raises(WatcheRobotError, match="sending LAST"): + stream.close() + + assert stream.closed + assert transport.live_audio_frames[-1] == ("cancel", stream.stream_id) + assert transport.commands[-1] == ("ctrl.audio.stop", {}) + + +def test_audio_stop_aborts_a_live_stream_only_once(): + transport = FakeTransport() + robot = WatcheRobot._from_transport(transport) + stream = robot.audio.open_stream() + + robot.audio.stop() + + assert stream.closed + assert transport.commands.count(("ctrl.audio.stop", {})) == 1 + + +def test_live_audio_stream_rejects_partial_samples_and_wrong_format(): + transport = FakeTransport() + robot = WatcheRobot._from_transport(transport) + stream = robot.audio.open_stream() + with pytest.raises(ValueError, match="partial sample"): + stream.write(b"\x01") + stream.abort() + + with pytest.raises(ValueError, match="24000 Hz"): + robot.audio.open_stream(sample_rate_hz=48000) + + +def test_live_audio_stream_is_not_limited_by_the_file_playback_size(): + transport = FakeTransport() + robot = WatcheRobot._from_transport(transport) + stream = robot.audio.open_stream() + block = b"\x01\x00" * (512 * 1024) + for _ in range(5): + stream.write(block) + stream.close() + + sent_bytes = sum( + len(frame[3]) + for frame in transport.live_audio_frames + if frame[0] == "data" + ) + assert sent_bytes == 5 * 1024 * 1024 + + +def test_opening_microphone_aborts_active_live_speaker_stream_first(): + transport = FakeTransport() + robot = WatcheRobot._from_transport(transport, opus_decoder_factory=FakeOpusDecoder) + stream = robot.audio.open_stream() + + microphone = robot.microphone.open() + + assert stream.closed + assert transport.commands[-2:] == [ + ("ctrl.audio.stop", {}), + ("ctrl.microphone.open", {"sample_rate_hz": 16000}), + ] + microphone.close() + + +def test_opening_live_speaker_stream_closes_active_robot_microphone_first(): + transport = FakeTransport() + robot = WatcheRobot._from_transport(transport, opus_decoder_factory=FakeOpusDecoder) + microphone = robot.microphone.open() + + stream = robot.audio.open_stream() + + assert microphone.closed + assert transport.commands[-2][0] == "ctrl.microphone.close" + assert transport.commands[-1][0] == "ctrl.audio.stream.begin" + stream.abort() + + @pytest.mark.parametrize("duration_ms", [0, -1, 1.5, True, 65536]) def test_motion_move_to_rejects_invalid_duration_ms(duration_ms): robot = WatcheRobot._from_transport(FakeTransport())