Gateway plugin: - media.upload (chunked binary) -> size/sha256 verify + MIME re-sniff -> cache_*_from_bytes -> media.upload.ack - media.offer / media.pull (chunked) for agent-sent media, delivery-path security re-checked at pull time - send_* overrides mint media_id and emit media.offer - message.send media_refs resolve to cached inbound media - per-send + per-chunk timeouts so a stalled peer can't starve the rest App (Kotlin CMP): - Protocol: media frame types/payloads/builders - GatewayClient: binary session, uploadMedia (chunked + streaming sha256), pullMedia serialized via Mutex so concurrent offers don't interleave - ChatStore/IrisController: MediaItem, attachments, auto-pull on offer - Platform media: SAF picker, ExoPlayer (audio mini-player + video), image loader, FileProvider document open (Android); AWT-free desktop actuals - ChatScreen: attach button + chips, media rendering, keyboard dismiss on send UI polish: - preserve image aspect ratio (no stretching), cap dominant dimension - adjustResize so only chat content squeezes for the keyboard - clear focus (hide keyboard) on send Docs: media.upload.ack in 04-wire-protocol.md + frames.schema.json + 07-media.md; M4 marked complete in 14-milestones.md. Tests: 17-test tests/gateway/test_android.py suite passes.
349 lines
14 KiB
Python
349 lines
14 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.
|
|
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 == "fcm.register":
|
|
fcm_token = frame.payload.get("fcm_token")
|
|
ntfy_topic = frame.payload.get("ntfy_topic")
|
|
if isinstance(fcm_token, str) or isinstance(ntfy_topic, str):
|
|
try:
|
|
self._devices.update_push_tokens(
|
|
device_id,
|
|
fcm_token=fcm_token if isinstance(fcm_token, str) else None,
|
|
ntfy_topic=ntfy_topic if isinstance(ntfy_topic, str) else None,
|
|
)
|
|
except Exception:
|
|
logger.warning("android: fcm.register update failed", exc_info=True)
|
|
# 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 |