"""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 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 # 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 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: Optional[str] = None ntfy_topic: Optional[str] = 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: 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()) # 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 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) -> 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)) 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). return True # ── Helpers ─────────────────────────────────────────────────────────── def _connection_for(self, ws: ServerConnection) -> Optional[DeviceConnection]: """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: 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