431 lines
18 KiB
Python
431 lines
18 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(),
|
|
# 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 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_FAVORITE:
|
|
await self._adapter.on_channel_favorite(frame, device_id)
|
|
elif frame.type == protocol.TYPE_CHANNEL_ICON:
|
|
await self._adapter.on_channel_icon(frame, device_id)
|
|
elif frame.type == protocol.TYPE_CHANNEL_SET_AUTOMATION:
|
|
await self._adapter.on_channel_set_automation(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_COMMANDS_CATALOG:
|
|
await self._adapter.on_commands_catalog(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 |