"""WebSocket server, connection registry, and frame routing. Runs on the gateway's asyncio loop (started in ``AndroidAdapter.connect()``). Uses the ``websockets`` core dep (v15): ``websockets.asyncio.server.serve( handler, host, port, ssl=ctx)``. Per-connection handler: 1. Await first frame (bounded); must be ``hello {token, device_id, device_name, caps, fcm_token?}``. Verify token (constant-time) + allowlist. On failure: send ``error {code:"auth"}`` and close. 2. On success: register in the device registry (SQLite) + connection registry (``device_id -> {ws, caps, fcm_token}``), send ``hello.ack {server_caps, sync_cursor, channels[]}``. 3. Loop: decode frames, dispatch to adapter inbound handlers. 4. On close: deregister. Routing: ``broadcast(frame)`` sends to ALL connected devices (single-user model). Heartbeat via WS ping/pong (websockets built-in) + app-level ``ping``/``pong`` frames. Milestone M1. """ import asyncio import logging import ssl import time from dataclasses import dataclass, field from typing import Any, Dict, Optional from websockets.asyncio.server import ServerConnection, serve from websockets.exceptions import ConnectionClosed from . import protocol from .pairing import DeviceRegistry, verify_token logger = logging.getLogger(__name__) # How long a new socket may take to present its ``hello`` before we drop it. HELLO_TIMEOUT_S = 10.0 # Max time a single outbound send may block on a peer's full write buffer # before we give up on that peer (so one stalled client can't starve the # rest of the broadcast). The peer's own ping timeout reaps it afterwards. SEND_TIMEOUT_S = 10.0 # Close codes (4000-4999 are reserved for applications). CLOSE_AUTH_FAILED = 4401 CLOSE_REPLACED = 4402 CLOSE_SHUTDOWN = 1001 @dataclass class DeviceConnection: """One live, authenticated device socket.""" device_id: str device_name: str ws: ServerConnection caps: Dict[str, Any] = field(default_factory=dict) fcm_token: Optional[str] = None ntfy_topic: Optional[str] = None connected_at: float = field(default_factory=time.time) class WsServer: """The plugin's WebSocket server + live connection registry.""" def __init__(self, adapter: Any, devices: DeviceRegistry): self._adapter = adapter self._devices = devices self._server: Optional[Any] = None self._connections: Dict[str, DeviceConnection] = {} self._lock = asyncio.Lock() # ── Lifecycle ───────────────────────────────────────────────────────── async def start(self) -> None: """Bind and start serving. Raises on bind failure (adapter maps it to a retryable fatal error).""" adapter = self._adapter ssl_ctx: Optional[ssl.SSLContext] = None if adapter.ws_cert and adapter.ws_key: try: ssl_ctx = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER) ssl_ctx.load_cert_chain(adapter.ws_cert, adapter.ws_key) except Exception as e: adapter._set_fatal_error( "tls_config", f"WS TLS cert/key invalid: {e}", retryable=False ) raise try: self._server = await serve( self._handler, adapter.host, adapter.port, ssl=ssl_ctx, # Media uploads (M4) are chunked binary frames; allow the # configured max upload size per frame. max_size=adapter.max_upload_bytes, # WS-level heartbeat: dead peers are reaped by websockets. ping_interval=20, ping_timeout=20, open_timeout=10, ) except OSError as e: adapter._set_fatal_error( "bind_failed", f"WS bind on {adapter.host}:{adapter.port} failed: {e}", retryable=True, ) raise scheme = "wss" if ssl_ctx else "ws" logger.info( "android: WS server listening on %s://%s:%s/ws", scheme, adapter.host, adapter.port, ) async def stop(self) -> None: """Stop serving and close all device sockets.""" if self._server is not None: self._server.close() try: await self._server.wait_closed() except Exception: pass self._server = None for conn in list(self._connections.values()): try: await conn.ws.close(code=CLOSE_SHUTDOWN, reason="gateway shutting down") except Exception: pass self._connections.clear() # ── Registry ────────────────────────────────────────────────────────── @property def connections(self) -> Dict[str, DeviceConnection]: return dict(self._connections) def has_devices(self) -> bool: return bool(self._connections) def device_ids(self) -> list: return list(self._connections.keys()) def connection(self, device_id: str) -> Optional[DeviceConnection]: return self._connections.get(device_id) # ── Outbound ────────────────────────────────────────────────────────── async def broadcast(self, frame: protocol.Frame) -> int: """Send a frame to every connected device. Returns devices reached. Best-effort: a dead or stalled socket is skipped (deregistered on its own close) so one slow peer can't starve the others.""" data = frame.to_json() sent = 0 for conn in list(self._connections.values()): try: await asyncio.wait_for(conn.ws.send(data), timeout=SEND_TIMEOUT_S) sent += 1 except Exception: pass return sent async def send_to(self, device_id: str, frame: protocol.Frame) -> bool: """Send a frame to one device (request responses / errors).""" conn = self._connections.get(device_id) if conn is None: return False try: await asyncio.wait_for(conn.ws.send(frame.to_json()), timeout=SEND_TIMEOUT_S) return True except Exception: return False # ── Per-connection handler ──────────────────────────────────────────── async def _handler(self, ws: ServerConnection) -> None: # 1. hello auth ----------------------------------------------------- try: raw = await asyncio.wait_for(ws.recv(), timeout=HELLO_TIMEOUT_S) except asyncio.TimeoutError: logger.warning("android: dropping socket with no hello (timeout)") await self._close_quiet(ws, 1000, "no hello") return except ConnectionClosed: return frame = protocol.Frame.from_json(raw) if frame is None or frame.type != protocol.TYPE_HELLO: await self._reject(ws, "first frame must be hello") return payload = frame.payload if not verify_token(payload.get("token"), self._adapter.token): peer = getattr(ws, "remote_address", None) logger.warning("android: hello rejected: invalid token (peer=%s)", peer) await self._reject(ws, "invalid token") return device_id = str(payload.get("device_id") or "").strip() if not device_id or len(device_id) > 128: await self._reject(ws, "device_id required") return if ( not self._adapter.allow_all and self._adapter.allowed_users and device_id not in self._adapter.allowed_users ): logger.warning("android: hello rejected: device %s not allowlisted", device_id) await self._reject(ws, "device not allowed") return device_name = str(payload.get("device_name") or device_id)[:120] caps = payload.get("caps") if not isinstance(caps, dict): caps = {} fcm_token = payload.get("fcm_token") if not isinstance(fcm_token, str): fcm_token = None ntfy_topic = payload.get("ntfy_topic") if not isinstance(ntfy_topic, str): ntfy_topic = None # 2. register -------------------------------------------------------- try: self._devices.upsert(device_id, device_name, caps, fcm_token, ntfy_topic) except Exception: logger.warning("android: device registry upsert failed", exc_info=True) conn = DeviceConnection( device_id=device_id, device_name=device_name, ws=ws, caps=caps, fcm_token=fcm_token, ntfy_topic=ntfy_topic, ) async with self._lock: old = self._connections.pop(device_id, None) self._connections[device_id] = conn if old is not None: # Same device re-paired from a new socket: the new one wins. try: await old.ws.close(code=CLOSE_REPLACED, reason="replaced by newer connection") except Exception: pass ack = protocol.hello_ack( server_caps=self._adapter.server_caps(), sync_cursor=self._adapter._outbox.latest_cursor(), channels=self._adapter.channel_list(), ) try: await ws.send(ack.to_json()) except Exception: return logger.info("android: device paired: %s (%s)", device_name, device_id) # 3. frame loop ------------------------------------------------------ try: async for raw in ws: await self._on_frame(ws, device_id, raw) except ConnectionClosed: pass except Exception: logger.warning("android: frame loop error for %s", device_id, exc_info=True) finally: async with self._lock: current = self._connections.get(device_id) if current is not None and current.ws is ws: self._connections.pop(device_id, None) # M4: drop in-flight upload temp files for this socket. try: self._adapter.on_connection_closed(device_id) except Exception: logger.warning("android: connection cleanup failed for %s", device_id, exc_info=True) logger.info("android: device disconnected: %s", device_id) # ── Inbound dispatch ────────────────────────────────────────────────── async def _on_frame(self, ws: ServerConnection, device_id: str, raw: Any) -> None: # M4: binary frames are media upload chunks (raw bytes, no JSON # envelope). Route them to the active upload session. if isinstance(raw, (bytes, bytearray, memoryview)): await self._adapter.on_media_chunk(device_id, bytes(raw)) return frame = protocol.Frame.from_json(raw) if frame is None: return # malformed JSON: ignore (forward-compat) if frame.type == protocol.TYPE_PING: ts = frame.payload.get("ts") await self._send_quiet(ws, protocol.pong(ts if isinstance(ts, int) else None)) elif frame.type == protocol.TYPE_MESSAGE_SEND: await self._adapter.on_message_send(frame, device_id) elif frame.type == protocol.TYPE_CHANNEL_CREATE: await self._adapter.on_channel_create(frame, device_id) elif frame.type == protocol.TYPE_CHANNEL_RENAME: await self._adapter.on_channel_rename(frame, device_id) elif frame.type == protocol.TYPE_CHANNEL_SET_DEFAULT: await self._adapter.on_channel_set_default(frame, device_id) elif frame.type == protocol.TYPE_CHANNEL_DELETE: await self._adapter.on_channel_delete(frame, device_id) elif frame.type == protocol.TYPE_CHANNEL_LIST: await self._adapter.on_channel_list(frame, device_id) elif frame.type == protocol.TYPE_SEARCH: await self._adapter.on_search(frame, device_id) elif frame.type == protocol.TYPE_SYNC: await self._adapter.on_sync(frame, device_id) elif frame.type == protocol.TYPE_MEDIA_UPLOAD_START: await self._adapter.on_media_upload_start(frame, device_id) elif frame.type == protocol.TYPE_MEDIA_UPLOAD_END: await self._adapter.on_media_upload_end(frame, device_id) elif frame.type == protocol.TYPE_MEDIA_PULL: await self._adapter.on_media_pull(frame, device_id) elif frame.type == protocol.TYPE_FCM_REGISTER: await self._adapter.on_fcm_register(frame, device_id) # Unknown types are ignored (forward-compat). # ── Helpers ─────────────────────────────────────────────────────────── async def _send_quiet(self, ws: ServerConnection, frame: protocol.Frame) -> None: try: await ws.send(frame.to_json()) except Exception: pass async def _reject(self, ws: ServerConnection, reason: str) -> None: await self._send_quiet(ws, protocol.error(protocol.ERR_AUTH, reason)) await self._close_quiet(ws, CLOSE_AUTH_FAILED, "auth failed") async def _close_quiet(self, ws: ServerConnection, code: int, reason: str) -> None: try: await ws.close(code=code, reason=reason) except Exception: pass