"""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. Inbound JSON frames are rate-limited per connection (token bucket, ``INBOUND_RATE_PER_S`` / ``INBOUND_BURST``); binary media-upload chunks are exempt. 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 contextlib import logging import ssl import time from dataclasses import dataclass, field from typing import Any 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 # Inbound JSON control-frame rate limit (per connection, token bucket). # A legitimate app sends pings + occasional user-initiated requests — far # below 20/s sustained. Binary media-upload chunks are EXEMPT (see # ``_on_frame``): a 100 MB upload is 400 x 256 KiB frames in a tight loop # and would exhaust any sane bucket; uploads are bounded instead by the # per-frame ``max_size`` and the per-upload total cap (``media.py``). INBOUND_RATE_PER_S = 20.0 INBOUND_BURST = 40 # Close codes (4000-4999 are reserved for applications). CLOSE_AUTH_FAILED = 4401 CLOSE_REPLACED = 4402 CLOSE_RATE_LIMITED = 4403 CLOSE_SHUTDOWN = 1001 # Max length of a client-supplied device_id. MAX_DEVICE_ID_LEN = 128 async def dispatch_frame(adapter: Any, frame: protocol.Frame, device_id: str) -> None: """Shared inbound frame dispatch for the WS and HTTP transports (docs/19 §19.4). Transport-specific frames (``ping``/``pong``, binary media chunks) are handled by their own server before this is called; unknown types are ignored (forward-compat).""" if frame.type == protocol.TYPE_MESSAGE_SEND: await adapter.on_message_send(frame, device_id) elif frame.type == protocol.TYPE_CHANNEL_CREATE: await adapter.on_channel_create(frame, device_id) elif frame.type == protocol.TYPE_CHANNEL_RENAME: await adapter.on_channel_rename(frame, device_id) elif frame.type == protocol.TYPE_CHANNEL_SET_DEFAULT: await adapter.on_channel_set_default(frame, device_id) elif frame.type == protocol.TYPE_CHANNEL_FAVORITE: await adapter.on_channel_favorite(frame, device_id) elif frame.type == protocol.TYPE_CHANNEL_ICON: await adapter.on_channel_icon(frame, device_id) elif frame.type == protocol.TYPE_CHANNEL_SET_AUTOMATION: await adapter.on_channel_set_automation(frame, device_id) elif frame.type == protocol.TYPE_CHANNEL_DELETE: await adapter.on_channel_delete(frame, device_id) elif frame.type == protocol.TYPE_CHANNEL_LIST: await adapter.on_channel_list(frame, device_id) elif frame.type == protocol.TYPE_COMMANDS_CATALOG: await adapter.on_commands_catalog(frame, device_id) elif frame.type == protocol.TYPE_SEARCH: await adapter.on_search(frame, device_id) elif frame.type == protocol.TYPE_SYNC: await adapter.on_sync(frame, device_id) elif frame.type == protocol.TYPE_HISTORY: await adapter.on_history(frame, device_id) elif frame.type == protocol.TYPE_MESSAGE_DELETE: await adapter.on_message_delete(frame, device_id) elif frame.type == protocol.TYPE_MEDIA_UPLOAD_START: await adapter.on_media_upload_start(frame, device_id) elif frame.type == protocol.TYPE_MEDIA_UPLOAD_END: await adapter.on_media_upload_end(frame, device_id) elif frame.type == protocol.TYPE_MEDIA_PULL: await adapter.on_media_pull(frame, device_id) elif frame.type == protocol.TYPE_FCM_REGISTER: await adapter.on_fcm_register(frame, device_id) # Unknown types are ignored (forward-compat). class _TokenBucket: """Minimal token bucket (stdlib only). One instance per connection.""" __slots__ = ("rate", "burst", "tokens", "updated_at") def __init__(self, rate: float, burst: int): self.rate = rate self.burst = burst self.tokens = float(burst) self.updated_at = time.monotonic() def consume(self) -> bool: """Try to take one token. Refills at ``rate``/s up to ``burst``.""" now = time.monotonic() elapsed = now - self.updated_at if elapsed > 0: self.tokens = min(self.burst, self.tokens + elapsed * self.rate) self.updated_at = now if self.tokens >= 1.0: self.tokens -= 1.0 return True return False @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: str | None = None ntfy_topic: str | None = None connected_at: float = field(default_factory=time.time) rate_bucket: _TokenBucket = field( default_factory=lambda: _TokenBucket(INBOUND_RATE_PER_S, INBOUND_BURST) ) 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: Any | None = 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: ssl.SSLContext | None = 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() # Best-effort: the server is already closing; a failure here is # not actionable (nothing left to clean up besides the registry). with contextlib.suppress(Exception): await self._server.wait_closed() self._server = None for conn in list(self._connections.values()): # Best-effort: a socket that is already gone needs no handling. with contextlib.suppress(Exception): await conn.ws.close(code=CLOSE_SHUTDOWN, reason="gateway shutting down") 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) -> DeviceConnection | None: 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()): # Best-effort: a dead or stalled socket is skipped (it is # deregistered on its own close); one slow peer must not starve # the rest of the broadcast. with contextlib.suppress(Exception): await asyncio.wait_for(conn.ws.send(data), timeout=SEND_TIMEOUT_S) sent += 1 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, ConnectionClosed) as e: if isinstance(e, asyncio.TimeoutError): logger.warning("android: dropping socket with no hello (timeout)") await self._close_quiet(ws, 1000, "no hello") # A peer that vanished before hello needs no further handling. 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) > MAX_DEVICE_ID_LEN: 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. # Best-effort close of the superseded socket. with contextlib.suppress(Exception): await old.ws.close(code=CLOSE_REPLACED, reason="replaced by newer connection") ack = protocol.hello_ack( server_caps=self._adapter.server_caps(), sync_cursor=self._adapter._outbox.latest_cursor(), channels=self._adapter.channel_list(), # M5: lets the app dedupe sync-replayed notifications that # already woke this device via push (docs/08 §8.7). last_pushed_cursor=self._adapter._devices.last_pushed_cursor(device_id), ) try: await ws.send(ack.to_json()) # M7: tell late-joining clients the current gateway health state # (the startup broadcast only reaches clients already connected). await ws.send(protocol.status(self._adapter.gateway_status()).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: # ``_on_frame`` returns False once it has closed the socket # (rate limit); stop draining the buffered frames so a # flood doesn't re-trigger the error+close per frame. if not await self._on_frame(ws, device_id, raw): break except Exception as e: # A clean disconnect (ConnectionClosed) is the normal path and is # not worth a warning; anything else is unexpected. if not isinstance(e, ConnectionClosed): 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) -> bool: """Dispatch one inbound frame. Returns False once the socket has been closed (rate limit) so the caller stops draining buffered frames.""" # M4: binary frames are media upload chunks (raw bytes, no JSON # envelope). Route them to the active upload session. They are # EXEMPT from the inbound rate limit: a 100 MB upload is 400 x # 256 KiB frames in a tight loop, which would exhaust any sane # frame bucket. Uploads are bounded instead by the per-frame # ``max_size`` and the per-upload total cap (``media.py``). if isinstance(raw, (bytes, bytearray, memoryview)): await self._adapter.on_media_chunk(device_id, bytes(raw)) return True # Inbound rate limit (JSON control frames only). On exceed: error + # close, same pattern as auth rejection. conn = self._connection_for(ws) if conn is not None and not conn.rate_bucket.consume(): logger.warning("android: inbound rate limit exceeded for %s; closing", device_id) await self._send_quiet( ws, protocol.error(protocol.ERR_RATE_LIMITED, "inbound frame rate limit exceeded"), ) await self._close_quiet(ws, CLOSE_RATE_LIMITED, "rate limited") return False frame = protocol.Frame.from_json(raw) if frame is None: return True # 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)) return True await dispatch_frame(self._adapter, frame, device_id) return True # ── Helpers ─────────────────────────────────────────────────────────── def _connection_for(self, ws: ServerConnection) -> DeviceConnection | None: """The live registry entry for this exact socket (identity match, so a replaced socket never consumes the new connection's bucket).""" for conn in self._connections.values(): if conn.ws is ws: return conn return None async def _send_quiet(self, ws: ServerConnection, frame: protocol.Frame) -> None: # "Quiet" by contract: the caller does not care whether the peer was # still there (e.g. an error frame right before the close). with contextlib.suppress(Exception): await ws.send(frame.to_json()) 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: # "Quiet" by contract: closing an already-closed socket is a no-op. with contextlib.suppress(Exception): await ws.close(code=code, reason=reason)