Files
iris_x_hermes/gateway-plugin/ws_server.py
T

420 lines
17 KiB
Python

"""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_HISTORY:
await self._adapter.on_history(frame, device_id)
elif frame.type == protocol.TYPE_MESSAGE_DELETE:
await self._adapter.on_message_delete(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