M1+M2: gateway core loop + agent transparency
M1 (gateway core loop / text round-trip): - WS server (ws_server.py): bind, hello auth (constant-time), hello.ack, heartbeat, connection registry - pairing.py: token generation, pairing store, QR payload - adapter.py: send() -> message frame; inbound message.send -> MessageEvent -> handle_message - app: Connect screen, GatewayClient (connect + reconnect), ChatScreen send/render, SecureStore (Android/Desktop) - tests/ws_probe.py: probe harness driving a real turn M2 (streaming + reasoning + tools + commentary): - protocol.py: M2 frame types (message.start/update/stop, tool.start/progress/end, commentary) - adapter.py: per-chat turn-state machine; classify outbound into frames; _split_reasoning; tool-line parsing - reasoning in streaming: capture via on_stream_delta hook (kind=reasoning, gated by plugins.stream_reasoning_deltas) with a FIFO barrier, attach to message.stop - app: live streaming bubble, ReasoningBlock (collapse + copy), ToolCard (Everything/Truncated/Nothing), dimmed commentary, typing - docs/14-milestones.md: M1/M2 marked done; reasoning note corrected
This commit is contained in:
1 parent
59acf66c89
commit
218c50d688
21 files changed
+3437
-129
No files matched your search
+300
-15
@@ -1,24 +1,309 @@
|
||||
"""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.serve(handler, host,
|
||||
port, ssl=ctx)``.
|
||||
Uses the ``websockets`` core dep (v15): ``websockets.asyncio.server.serve(
|
||||
handler, host, port, ssl=ctx)``.
|
||||
|
||||
Per-connection handler:
|
||||
1. Await first frame; 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 connection registry (``device_id ->
|
||||
{ws, caps, fcm_token}``), send ``hello.ack {server_caps, sync_cursor,
|
||||
channels[]}``.
|
||||
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; if no devices remain, ensure pending outbox
|
||||
frames have push fired.
|
||||
4. On close: deregister.
|
||||
|
||||
Routing: ``emit(chat_id, frame)`` broadcasts to ALL connected devices
|
||||
(single-user model). Heartbeat via WS ping/pong + app-level ping/pong.
|
||||
Backpressure: bounded per-connection send queue; coalesce ``message.update``
|
||||
under pressure, never drop ``message``/``tool.end``/``notification``.
|
||||
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
|
||||
|
||||
# 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())
|
||||
|
||||
# ── Outbound ──────────────────────────────────────────────────────────
|
||||
|
||||
async def broadcast(self, frame: protocol.Frame) -> int:
|
||||
"""Send a frame to every connected device. Returns devices reached.
|
||||
Best-effort: a dead socket is skipped (deregistered on its own close)."""
|
||||
data = frame.to_json()
|
||||
sent = 0
|
||||
for conn in list(self._connections.values()):
|
||||
try:
|
||||
await conn.ws.send(data)
|
||||
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 conn.ws.send(frame.to_json())
|
||||
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=0, # outbox lands in M3; cursor starts at 0
|
||||
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)
|
||||
logger.info("android: device disconnected: %s", device_id)
|
||||
|
||||
# ── Inbound dispatch ──────────────────────────────────────────────────
|
||||
|
||||
async def _on_frame(self, ws: ServerConnection, device_id: str, raw: Any) -> None:
|
||||
frame = protocol.Frame.from_json(raw)
|
||||
if frame is None:
|
||||
return # malformed / unknown binary: 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 == "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
|
||||
Reference in new issue
Block a user