HTTP transport: drop WS server, offline send queue + dead-stream watchdog
Gateway (docs/19): - Remove ws_server.py; frame dispatch factored into dispatch.py - http_server: media upload/pull, pairing over HTTP - protocol: media frames mirrored; tests + ws_probe updated for HTTP App: - HttpGateway: postFrame/uploadMedia/pullMedia no longer throw on network failure (PostResult ok=false / Result.failure) — uncaught SocketTimeoutException on Dispatchers.Default crashed the app - GatewayClient: dead-stream watchdog (health probe every 10s, 2 failures -> redial in ~20s instead of the 45s SSE read timeout); state flips to Reconnecting when the stream dies, restored from the last hello.ack on long-poll success; poke() + backoff reset on app resume (MainActivity.onResume) - Offline sends: composer enabled while disconnected; a send with no response (status 0) stays queued (Pending) and is auto-resent on the next (re)connect after a 2s outbox-replay grace; gateway 4xx rejections fail the bubble (tap to retry, no auto-loop) - ChatStore: echo-replace and thread-relocate also match Failed bubbles (POST response lost in a network drop); loadHistory dedupes local failed bubbles the server already has; failMessage() - MainActivity: poke() on resume so a backgrounded app reconnects promptly instead of waiting out the backoff
This commit is contained in:
1 parent
2349a95dd4
commit
e6015033b6
22 files changed
+2804
-2439
No files matched your search
+78
-272
@@ -1,18 +1,20 @@
|
||||
"""
|
||||
Android Platform Adapter for Hermes Agent (Iris x Hermes).
|
||||
|
||||
A plugin-based gateway adapter that runs a WebSocket server *inside* the
|
||||
A plugin-based gateway adapter that runs an HTTP server *inside* the
|
||||
``hermes gateway`` process. The native Android / Desktop app connects to it
|
||||
with a pairing token and talks to the agent over a single WS transport
|
||||
(chat, streaming, tools, media, pairing, push-token).
|
||||
with a pairing token and talks to the agent over a single HTTP transport
|
||||
(chat, streaming, tools, media, pairing, push-token): JSON frames via
|
||||
``POST /v1/frame``, events via SSE ``GET /v1/events`` (or long-poll), and
|
||||
media via ``POST /v1/media`` / ``GET /v1/media/{id}`` (docs/19).
|
||||
|
||||
Zero new Python dependencies: ``websockets`` and ``httpx`` are hermes core
|
||||
deps. Zero hermes-core changes.
|
||||
Zero new Python dependencies: ``httpx`` is a hermes core dep. Zero
|
||||
hermes-core changes.
|
||||
|
||||
Milestone M1: the gateway core loop (text round-trip). The WS server binds
|
||||
and authenticates devices (``hello`` with constant-time token check), the
|
||||
adapter emits ``message`` frames from ``send()`` and turns inbound
|
||||
``message.send`` frames into ``MessageEvent``s for ``handle_message()``.
|
||||
Milestone M1: the gateway core loop (text round-trip). The server binds and
|
||||
authenticates devices (constant-time token check), the adapter emits
|
||||
``message`` frames from ``send()`` and turns inbound ``message.send`` frames
|
||||
into ``MessageEvent``s for ``handle_message()``.
|
||||
|
||||
Milestone M2: agent transparency. ``send()``/``edit_message()`` are mapped to
|
||||
``message.start``/``message.update``/``message.stop`` (streaming), tool
|
||||
@@ -22,13 +24,13 @@ reasoning prefix is split into a ``reasoning`` field. Outbox and search land
|
||||
in M3; media, push, and desktop land in later milestones (see
|
||||
``docs/14-milestones.md``).
|
||||
|
||||
Milestone M4: media. Inbound ``media.upload`` (chunked binary frames) is
|
||||
reassembled in a temp file, verified (size + sha256), re-sniffed, and cached
|
||||
via hermes ``cache_*_from_bytes``; the resulting refs attach to the next
|
||||
Milestone M4: media. Inbound uploads (``POST /v1/media``) are streamed to a
|
||||
temp file, verified (size + sha256), re-sniffed, and cached via hermes
|
||||
``cache_*_from_bytes``; the resulting refs attach to the next
|
||||
``message.send`` as ``MessageEvent.media_urls``. Outbound ``send_*`` calls
|
||||
register the (delivery-validated) file in the media registry and emit
|
||||
``media.offer``; ``media.pull`` streams the file back as chunked binary
|
||||
frames, re-checking ``validate_media_delivery_path`` at pull time.
|
||||
``media.offer``; ``GET /v1/media/{id}`` streams the file back, re-checking
|
||||
``validate_media_delivery_path`` at pull time.
|
||||
|
||||
Milestone M5: push + offline. Frames with no live subscriber are parked in
|
||||
the outbox (M3) AND wake the device via the push backend (``push.py``: FCM
|
||||
@@ -127,7 +129,6 @@ from .pairing import ( # noqa: E402
|
||||
qr_payload,
|
||||
)
|
||||
from .push import NtfyBackend, PushBackend, build_push_backend # noqa: E402
|
||||
from .ws_server import WsServer # noqa: E402
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Slash-command catalog (the app's "/" drawer)
|
||||
@@ -459,8 +460,6 @@ DEFAULT_PUSH_BACKEND = "fcm"
|
||||
DEFAULT_OUTBOX_RETENTION_HOURS = 72
|
||||
DEFAULT_MAX_UPLOAD_BYTES = 100 * 1024 * 1024 # 100 MB
|
||||
|
||||
# Max length of a client-supplied media_ref (mu_*/md_* ids are short).
|
||||
MAX_MEDIA_REF_LEN = 64
|
||||
# How often (seconds) the outbox-prune "storage reclaimed" notice may repeat.
|
||||
_PRUNE_NOTIFY_INTERVAL_S = 3600.0
|
||||
|
||||
@@ -802,15 +801,12 @@ class _TurnState:
|
||||
|
||||
|
||||
def check_requirements() -> bool:
|
||||
"""PASSIVE dependency probe: ``websockets`` importable + token set.
|
||||
"""PASSIVE dependency probe: token set.
|
||||
|
||||
Must be side-effect free (called from ``hermes setup`` / ``status`` /
|
||||
dashboard readiness). Never installs.
|
||||
dashboard readiness). Never installs. The HTTP transport is stdlib-only,
|
||||
so there is no extra dependency to probe.
|
||||
"""
|
||||
try:
|
||||
import websockets # noqa: F401 (core dep)
|
||||
except Exception:
|
||||
return False
|
||||
return bool(_get_scoped_secret("ANDROID_TOKEN"))
|
||||
|
||||
|
||||
@@ -856,9 +852,9 @@ def _env_enablement() -> dict | None:
|
||||
host = os.getenv("ANDROID_WS_HOST", "").strip()
|
||||
if host:
|
||||
seed["host"] = host
|
||||
port_raw = os.getenv("ANDROID_WS_PORT", "").strip()
|
||||
if port_raw:
|
||||
seed["port"] = _parse_port(port_raw)
|
||||
http_port_raw = os.getenv("ANDROID_HTTP_PORT", "").strip()
|
||||
if http_port_raw:
|
||||
seed["http_port"] = _parse_port(http_port_raw)
|
||||
push = os.getenv("ANDROID_PUSH_BACKEND", "").strip().lower()
|
||||
if push:
|
||||
seed["push_backend"] = push
|
||||
@@ -1032,10 +1028,16 @@ def interactive_setup() -> None:
|
||||
else:
|
||||
print_info("Existing ANDROID_TOKEN found (not shown).")
|
||||
|
||||
host = prompt("WS bind host", default=get_env_value("ANDROID_WS_HOST") or DEFAULT_HOST)
|
||||
host = prompt("Bind host", default=get_env_value("ANDROID_WS_HOST") or DEFAULT_HOST)
|
||||
save_env_value("ANDROID_WS_HOST", host or DEFAULT_HOST)
|
||||
port = prompt("WS port", default=str(_parse_port(get_env_value("ANDROID_WS_PORT") or "")))
|
||||
save_env_value("ANDROID_WS_PORT", str(_parse_port(port)))
|
||||
# _parse_port falls back to DEFAULT_PORT (8790) for empty input, so the
|
||||
# HTTP default must be applied explicitly (docs/19: 8791).
|
||||
http_port_raw = (get_env_value("ANDROID_HTTP_PORT") or "").strip()
|
||||
port = prompt(
|
||||
"HTTP port",
|
||||
default=str(int(http_port_raw) if http_port_raw.isdigit() else DEFAULT_HTTP_PORT),
|
||||
)
|
||||
save_env_value("ANDROID_HTTP_PORT", str(_parse_port(port)))
|
||||
backend = prompt(
|
||||
"Push backend (fcm/ntfy)",
|
||||
default=get_env_value("ANDROID_PUSH_BACKEND") or DEFAULT_PUSH_BACKEND,
|
||||
@@ -1063,19 +1065,20 @@ def interactive_setup() -> None:
|
||||
|
||||
|
||||
class AndroidAdapter(BasePlatformAdapter):
|
||||
"""WebSocket-backed adapter for the native Iris Android / Desktop app.
|
||||
"""HTTP-backed adapter for the native Iris Android / Desktop app.
|
||||
|
||||
M1: the WS server (``ws_server.WsServer``) authenticates devices with the
|
||||
pairing token, the connection registry tracks live sockets, ``send()``
|
||||
The HTTP server (``http_server.HttpServer``) authenticates devices with
|
||||
the pairing token, the device registry tracks live subscribers, ``send()``
|
||||
emits ``message`` frames, and inbound ``message.send`` frames become
|
||||
``MessageEvent``s for ``handle_message()``.
|
||||
"""
|
||||
|
||||
# WS has no message-size limit. The stream consumer resolves its per-chat
|
||||
# chunking budget via ``max_message_length_for_chat`` -> this attribute
|
||||
# (defaulting to 4096 when unset), which would split long replies — and
|
||||
# complete HTML artifacts — across multiple fence-reopened messages. A
|
||||
# large cap disables that chunking so a reply arrives as a single message.
|
||||
# The HTTP transport has no per-message size limit. The stream consumer
|
||||
# resolves its per-chat chunking budget via ``max_message_length_for_chat``
|
||||
# -> this attribute (defaulting to 4096 when unset), which would split
|
||||
# long replies — and complete HTML artifacts — across multiple
|
||||
# fence-reopened messages. A large cap disables that chunking so a reply
|
||||
# arrives as a single message.
|
||||
MAX_MESSAGE_LENGTH = 1_000_000
|
||||
|
||||
def __init__(self, config, **kwargs):
|
||||
@@ -1088,12 +1091,10 @@ class AndroidAdapter(BasePlatformAdapter):
|
||||
|
||||
extra = getattr(config, "extra", {}) or {}
|
||||
|
||||
# Connection settings (env vars override config.yaml)
|
||||
# Connection settings (env vars override config.yaml). The bind host
|
||||
# is shared with the (legacy) WS-era env var name for compatibility.
|
||||
self.host = os.getenv("ANDROID_WS_HOST", "").strip() or extra.get("host", DEFAULT_HOST)
|
||||
self.port = _parse_port(
|
||||
os.getenv("ANDROID_WS_PORT", "") or str(extra.get("port", DEFAULT_PORT))
|
||||
)
|
||||
# docs/19: HTTP fallback leg (same bind host as the WS; optional TLS).
|
||||
# docs/19: HTTP transport (the only device-facing transport; optional TLS).
|
||||
self.http_port = _parse_port(
|
||||
os.getenv("ANDROID_HTTP_PORT", "") or str(extra.get("http_port", DEFAULT_HTTP_PORT))
|
||||
)
|
||||
@@ -1127,8 +1128,6 @@ class AndroidAdapter(BasePlatformAdapter):
|
||||
self.home_channel_name = DEFAULT_HOME_CHANNEL_NAME
|
||||
|
||||
# TLS (optional)
|
||||
self.ws_cert = _get_scoped_secret("ANDROID_WS_CERT") or extra.get("ws_cert", "")
|
||||
self.ws_key = _get_scoped_secret("ANDROID_WS_KEY") or extra.get("ws_key", "")
|
||||
self.http_cert = _get_scoped_secret("ANDROID_HTTP_CERT") or extra.get("http_cert", "")
|
||||
self.http_key = _get_scoped_secret("ANDROID_HTTP_KEY") or extra.get("http_key", "")
|
||||
|
||||
@@ -1141,13 +1140,12 @@ class AndroidAdapter(BasePlatformAdapter):
|
||||
|
||||
# Runtime state
|
||||
self._devices = DeviceRegistry(get_hermes_home() / "android" / "devices.db")
|
||||
self._ws_server = WsServer(self, self._devices)
|
||||
# docs/19: HTTP fallback leg (inert until the app uses it; a bind
|
||||
# failure disables it without affecting the WS).
|
||||
# docs/19: HTTP transport (the only device-facing transport).
|
||||
self._http_server = HttpServer(self, self._devices)
|
||||
# docs/19 §19.7: reply sinks for in-flight HTTP requests — while a
|
||||
# POST /v1/frame is being dispatched, the handler's point-to-point
|
||||
# replies are captured here and returned as the HTTP response.
|
||||
# Entry shape: (sink queue, abandoned event).
|
||||
self._http_reply_sinks: dict[str, tuple[queue.Queue, threading.Event]] = {}
|
||||
self._connected = False
|
||||
# M2: per-chat turn state for outbound frame classification.
|
||||
@@ -1191,7 +1189,7 @@ class AndroidAdapter(BasePlatformAdapter):
|
||||
# ── Connection lifecycle ──────────────────────────────────────────────
|
||||
|
||||
async def connect(self, *, is_reconnect: bool = False) -> bool:
|
||||
"""Bring the platform up: bind the WS server on host:port."""
|
||||
"""Bring the platform up: bind the HTTP server on host:http_port."""
|
||||
if not self.token:
|
||||
logger.error("android: ANDROID_TOKEN must be set")
|
||||
self._set_fatal_error(
|
||||
@@ -1201,41 +1199,24 @@ class AndroidAdapter(BasePlatformAdapter):
|
||||
)
|
||||
return False
|
||||
|
||||
# Prevent two profiles from binding the same port/identity.
|
||||
try:
|
||||
from gateway.status import acquire_scoped_lock
|
||||
|
||||
lock_key = f"{self.host}:{self.port}"
|
||||
if not acquire_scoped_lock("android", lock_key):
|
||||
logger.error(
|
||||
"android: %s:%s already in use by another profile", self.host, self.port
|
||||
)
|
||||
self._set_fatal_error(
|
||||
"lock_conflict",
|
||||
"WS port in use by another profile",
|
||||
retryable=False,
|
||||
)
|
||||
return False
|
||||
self._lock_key = lock_key
|
||||
except ImportError:
|
||||
self._lock_key = None # status module not available (e.g. tests)
|
||||
|
||||
try:
|
||||
await self._ws_server.start()
|
||||
except Exception:
|
||||
self._connected = False
|
||||
return False
|
||||
|
||||
# docs/19: start the HTTP fallback leg next to the WS. Bind failure
|
||||
# is NON-fatal (unlike the WS): the plugin keeps working WS-only.
|
||||
# The HTTP server is the only device-facing transport, so a bind
|
||||
# failure is fatal (the app has no other way to reach the gateway).
|
||||
# start() never raises; it disables the leg and logs on failure.
|
||||
await self._http_server.start()
|
||||
if not self._http_server.enabled:
|
||||
logger.error("android: HTTP server failed to bind %s:%s", self.host, self.http_port)
|
||||
self._set_fatal_error(
|
||||
"bind_failed",
|
||||
f"HTTP port {self.http_port} unavailable",
|
||||
retryable=False,
|
||||
)
|
||||
return False
|
||||
|
||||
# M5: announce gateway health to connected clients (none yet at
|
||||
# startup; the frame + plumbing exist for future transitions).
|
||||
# Reset in case this adapter instance previously went down (the
|
||||
# gateway may reconnect the same adapter after a fatal error).
|
||||
self._gateway_status = protocol.STATUS_ONLINE
|
||||
await self._ws_server.broadcast(protocol.status(self._gateway_status))
|
||||
await self._http_server.fanout(protocol.status(self._gateway_status), cursor=None)
|
||||
|
||||
# M3: ensure the default (home) channel exists in the directory so the
|
||||
@@ -1257,29 +1238,18 @@ class AndroidAdapter(BasePlatformAdapter):
|
||||
|
||||
self._connected = True
|
||||
self._mark_connected()
|
||||
logger.info("android: connected; WS server on %s:%s", self.host, self.port)
|
||||
logger.info("android: connected; HTTP server on %s:%s", self.host, self.http_port)
|
||||
return True
|
||||
|
||||
async def disconnect(self) -> None:
|
||||
"""Tear down the platform: stop the server, close device sockets."""
|
||||
"""Tear down the platform: stop the server, close device streams."""
|
||||
# Tell live clients the gateway is going away (restart/shutdown) so
|
||||
# the app can distinguish a clean gateway teardown from a plain
|
||||
# network drop: the "Gateway restarting" chat notice is shown only
|
||||
# when this frame was received (docs/04 §status).
|
||||
self._gateway_status = protocol.STATUS_RESTARTING
|
||||
with contextlib.suppress(Exception):
|
||||
await self._ws_server.broadcast(protocol.status(self._gateway_status))
|
||||
await self._http_server.fanout(protocol.status(self._gateway_status), cursor=None)
|
||||
with contextlib.suppress(ImportError):
|
||||
from gateway.status import release_scoped_lock
|
||||
|
||||
lock_key = getattr(self, "_lock_key", None)
|
||||
if lock_key:
|
||||
release_scoped_lock("android", lock_key)
|
||||
try:
|
||||
await self._ws_server.stop()
|
||||
except Exception:
|
||||
logger.warning("android: WS server stop failed", exc_info=True)
|
||||
try:
|
||||
await self._http_server.stop()
|
||||
except Exception:
|
||||
@@ -1689,31 +1659,24 @@ class AndroidAdapter(BasePlatformAdapter):
|
||||
return self._http_reply_sinks.pop(device_id, None)
|
||||
|
||||
async def _broadcast_both(self, frame: "protocol.Frame") -> None:
|
||||
"""Bare (non-outbox) broadcast to both transports (docs/19): the
|
||||
frame reaches WS devices and live SSE/long-poll subscribers."""
|
||||
await self._ws_server.broadcast(frame)
|
||||
"""Bare (non-outbox) broadcast to live subscribers (docs/19): the
|
||||
frame reaches every live SSE/long-poll subscriber."""
|
||||
await self._http_server.fanout(frame, cursor=None)
|
||||
|
||||
async def _reply(self, device_id: str, frame: "protocol.Frame") -> None:
|
||||
"""Point-to-point reply with HTTP-leg fallback (docs/19 §19.7).
|
||||
"""Point-to-point reply with broadcast fallback (docs/19 §19.7).
|
||||
|
||||
WS-originated requests keep point-to-point delivery. For an
|
||||
in-flight HTTP request (a reply sink is registered) the frame goes
|
||||
into the HTTP response. If the device has no live WS and no sink
|
||||
(e.g. it dropped mid-request), the frame is broadcast so the SSE
|
||||
stream delivers it (single-user model).
|
||||
For an in-flight HTTP request (a reply sink is registered) the frame
|
||||
goes into the HTTP response. Otherwise it is broadcast so the
|
||||
device's SSE stream delivers it (single-user model).
|
||||
"""
|
||||
entry = self._http_reply_sinks.get(device_id)
|
||||
if entry is not None:
|
||||
entry[0].put(frame)
|
||||
return
|
||||
if await self._ws_server.send_to(device_id, frame):
|
||||
return
|
||||
await self._ws_server.broadcast(frame)
|
||||
await self._http_server.fanout(frame, cursor=None)
|
||||
await self._broadcast_both(frame)
|
||||
|
||||
async def _broadcast_or_log(self, chat_id: str, frame: "protocol.Frame") -> None:
|
||||
delivered = await self._ws_server.broadcast(frame)
|
||||
# M3/M5: always append to the outbox so a reconnecting app can catch
|
||||
# up on *all* recent frames, not just the ones that were parked. This
|
||||
# covers the case where the app's in-memory ChatStore is reset (e.g.
|
||||
@@ -1729,7 +1692,7 @@ class AndroidAdapter(BasePlatformAdapter):
|
||||
# docs/19 §19.8: a device reading SSE/long-poll IS a live subscriber
|
||||
# — count it in the delivery total or every message would push AND
|
||||
# stream to a device that is already receiving it.
|
||||
delivered += await self._http_server.fanout(frame, cursor)
|
||||
delivered = await self._http_server.fanout(frame, cursor)
|
||||
if delivered == 0:
|
||||
logger.info(
|
||||
"android: no live devices for %s; %s frame parked in outbox (cursor=%s)",
|
||||
@@ -1827,12 +1790,7 @@ class AndroidAdapter(BasePlatformAdapter):
|
||||
device_id = device.get("device_id")
|
||||
if not device_id:
|
||||
continue
|
||||
# Prefer the live connection's token (fcm.register refreshes it
|
||||
# in memory) over the possibly-stale registry row.
|
||||
conn = self._ws_server.connection(device_id)
|
||||
token = getattr(conn, backend.token_field, None) if conn is not None else None
|
||||
if not token:
|
||||
token = device.get(backend.token_field)
|
||||
token = device.get(backend.token_field)
|
||||
if not token:
|
||||
continue
|
||||
try:
|
||||
@@ -1894,13 +1852,11 @@ class AndroidAdapter(BasePlatformAdapter):
|
||||
if isinstance(tid, str) and tid:
|
||||
thread_id = tid
|
||||
frame = protocol.typing(chat_id, True, thread_id=thread_id)
|
||||
await self._ws_server.broadcast(frame)
|
||||
await self._http_server.fanout(frame, cursor=None)
|
||||
|
||||
async def stop_typing(self, chat_id: str) -> None:
|
||||
"""Clear the typing indicator (``typing`` frame, on=false)."""
|
||||
frame = protocol.typing(chat_id, False)
|
||||
await self._ws_server.broadcast(frame)
|
||||
await self._http_server.fanout(frame, cursor=None)
|
||||
|
||||
# ── M4: outbound media (agent -> app) ─────────────────────────────────
|
||||
@@ -2243,150 +2199,6 @@ class AndroidAdapter(BasePlatformAdapter):
|
||||
|
||||
threading.Thread(target=_work, daemon=True, name="android-thread-title").start()
|
||||
|
||||
# ── M4: inbound media (app -> agent) ──────────────────────────────────
|
||||
#
|
||||
# ``media.upload.start`` -> raw binary frames (one at a time per
|
||||
# connection) -> ``media.upload.end``. The session streams to a temp
|
||||
# file (bounded RAM); on end we verify size + sha256, re-sniff the kind,
|
||||
# and cache via hermes ``cache_*_from_bytes``. ``media.pull`` serves an
|
||||
# outbound offer as chunked binary frames, re-checking the delivery-path
|
||||
# validation at pull time.
|
||||
|
||||
async def on_media_upload_start(self, frame: protocol.Frame, device_id: str) -> None:
|
||||
payload = frame.payload
|
||||
media_ref = str(payload.get("media_ref") or "").strip()
|
||||
if not media_ref or len(media_ref) > MAX_MEDIA_REF_LEN:
|
||||
await self._reply(
|
||||
device_id,
|
||||
protocol.error(
|
||||
protocol.ERR_UNSUPPORTED, "media.upload.start requires media_ref", id=frame.id
|
||||
),
|
||||
)
|
||||
return
|
||||
kind = payload.get("kind")
|
||||
if kind not in media_bridge.KINDS:
|
||||
await self._reply(
|
||||
device_id,
|
||||
protocol.error(
|
||||
protocol.ERR_UNSUPPORTED, f"unsupported media kind {kind!r}", id=frame.id
|
||||
),
|
||||
)
|
||||
return
|
||||
mime = str(payload.get("mime") or "application/octet-stream")[:128]
|
||||
filename = str(payload.get("filename") or "upload")[:255]
|
||||
size = payload.get("size")
|
||||
try:
|
||||
size = int(size) if size is not None else -1
|
||||
except (TypeError, ValueError):
|
||||
size = -1
|
||||
if size <= 0:
|
||||
await self._reply(
|
||||
device_id,
|
||||
protocol.error(
|
||||
protocol.ERR_UNSUPPORTED,
|
||||
"media.upload.start requires a positive size",
|
||||
id=frame.id,
|
||||
),
|
||||
)
|
||||
return
|
||||
if size > self.max_upload_bytes:
|
||||
await self._reply(
|
||||
device_id,
|
||||
protocol.error(
|
||||
protocol.ERR_MEDIA_TOO_LARGE,
|
||||
f"upload of {size} bytes exceeds limit ({self.max_upload_bytes})",
|
||||
id=frame.id,
|
||||
),
|
||||
)
|
||||
return
|
||||
try:
|
||||
self._media.create_upload(
|
||||
device_id,
|
||||
media_ref,
|
||||
kind,
|
||||
mime,
|
||||
filename,
|
||||
size,
|
||||
frame.id,
|
||||
self.max_upload_bytes,
|
||||
)
|
||||
except media_bridge.MediaError as e:
|
||||
await self._reply(device_id, protocol.error(e.code, e.message, id=frame.id))
|
||||
return
|
||||
# No ack: WS ordering guarantees the server processes this before the
|
||||
# first binary chunk; failures arrive as ``error`` frames.
|
||||
|
||||
async def on_media_chunk(self, device_id: str, chunk: bytes) -> None:
|
||||
session = self._media.get_upload(device_id)
|
||||
if session is None:
|
||||
return # stray binary frame: ignore (forward-compat)
|
||||
session.feed(chunk)
|
||||
if session.failed:
|
||||
await self._reply(
|
||||
device_id,
|
||||
protocol.error(session.error_code, session.error_message, id=session.request_id),
|
||||
)
|
||||
self._media.discard_upload(device_id, session.media_ref)
|
||||
|
||||
async def on_media_upload_end(self, frame: protocol.Frame, device_id: str) -> None:
|
||||
payload = frame.payload
|
||||
media_ref = str(payload.get("media_ref") or "").strip()
|
||||
sha256 = str(payload.get("sha256") or "").strip().lower()
|
||||
if not media_ref:
|
||||
await self._reply(
|
||||
device_id,
|
||||
protocol.error(
|
||||
protocol.ERR_UNSUPPORTED, "media.upload.end requires media_ref", id=frame.id
|
||||
),
|
||||
)
|
||||
return
|
||||
try:
|
||||
entry = self._media.complete_upload(device_id, media_ref, sha256)
|
||||
except media_bridge.MediaError as e:
|
||||
await self._reply(device_id, protocol.error(e.code, e.message, id=frame.id))
|
||||
return
|
||||
await self._reply(
|
||||
device_id, protocol.media_upload_ack(True, entry.media_id, id=frame.id)
|
||||
)
|
||||
|
||||
async def on_media_pull(self, frame: protocol.Frame, device_id: str) -> None:
|
||||
payload = frame.payload
|
||||
media_id = str(payload.get("media_id") or "").strip()
|
||||
entry = self._media.get_outbound(media_id) if media_id else None
|
||||
if entry is None:
|
||||
await self._reply(
|
||||
device_id,
|
||||
protocol.error(
|
||||
protocol.ERR_NOT_FOUND, f"unknown media_id {media_id!r}", id=frame.id
|
||||
),
|
||||
)
|
||||
return
|
||||
# Delivery-path security: re-validate at pull time (the file may have
|
||||
# moved / been replaced since the offer).
|
||||
safe = validate_media_delivery_path(entry.path)
|
||||
if safe is None:
|
||||
await self._reply(
|
||||
device_id,
|
||||
protocol.error(protocol.ERR_NOT_FOUND, "media no longer deliverable", id=frame.id),
|
||||
)
|
||||
return
|
||||
conn = self._ws_server.connection(device_id)
|
||||
if conn is None:
|
||||
return
|
||||
try:
|
||||
await media_bridge.stream_file(conn.ws, safe, media_bridge.DEFAULT_CHUNK_BYTES)
|
||||
except Exception as e:
|
||||
logger.warning("android: media.pull stream failed for %s: %s", media_id, e)
|
||||
await self._reply(
|
||||
device_id, protocol.error(protocol.ERR_INTERNAL, f"pull failed: {e}", id=frame.id)
|
||||
)
|
||||
return
|
||||
await self._reply(device_id, protocol.media_pull_end(True, id=frame.id))
|
||||
|
||||
def on_connection_closed(self, device_id: str) -> None:
|
||||
"""M4: drop in-flight upload temp files for a disconnected device."""
|
||||
self._media.discard_device(device_id)
|
||||
|
||||
# ── M3: channel directory management (app -> agent) ───────────────────
|
||||
#
|
||||
# Each request is answered by broadcasting the matching ``channel.*``
|
||||
@@ -2421,7 +2233,7 @@ class AndroidAdapter(BasePlatformAdapter):
|
||||
return
|
||||
resp = protocol.channel_created(entry)
|
||||
resp.id = frame.id
|
||||
await self._ws_server.broadcast(resp)
|
||||
await self._http_server.fanout(resp)
|
||||
# M5: banner + push mirror (parked in the outbox when offline).
|
||||
await self._broadcast_or_log(
|
||||
entry["chat_id"],
|
||||
@@ -2467,7 +2279,7 @@ class AndroidAdapter(BasePlatformAdapter):
|
||||
return
|
||||
resp = protocol.channel_renamed(entry)
|
||||
resp.id = frame.id
|
||||
await self._ws_server.broadcast(resp)
|
||||
await self._http_server.fanout(resp)
|
||||
# M5: banner + push mirror (parked in the outbox when offline).
|
||||
await self._broadcast_or_log(
|
||||
chat_id,
|
||||
@@ -2500,7 +2312,7 @@ class AndroidAdapter(BasePlatformAdapter):
|
||||
# new is_default flag) so every device reconciles the default change.
|
||||
resp = protocol.channel_renamed(entry)
|
||||
resp.id = frame.id
|
||||
await self._ws_server.broadcast(resp)
|
||||
await self._http_server.fanout(resp)
|
||||
|
||||
async def on_channel_favorite(self, frame: protocol.Frame, device_id: str) -> None:
|
||||
chat_id = frame.chat_id or frame.payload.get("chat_id")
|
||||
@@ -2524,7 +2336,7 @@ class AndroidAdapter(BasePlatformAdapter):
|
||||
# new favorite flag) so every device reconciles the change.
|
||||
resp = protocol.channel_renamed(entry)
|
||||
resp.id = frame.id
|
||||
await self._ws_server.broadcast(resp)
|
||||
await self._http_server.fanout(resp)
|
||||
|
||||
async def on_channel_icon(self, frame: protocol.Frame, device_id: str) -> None:
|
||||
chat_id = frame.chat_id or frame.payload.get("chat_id")
|
||||
@@ -2557,7 +2369,7 @@ class AndroidAdapter(BasePlatformAdapter):
|
||||
return
|
||||
resp = protocol.channel_renamed(entry)
|
||||
resp.id = frame.id
|
||||
await self._ws_server.broadcast(resp)
|
||||
await self._http_server.fanout(resp)
|
||||
|
||||
async def on_channel_set_automation(self, frame: protocol.Frame, device_id: str) -> None:
|
||||
chat_id = frame.chat_id or frame.payload.get("chat_id")
|
||||
@@ -2585,7 +2397,7 @@ class AndroidAdapter(BasePlatformAdapter):
|
||||
# new automation flag) so every device reconciles the change.
|
||||
resp = protocol.channel_renamed(entry)
|
||||
resp.id = frame.id
|
||||
await self._ws_server.broadcast(resp)
|
||||
await self._http_server.fanout(resp)
|
||||
|
||||
async def on_channel_delete(self, frame: protocol.Frame, device_id: str) -> None:
|
||||
chat_id = frame.chat_id or frame.payload.get("chat_id")
|
||||
@@ -2632,7 +2444,7 @@ class AndroidAdapter(BasePlatformAdapter):
|
||||
)
|
||||
resp = protocol.channel_deleted(chat_id)
|
||||
resp.id = frame.id
|
||||
await self._ws_server.broadcast(resp)
|
||||
await self._http_server.fanout(resp)
|
||||
# M5: banner + push mirror (parked in the outbox when offline).
|
||||
await self._broadcast_or_log(
|
||||
chat_id,
|
||||
@@ -2844,8 +2656,8 @@ class AndroidAdapter(BasePlatformAdapter):
|
||||
async def on_fcm_register(self, frame: protocol.Frame, device_id: str) -> None:
|
||||
"""Update the device's push tokens (FCM rotation / ntfy topic).
|
||||
|
||||
Persists to the device registry AND refreshes the live connection so
|
||||
the next push targets the current token without a stale read.
|
||||
Persists to the device registry so the next push targets the current
|
||||
token without a stale read.
|
||||
"""
|
||||
fcm_token = frame.payload.get("fcm_token")
|
||||
ntfy_topic = frame.payload.get("ntfy_topic")
|
||||
@@ -2858,12 +2670,6 @@ class AndroidAdapter(BasePlatformAdapter):
|
||||
except Exception:
|
||||
logger.warning("android: fcm.register update failed", exc_info=True)
|
||||
return
|
||||
conn = self._ws_server.connection(device_id)
|
||||
if conn is not None:
|
||||
if fcm_token is not None:
|
||||
conn.fcm_token = fcm_token
|
||||
if ntfy_topic is not None:
|
||||
conn.ntfy_topic = ntfy_topic
|
||||
logger.info("android: push tokens updated for %s", device_id)
|
||||
|
||||
# ── M5: approval / clarify banners ────────────────────────────────────
|
||||
@@ -3107,7 +2913,7 @@ def register(ctx):
|
||||
validate_config=validate_config,
|
||||
is_connected=is_connected,
|
||||
required_env=["ANDROID_TOKEN"],
|
||||
install_hint="No extra packages needed (websockets + httpx are core deps)",
|
||||
install_hint="No extra packages needed (httpx is a core dep)",
|
||||
setup_fn=interactive_setup,
|
||||
# Env-driven auto-configuration: seeds PlatformConfig.extra with
|
||||
# host/port/push_backend + home_channel so env-only setups show up in
|
||||
|
||||
@@ -0,0 +1,84 @@
|
||||
"""Shared inbound frame dispatch + inbound rate limit.
|
||||
|
||||
Extracted from the (now-removed) WS server so the HTTP transport has a
|
||||
single home for the transport-agnostic dispatch chain and the per-device
|
||||
token bucket. The HTTP leg (``http_server.py``) is the only transport;
|
||||
this module is transport-neutral.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
from . import protocol
|
||||
|
||||
# Inbound JSON control-frame rate limit (per device, token bucket).
|
||||
# A legitimate app sends occasional user-initiated requests — far below
|
||||
# 20/s sustained. Media uploads are exempt (they travel via
|
||||
# ``POST /v1/media``, not the frame endpoint).
|
||||
INBOUND_RATE_PER_S = 20.0
|
||||
INBOUND_BURST = 40
|
||||
|
||||
# Max length of a client-supplied device_id.
|
||||
MAX_DEVICE_ID_LEN = 128
|
||||
|
||||
|
||||
async def dispatch_frame(adapter: Any, frame: protocol.Frame, device_id: str) -> None:
|
||||
"""Shared inbound frame dispatch (docs/19 §19.4). Unknown types are
|
||||
ignored (forward-compat)."""
|
||||
if frame.type == protocol.TYPE_MESSAGE_SEND:
|
||||
await adapter.on_message_send(frame, device_id)
|
||||
elif frame.type == protocol.TYPE_CHANNEL_CREATE:
|
||||
await adapter.on_channel_create(frame, device_id)
|
||||
elif frame.type == protocol.TYPE_CHANNEL_RENAME:
|
||||
await adapter.on_channel_rename(frame, device_id)
|
||||
elif frame.type == protocol.TYPE_CHANNEL_SET_DEFAULT:
|
||||
await adapter.on_channel_set_default(frame, device_id)
|
||||
elif frame.type == protocol.TYPE_CHANNEL_FAVORITE:
|
||||
await adapter.on_channel_favorite(frame, device_id)
|
||||
elif frame.type == protocol.TYPE_CHANNEL_ICON:
|
||||
await adapter.on_channel_icon(frame, device_id)
|
||||
elif frame.type == protocol.TYPE_CHANNEL_SET_AUTOMATION:
|
||||
await adapter.on_channel_set_automation(frame, device_id)
|
||||
elif frame.type == protocol.TYPE_CHANNEL_DELETE:
|
||||
await adapter.on_channel_delete(frame, device_id)
|
||||
elif frame.type == protocol.TYPE_CHANNEL_LIST:
|
||||
await adapter.on_channel_list(frame, device_id)
|
||||
elif frame.type == protocol.TYPE_COMMANDS_CATALOG:
|
||||
await adapter.on_commands_catalog(frame, device_id)
|
||||
elif frame.type == protocol.TYPE_SEARCH:
|
||||
await adapter.on_search(frame, device_id)
|
||||
elif frame.type == protocol.TYPE_SYNC:
|
||||
await adapter.on_sync(frame, device_id)
|
||||
elif frame.type == protocol.TYPE_HISTORY:
|
||||
await adapter.on_history(frame, device_id)
|
||||
elif frame.type == protocol.TYPE_MESSAGE_DELETE:
|
||||
await adapter.on_message_delete(frame, device_id)
|
||||
elif frame.type == protocol.TYPE_FCM_REGISTER:
|
||||
await adapter.on_fcm_register(frame, device_id)
|
||||
# Unknown types are ignored (forward-compat).
|
||||
|
||||
|
||||
class _TokenBucket:
|
||||
"""Minimal token bucket (stdlib only). One instance per device."""
|
||||
|
||||
__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
|
||||
+274
-97
@@ -1,11 +1,10 @@
|
||||
"""HTTP fallback transport (docs/19): the "HTTP leg".
|
||||
"""HTTP transport (docs/19): the gateway's device-facing server.
|
||||
|
||||
A second, short-lived-connection transport next to the WebSocket: the same
|
||||
JSON frames, the same outbox/cursor, the same token — served over plain HTTP
|
||||
by the gateway. When the WS is down (flaky network, NAT timeout, app just
|
||||
relaunched), the app sends over ``POST /v1/frame`` and receives over
|
||||
``GET /v1/events`` (SSE) or ``GET /v1/poll`` (long-poll) instead of waiting
|
||||
for a WS redial.
|
||||
Short-lived-connection transport: the same JSON frames, the same
|
||||
outbox/cursor, the same token — served over plain HTTP by the gateway.
|
||||
The app sends over ``POST /v1/frame`` and receives over
|
||||
``GET /v1/events`` (SSE) or ``GET /v1/poll`` (long-poll); media travels
|
||||
via ``POST /v1/media`` / ``GET /v1/media/{id}`` (v2, docs/19 §19.15).
|
||||
|
||||
Zero new Python dependencies: stdlib ``http.server`` (a
|
||||
``ThreadingHTTPServer`` in a daemon thread) bridged into the gateway's
|
||||
@@ -13,19 +12,30 @@ asyncio loop with ``asyncio.run_coroutine_threadsafe``.
|
||||
|
||||
Endpoints (docs/19 §19.4):
|
||||
* ``GET /v1/health`` — unauthenticated liveness probe.
|
||||
* ``POST /v1/frame`` — accept-and-ack for any JSON frame the WS
|
||||
accepts (except binary media, which stays
|
||||
WS-only in v1).
|
||||
* ``POST /v1/frame`` — accept-and-ack for any JSON frame the
|
||||
app sends (media uses the /v1/media
|
||||
endpoints; hello/ping are
|
||||
transport-specific).
|
||||
* ``GET /v1/events?cursor=N`` — SSE stream: outbox catch-up, then live
|
||||
frames (``id`` = outbox cursor, so resume
|
||||
is just ``Last-Event-ID``).
|
||||
* ``GET /v1/poll?cursor=N`` — long-poll fallback where SSE is blocked.
|
||||
* ``POST /v1/media`` — media upload (docs/19 §19.15, v2): the
|
||||
whole file as the request body; metadata
|
||||
in ``X-Iris-Media-*`` headers; sha256
|
||||
contract per docs/07 §7.2.
|
||||
* ``GET /v1/media/{media_id}`` — media pull (docs/19 §19.15, v2): streams
|
||||
an outbound offer (``media.offer`` id)
|
||||
as the response body.
|
||||
|
||||
Auth: ``Authorization: Bearer <token>`` (constant-time ``verify_token``) +
|
||||
``X-Iris-Device`` header (same device id / allowlist as the WS ``hello``).
|
||||
``X-Iris-Device`` header (device id / allowlist). Device registration
|
||||
(name + push tokens) rides on the SSE open via ``X-Iris-Device-Name`` /
|
||||
``X-Iris-Fcm-Token`` / ``X-Iris-Ntfy-Topic`` headers (the HTTP equivalent
|
||||
of the old WS ``hello`` upsert).
|
||||
|
||||
Bind failure is NON-fatal (unlike the WS): the plugin keeps working
|
||||
WS-only.
|
||||
HTTP is the ONLY transport: a bind failure is FATAL (the app has no other
|
||||
way to reach the gateway).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -43,24 +53,25 @@ from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
|
||||
from typing import Any
|
||||
from urllib.parse import parse_qs, urlparse
|
||||
|
||||
from . import protocol
|
||||
from . import dispatch, protocol
|
||||
from . import media as media_bridge
|
||||
from .pairing import verify_token
|
||||
from .ws_server import (
|
||||
INBOUND_BURST,
|
||||
INBOUND_RATE_PER_S,
|
||||
MAX_DEVICE_ID_LEN,
|
||||
_TokenBucket,
|
||||
dispatch_frame,
|
||||
)
|
||||
|
||||
try: # main-repo import (same as adapter.py); absent in bare unit contexts
|
||||
from gateway.platforms.base import validate_media_delivery_path
|
||||
except ImportError: # pragma: no cover
|
||||
validate_media_delivery_path = None # type: ignore[assignment]
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Default port for the HTTP leg (WS default is 8790).
|
||||
# Default port for the HTTP transport (the WS-era default was 8790).
|
||||
DEFAULT_HTTP_PORT = 8791
|
||||
|
||||
# Request body cap for POST /v1/frame (frames are small; media never
|
||||
# travels here in v1).
|
||||
MAX_BODY_BYTES = 64 * 1024
|
||||
# Request body cap for POST /v1/frame. Frames are usually small, but a
|
||||
# ``channel.icon`` carries a base64 blob up to 512 KiB (docs/10), so the cap
|
||||
# must clear that with headroom. Media never travels here (it uses
|
||||
# POST /v1/media).
|
||||
MAX_BODY_BYTES = 1024 * 1024
|
||||
|
||||
# Per-subscriber live-frame queue. A subscriber that can't keep up is
|
||||
# dropped; it reconnects with Last-Event-ID and catches up from the outbox.
|
||||
@@ -79,18 +90,16 @@ ACCEPT_ACK_TIMEOUT_S = 5.0
|
||||
# Sentinel pushed into subscriber queues on shutdown.
|
||||
_STOP = object()
|
||||
|
||||
# Frame types POST /v1/frame must not accept (docs/19 §19.3): media is
|
||||
# inherently binary/streaming (WS-only in v1); hello/ping are
|
||||
# transport-specific (auth is via headers, liveness via /v1/health).
|
||||
HTTP_REJECTED_TYPES = frozenset(
|
||||
{
|
||||
protocol.TYPE_HELLO,
|
||||
protocol.TYPE_PING,
|
||||
protocol.TYPE_MEDIA_UPLOAD_START,
|
||||
protocol.TYPE_MEDIA_UPLOAD_END,
|
||||
protocol.TYPE_MEDIA_PULL,
|
||||
}
|
||||
)
|
||||
# Max length of a client-supplied media_ref (same as the WS path).
|
||||
MAX_MEDIA_REF_LEN = 64
|
||||
|
||||
# MediaError code -> HTTP status for the /v1/media endpoints.
|
||||
_MEDIA_STATUS = {
|
||||
protocol.ERR_MEDIA_TOO_LARGE: 413,
|
||||
protocol.ERR_NOT_FOUND: 404,
|
||||
protocol.ERR_UNSUPPORTED: 400,
|
||||
protocol.ERR_INTERNAL: 500,
|
||||
}
|
||||
|
||||
|
||||
def _with_cursor(frame: dict[str, Any], cursor: int) -> str:
|
||||
@@ -151,12 +160,12 @@ class _Subscriber:
|
||||
|
||||
|
||||
class HttpServer:
|
||||
"""The plugin's HTTP fallback server + live subscriber registry.
|
||||
"""The plugin's HTTP server + live subscriber registry.
|
||||
|
||||
The handler threads never touch adapter state directly: inbound frames
|
||||
are bridged into the gateway's asyncio loop (captured at ``start()``)
|
||||
with ``asyncio.run_coroutine_threadsafe`` and dispatched through the
|
||||
same ``dispatch_frame`` the WS server uses.
|
||||
with ``asyncio.run_coroutine_threadsafe`` and dispatched through
|
||||
``dispatch_frame`` (``dispatch.py``).
|
||||
"""
|
||||
|
||||
def __init__(self, adapter: Any, devices: Any):
|
||||
@@ -167,7 +176,7 @@ class HttpServer:
|
||||
self._thread: threading.Thread | None = None
|
||||
self._subs: dict[str, list[_Subscriber]] = {}
|
||||
self._subs_lock = threading.Lock()
|
||||
self._buckets: dict[str, _TokenBucket] = {}
|
||||
self._buckets: dict[str, dispatch._TokenBucket] = {}
|
||||
self._buckets_lock = threading.Lock()
|
||||
self._lock_key: str | None = None
|
||||
self.enabled = False
|
||||
@@ -176,8 +185,9 @@ class HttpServer:
|
||||
# ── Lifecycle ─────────────────────────────────────────────────────────
|
||||
|
||||
async def start(self) -> None:
|
||||
"""Bind and start serving. NEVER raises: a bind failure disables the
|
||||
HTTP leg (the plugin keeps working WS-only, docs/19 §19.4)."""
|
||||
"""Bind and start serving. NEVER raises: a bind failure leaves
|
||||
``enabled`` False, which the adapter treats as a fatal error
|
||||
(HTTP is the only transport, docs/19 §19.4)."""
|
||||
if self.enabled:
|
||||
return
|
||||
self._loop = asyncio.get_running_loop()
|
||||
@@ -191,7 +201,7 @@ class HttpServer:
|
||||
lock_key = f"http:{host}:{port}"
|
||||
if not acquire_scoped_lock("android", lock_key):
|
||||
logger.warning(
|
||||
"android: HTTP port %s:%s in use by another profile; HTTP leg disabled",
|
||||
"android: HTTP port %s:%s in use by another profile; server disabled",
|
||||
host,
|
||||
port,
|
||||
)
|
||||
@@ -208,9 +218,7 @@ class HttpServer:
|
||||
ctx.load_cert_chain(self._adapter.http_cert, self._adapter.http_key)
|
||||
httpd.socket = ctx.wrap_socket(httpd.socket, server_side=True)
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
"android: HTTP fallback leg disabled (bind %s:%s failed: %s)", host, port, e
|
||||
)
|
||||
logger.warning("android: HTTP server disabled (bind %s:%s failed: %s)", host, port, e)
|
||||
self._release_lock()
|
||||
return
|
||||
|
||||
@@ -221,9 +229,7 @@ class HttpServer:
|
||||
self._thread.start()
|
||||
self.enabled = True
|
||||
scheme = "https" if (self._adapter.http_cert and self._adapter.http_key) else "http"
|
||||
logger.info(
|
||||
"android: HTTP fallback leg listening on %s://%s:%s", scheme, host, self.bound_port
|
||||
)
|
||||
logger.info("android: HTTP server listening on %s://%s:%s", scheme, host, self.bound_port)
|
||||
|
||||
async def stop(self) -> None:
|
||||
"""Stop serving and unblock all subscribers."""
|
||||
@@ -317,7 +323,7 @@ class HttpServer:
|
||||
_send_json(handler, 401, {"error": "unauthorized"})
|
||||
return None
|
||||
device_id = (handler.headers.get("X-Iris-Device") or "").strip()
|
||||
if not device_id or len(device_id) > MAX_DEVICE_ID_LEN:
|
||||
if not device_id or len(device_id) > dispatch.MAX_DEVICE_ID_LEN:
|
||||
_send_json(handler, 401, {"error": "X-Iris-Device header required"})
|
||||
return None
|
||||
if (
|
||||
@@ -333,11 +339,13 @@ class HttpServer:
|
||||
return device_id
|
||||
|
||||
def _rate_limited(self, device_id: str) -> bool:
|
||||
"""Per-device token bucket, same parameters as the WS inbound limit."""
|
||||
"""Per-device token bucket, same parameters as the frame limit."""
|
||||
with self._buckets_lock:
|
||||
b = self._buckets.get(device_id)
|
||||
if b is None:
|
||||
b = self._buckets[device_id] = _TokenBucket(INBOUND_RATE_PER_S, INBOUND_BURST)
|
||||
b = self._buckets[device_id] = dispatch._TokenBucket(
|
||||
dispatch.INBOUND_RATE_PER_S, dispatch.INBOUND_BURST
|
||||
)
|
||||
return not b.consume()
|
||||
|
||||
# ── POST /v1/frame ────────────────────────────────────────────────────
|
||||
@@ -379,31 +387,35 @@ class HttpServer:
|
||||
frames: list[protocol.Frame] = []
|
||||
deadline = time.monotonic() + ACCEPT_ACK_TIMEOUT_S
|
||||
while True:
|
||||
try:
|
||||
frames.append(sink.get(timeout=0.05))
|
||||
# If the handler is done, drain any replies and stop (no wait).
|
||||
# This keeps fast/ignored frames from incurring the sink timeout.
|
||||
if task.done():
|
||||
while True:
|
||||
try:
|
||||
frames.append(sink.get_nowait())
|
||||
except queue.Empty:
|
||||
break
|
||||
break
|
||||
try:
|
||||
frames.append(sink.get(timeout=0.01))
|
||||
except queue.Empty:
|
||||
if task.done():
|
||||
# All replies are in the sink now (the handler finished);
|
||||
# drain them all.
|
||||
while True:
|
||||
try:
|
||||
frames.append(sink.get_nowait())
|
||||
except queue.Empty:
|
||||
break
|
||||
break
|
||||
if time.monotonic() >= deadline:
|
||||
# Long-running handler (the agent turn): ack now; late
|
||||
# replies go to the event stream (the dispatch's finally
|
||||
# sees ``abandoned`` and delivers them there).
|
||||
abandoned.set()
|
||||
break
|
||||
continue
|
||||
# Got a frame; loop back to check task.done() (drain the rest if
|
||||
# the handler finished, e.g. a sync replay).
|
||||
if not frames:
|
||||
_send_json(handler, 202, {"ok": True})
|
||||
elif len(frames) == 1:
|
||||
f = frames[0]
|
||||
status = 429 if f.payload.get("code") == protocol.ERR_RATE_LIMITED else (
|
||||
400 if f.type == protocol.TYPE_ERROR else 200
|
||||
status = (
|
||||
429
|
||||
if f.payload.get("code") == protocol.ERR_RATE_LIMITED
|
||||
else (400 if f.type == protocol.TYPE_ERROR else 200)
|
||||
)
|
||||
_send_frame_json(handler, status, f.to_json())
|
||||
else:
|
||||
@@ -411,9 +423,7 @@ class HttpServer:
|
||||
# event stream; the ack stays plain.
|
||||
for f in frames:
|
||||
with contextlib.suppress(Exception):
|
||||
asyncio.run_coroutine_threadsafe(
|
||||
self._deliver_via_stream(f), loop
|
||||
)
|
||||
asyncio.run_coroutine_threadsafe(self._deliver_via_stream(f), loop)
|
||||
_send_json(handler, 202, {"ok": True})
|
||||
|
||||
async def _dispatch_guarded(
|
||||
@@ -424,11 +434,9 @@ class HttpServer:
|
||||
abandoned: threading.Event,
|
||||
) -> None:
|
||||
try:
|
||||
await dispatch_frame(self._adapter, frame, device_id)
|
||||
await dispatch.dispatch_frame(self._adapter, frame, device_id)
|
||||
except Exception:
|
||||
logger.warning(
|
||||
"android: HTTP dispatch failed for %s", frame.type, exc_info=True
|
||||
)
|
||||
logger.warning("android: HTTP dispatch failed for %s", frame.type, exc_info=True)
|
||||
finally:
|
||||
# Pop our sink entry (a newer request from the same device may
|
||||
# have replaced it). If the HTTP response was already sent
|
||||
@@ -438,25 +446,39 @@ class HttpServer:
|
||||
# left to deliver.
|
||||
popped = self._adapter._http_pop_sink_if(device_id, sink)
|
||||
if popped is not None and abandoned.is_set():
|
||||
while True:
|
||||
try:
|
||||
f = sink.get_nowait()
|
||||
except queue.Empty:
|
||||
break
|
||||
with contextlib.suppress(Exception):
|
||||
await self._deliver_via_stream(f)
|
||||
while True:
|
||||
try:
|
||||
f = sink.get_nowait()
|
||||
except queue.Empty:
|
||||
break
|
||||
with contextlib.suppress(Exception):
|
||||
await self._deliver_via_stream(f)
|
||||
|
||||
async def _deliver_via_stream(self, frame: protocol.Frame) -> None:
|
||||
await self._adapter._ws_server.broadcast(frame)
|
||||
await self.fanout(frame, cursor=None)
|
||||
|
||||
# ── GET /v1/events (SSE) ──────────────────────────────────────────────
|
||||
|
||||
def _handle_sse(
|
||||
self, handler: BaseHTTPRequestHandler, device_id: str, parsed: Any
|
||||
) -> None:
|
||||
def _handle_sse(self, handler: BaseHTTPRequestHandler, device_id: str, parsed: Any) -> None:
|
||||
qs = parse_qs(parsed.query)
|
||||
cursor = _parse_cursor(qs.get("cursor", [None])[0], handler.headers.get("Last-Event-ID"))
|
||||
# Device registration (the HTTP equivalent of the WS hello upsert):
|
||||
# the SSE open carries the device name + push tokens as optional
|
||||
# headers; upsert is idempotent and COALESCEs absent tokens, so a
|
||||
# re-open never clobbers a newer fcm.register value.
|
||||
device_name = (handler.headers.get("X-Iris-Device-Name") or "").strip()[:120]
|
||||
fcm_token = handler.headers.get("X-Iris-Fcm-Token") or None
|
||||
ntfy_topic = handler.headers.get("X-Iris-Ntfy-Topic") or None
|
||||
try:
|
||||
self._devices.upsert(
|
||||
device_id,
|
||||
device_name or device_id,
|
||||
None,
|
||||
fcm_token,
|
||||
ntfy_topic,
|
||||
)
|
||||
except Exception:
|
||||
logger.warning("android: device registry upsert failed", exc_info=True)
|
||||
sub = _Subscriber(device_id=device_id, kind="sse")
|
||||
# Register BEFORE the replay so a frame appended in between is
|
||||
# fanned out to us (and de-duped by cursor below) instead of lost.
|
||||
@@ -499,7 +521,9 @@ class HttpServer:
|
||||
continue # already replayed above
|
||||
self._write_sse(handler, "frame", c, data)
|
||||
except (BrokenPipeError, ConnectionResetError, OSError):
|
||||
pass # client went away: normal
|
||||
# Client went away mid-stream: normal (the app reconnects with
|
||||
# Last-Event-ID and catches up from the outbox).
|
||||
pass
|
||||
finally:
|
||||
self._remove_sub(sub)
|
||||
|
||||
@@ -519,11 +543,137 @@ class HttpServer:
|
||||
handler.wfile.write(text.encode("utf-8"))
|
||||
handler.wfile.flush()
|
||||
|
||||
# ── POST /v1/media (upload, docs/19 §19.15) ───────────────────────────
|
||||
|
||||
def _handle_media_upload(self, handler: BaseHTTPRequestHandler, device_id: str) -> None:
|
||||
"""Whole-file upload: metadata in headers, file bytes as the body.
|
||||
|
||||
Mirrors the WS ``media.upload`` contract (docs/07 §7.2) in one
|
||||
request: the body is streamed to a temp file (bounded RAM), then
|
||||
size + sha256 are verified and the file cached via the hermes
|
||||
``cache_*_from_bytes`` helpers. Runs entirely on the handler thread
|
||||
(plain file IO — no asyncio bridge needed)."""
|
||||
media_ref = (handler.headers.get("X-Iris-Media-Ref") or "").strip()
|
||||
kind = (handler.headers.get("X-Iris-Media-Kind") or "").strip()
|
||||
filename = (handler.headers.get("X-Iris-Media-Filename") or "upload")[:255]
|
||||
sha256 = (handler.headers.get("X-Iris-Media-Sha256") or "").strip().lower()
|
||||
mime = handler.headers.get("Content-Type") or "application/octet-stream"
|
||||
mime = mime.split(";")[0].strip()[:128]
|
||||
try:
|
||||
length = int(handler.headers.get("Content-Length") or 0)
|
||||
except ValueError:
|
||||
length = 0
|
||||
|
||||
def reject(code: str, message: str) -> None:
|
||||
_send_frame_json(
|
||||
handler,
|
||||
_MEDIA_STATUS.get(code, 400),
|
||||
protocol.error(code, message).to_json(),
|
||||
)
|
||||
|
||||
# Same validation rules as the WS media.upload.start handler.
|
||||
if not media_ref or len(media_ref) > MAX_MEDIA_REF_LEN:
|
||||
reject(protocol.ERR_UNSUPPORTED, "X-Iris-Media-Ref header required")
|
||||
return
|
||||
if kind not in media_bridge.KINDS:
|
||||
reject(protocol.ERR_UNSUPPORTED, f"unsupported media kind {kind!r}")
|
||||
return
|
||||
if length <= 0:
|
||||
reject(protocol.ERR_UNSUPPORTED, "empty body")
|
||||
return
|
||||
if length > self._adapter.max_upload_bytes:
|
||||
reject(
|
||||
protocol.ERR_MEDIA_TOO_LARGE,
|
||||
f"upload of {length} bytes exceeds limit ({self._adapter.max_upload_bytes})",
|
||||
)
|
||||
return
|
||||
try:
|
||||
sess = self._adapter._media.create_upload(
|
||||
device_id,
|
||||
media_ref,
|
||||
kind,
|
||||
mime,
|
||||
filename,
|
||||
length,
|
||||
None,
|
||||
self._adapter.max_upload_bytes,
|
||||
)
|
||||
except media_bridge.MediaError as e:
|
||||
reject(e.code, e.message)
|
||||
return
|
||||
try:
|
||||
remaining = length
|
||||
while remaining > 0:
|
||||
chunk = handler.rfile.read(min(media_bridge.DEFAULT_CHUNK_BYTES, remaining))
|
||||
if not chunk:
|
||||
raise media_bridge.MediaError(
|
||||
protocol.ERR_INTERNAL, "client disconnected mid-upload"
|
||||
)
|
||||
sess.feed(chunk)
|
||||
remaining -= len(chunk)
|
||||
if sess.received != length:
|
||||
raise media_bridge.MediaError(
|
||||
protocol.ERR_INTERNAL,
|
||||
f"size mismatch (declared {length}, received {sess.received})",
|
||||
)
|
||||
entry = self._adapter._media.complete_upload(device_id, media_ref, sha256)
|
||||
except media_bridge.MediaError as e:
|
||||
# complete_upload already popped the session; discard is a no-op
|
||||
# in that case (feed/short-read failures leave it active).
|
||||
self._adapter._media.discard_upload(device_id, media_ref)
|
||||
reject(e.code, e.message)
|
||||
return
|
||||
except (BrokenPipeError, ConnectionResetError, OSError):
|
||||
self._adapter._media.discard_upload(device_id, media_ref)
|
||||
return # client went away: nothing to answer
|
||||
_send_frame_json(handler, 201, protocol.media_upload_ack(True, entry.media_id).to_json())
|
||||
|
||||
# ── GET /v1/media/{id} (pull, docs/19 §19.15) ─────────────────────────
|
||||
|
||||
def _handle_media_pull(
|
||||
self, handler: BaseHTTPRequestHandler, device_id: str, media_id: str
|
||||
) -> None:
|
||||
"""Stream an outbound offer as the response body (docs/07 §7.3).
|
||||
|
||||
The delivery-path validation is re-checked at pull time, exactly as
|
||||
the WS ``media.pull`` handler does (the file may have moved since
|
||||
the offer)."""
|
||||
entry = self._adapter._media.get_outbound(media_id)
|
||||
if entry is None:
|
||||
_send_frame_json(
|
||||
handler,
|
||||
404,
|
||||
protocol.error(protocol.ERR_NOT_FOUND, f"unknown media_id {media_id!r}").to_json(),
|
||||
)
|
||||
return
|
||||
safe = validate_media_delivery_path(entry.path) if validate_media_delivery_path else None
|
||||
if safe is None:
|
||||
_send_frame_json(
|
||||
handler,
|
||||
404,
|
||||
protocol.error(protocol.ERR_NOT_FOUND, "media no longer deliverable").to_json(),
|
||||
)
|
||||
return
|
||||
filename = entry.filename.replace('"', "")
|
||||
handler.send_response(200)
|
||||
handler.send_header("Content-Type", entry.mime)
|
||||
handler.send_header("Content-Length", str(entry.size))
|
||||
handler.send_header("Content-Disposition", f'attachment; filename="{filename}"')
|
||||
handler.end_headers()
|
||||
try:
|
||||
with open(safe, "rb") as f: # pi-lens-ignore: python-path-traversal
|
||||
while True:
|
||||
chunk = f.read(media_bridge.DEFAULT_CHUNK_BYTES)
|
||||
if not chunk:
|
||||
break
|
||||
handler.wfile.write(chunk)
|
||||
handler.wfile.flush()
|
||||
except (BrokenPipeError, ConnectionResetError, OSError):
|
||||
pass # client went away mid-pull, or the file vanished: normal
|
||||
|
||||
# ── GET /v1/poll (long-poll) ──────────────────────────────────────────
|
||||
|
||||
def _handle_poll(
|
||||
self, handler: BaseHTTPRequestHandler, device_id: str, parsed: Any
|
||||
) -> None:
|
||||
def _handle_poll(self, handler: BaseHTTPRequestHandler, device_id: str, parsed: Any) -> None:
|
||||
qs = parse_qs(parsed.query)
|
||||
cursor = _parse_cursor(qs.get("cursor", [None])[0])
|
||||
sub = _Subscriber(device_id=device_id, kind="poll")
|
||||
@@ -552,6 +702,7 @@ class HttpServer:
|
||||
hwm = max(max_cursor, self._adapter._outbox.latest_cursor())
|
||||
_send_json(handler, 200, {"cursor": hwm, "frames": frames})
|
||||
except (BrokenPipeError, ConnectionResetError, OSError):
|
||||
# Client went away while we held the poll: normal.
|
||||
pass
|
||||
finally:
|
||||
self._remove_sub(sub)
|
||||
@@ -601,6 +752,17 @@ class _Handler(BaseHTTPRequestHandler):
|
||||
if device_id is not None:
|
||||
hs._handle_poll(self, device_id, parsed)
|
||||
return
|
||||
if parsed.path.startswith("/v1/media/"):
|
||||
media_id = parsed.path[len("/v1/media/") :]
|
||||
# The id is looked up in an exact-match dict; reject anything
|
||||
# path-shaped so a bad URL can't be mistaken for an id.
|
||||
if media_id and "/" not in media_id:
|
||||
device_id = hs._authenticate(self)
|
||||
if device_id is not None:
|
||||
hs._handle_media_pull(self, device_id, media_id)
|
||||
else:
|
||||
_send_json(self, 404, {"error": "not found"})
|
||||
return
|
||||
_send_json(self, 404, {"error": "not found"})
|
||||
|
||||
def do_POST(self) -> None: # noqa: N802
|
||||
@@ -609,6 +771,21 @@ class _Handler(BaseHTTPRequestHandler):
|
||||
_send_json(self, 503, {"error": "http leg disabled"})
|
||||
return
|
||||
parsed = urlparse(self.path)
|
||||
if parsed.path == "/v1/media":
|
||||
device_id = hs._authenticate(self)
|
||||
if device_id is None:
|
||||
return
|
||||
if hs._rate_limited(device_id):
|
||||
_send_frame_json(
|
||||
self,
|
||||
429,
|
||||
protocol.error(
|
||||
protocol.ERR_RATE_LIMITED, "http media rate limit exceeded"
|
||||
).to_json(),
|
||||
)
|
||||
return
|
||||
hs._handle_media_upload(self, device_id)
|
||||
return
|
||||
if parsed.path != "/v1/frame":
|
||||
_send_json(self, 404, {"error": "not found"})
|
||||
return
|
||||
@@ -639,6 +816,15 @@ class _Handler(BaseHTTPRequestHandler):
|
||||
except ValueError:
|
||||
length = 0
|
||||
if length <= 0 or length > MAX_BODY_BYTES:
|
||||
# Drain the (oversize) body so the connection stays clean; cap the
|
||||
# drain at MAX_BODY_BYTES so a runaway body can't wedge the thread.
|
||||
if length > 0:
|
||||
to_drain = min(length, MAX_BODY_BYTES)
|
||||
while to_drain > 0:
|
||||
chunk = self.rfile.read(min(65536, to_drain))
|
||||
if not chunk:
|
||||
break
|
||||
to_drain -= len(chunk)
|
||||
_send_frame_json(
|
||||
self,
|
||||
413,
|
||||
@@ -654,13 +840,4 @@ class _Handler(BaseHTTPRequestHandler):
|
||||
self, 400, protocol.error(protocol.ERR_INTERNAL, "invalid frame").to_json()
|
||||
)
|
||||
return
|
||||
if frame.type in HTTP_REJECTED_TYPES:
|
||||
_send_frame_json(
|
||||
self,
|
||||
400,
|
||||
protocol.error(
|
||||
protocol.ERR_UNSUPPORTED, f"{frame.type} requires the live connection"
|
||||
).to_json(),
|
||||
)
|
||||
return
|
||||
hs._handle_frame(self, device_id, frame)
|
||||
@@ -20,7 +20,6 @@ live under ``get_hermes_home()/"android"/media/tmp``.
|
||||
Milestone M4.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import contextlib
|
||||
import hashlib
|
||||
import logging
|
||||
@@ -434,27 +433,3 @@ class MediaStore:
|
||||
for k in stale:
|
||||
del self._outbound[k]
|
||||
return len(stale)
|
||||
|
||||
|
||||
async def stream_file(
|
||||
ws, path: str, chunk_bytes: int = DEFAULT_CHUNK_BYTES, timeout: float = 10.0
|
||||
) -> int:
|
||||
"""Stream *path* to *ws* as binary frames. Returns bytes sent.
|
||||
|
||||
Ordering is guaranteed by the WebSocket; the caller sends the terminal
|
||||
``media.pull.end`` frame afterwards. Each chunk send is bounded by
|
||||
*timeout* so a stalled puller can't wedge the handler forever (the
|
||||
caller treats the raised error as an aborted pull).
|
||||
"""
|
||||
sent = 0
|
||||
# Safe: ``path`` is produced by hermes ``cache_*_from_bytes`` (a path inside
|
||||
# hermes's own media cache dir), never derived from raw user input.
|
||||
# pi-lens-ignore: python-path-traversal
|
||||
with open(path, "rb") as f:
|
||||
while True:
|
||||
chunk = f.read(chunk_bytes)
|
||||
if not chunk:
|
||||
break
|
||||
await asyncio.wait_for(ws.send(chunk), timeout=timeout)
|
||||
sent += len(chunk)
|
||||
return sent
|
||||
@@ -45,7 +45,7 @@ def verify_token(provided: str | None, expected: str | None) -> bool:
|
||||
def qr_payload(host: str, port: int, token: str, secure: bool = False) -> str:
|
||||
"""Pairing URL encoded into the QR / pre-filled into the app.
|
||||
|
||||
``iris://pair?host=<lan-ip>&port=8790&token=<token>`` — the app's
|
||||
``iris://pair?host=<lan-ip>&port=8791&token=<token>`` — the app's
|
||||
Connect screen parses this to pre-fill settings (docs/09 §9.2).
|
||||
"""
|
||||
return (
|
||||
@@ -57,9 +57,9 @@ def qr_payload(host: str, port: int, token: str, secure: bool = False) -> str:
|
||||
|
||||
|
||||
def pairing_url(host: str, port: int, secure: bool = False) -> str:
|
||||
"""Plain ws(s) URL the app connects to (shown next to the QR)."""
|
||||
scheme = "wss" if secure else "ws"
|
||||
return f"{scheme}://{host}:{int(port)}/ws"
|
||||
"""Plain http(s) URL the app connects to (shown next to the QR)."""
|
||||
scheme = "https" if secure else "http"
|
||||
return f"{scheme}://{host}:{int(port)}"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@@ -26,11 +26,8 @@ PROTOCOL_VERSION = 1
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
# Pairing / lifecycle
|
||||
TYPE_HELLO = "hello"
|
||||
TYPE_HELLO_ACK = "hello.ack"
|
||||
TYPE_ERROR = "error"
|
||||
TYPE_PING = "ping"
|
||||
TYPE_PONG = "pong"
|
||||
|
||||
# Chat
|
||||
TYPE_MESSAGE = "message"
|
||||
@@ -82,12 +79,8 @@ TYPE_SYNC_DONE = "sync.done"
|
||||
TYPE_HISTORY = "history"
|
||||
|
||||
# Media (M4)
|
||||
TYPE_MEDIA_UPLOAD_START = "media.upload.start"
|
||||
TYPE_MEDIA_UPLOAD_END = "media.upload.end"
|
||||
TYPE_MEDIA_UPLOAD_ACK = "media.upload.ack"
|
||||
TYPE_MEDIA_OFFER = "media.offer"
|
||||
TYPE_MEDIA_PULL = "media.pull"
|
||||
TYPE_MEDIA_PULL_END = "media.pull.end"
|
||||
|
||||
# Push / notifications (M5)
|
||||
TYPE_NOTIFICATION = "notification"
|
||||
@@ -749,7 +742,8 @@ def media_offer(
|
||||
thread_id: str | None = None,
|
||||
message_id: str | None = None,
|
||||
) -> Frame:
|
||||
"""Event: the agent produced media the app can fetch via ``media.pull``.
|
||||
"""Event: the agent produced media the app can fetch via
|
||||
``GET /v1/media/{media_id}`` (docs/19 §19.15).
|
||||
|
||||
``message_id`` (optional) associates the offer with the assistant message
|
||||
it belongs to (the app falls back to the lane's last assistant message).
|
||||
@@ -766,14 +760,9 @@ def media_offer(
|
||||
return Frame(type=TYPE_MEDIA_OFFER, chat_id=chat_id, thread_id=thread_id, payload=payload)
|
||||
|
||||
|
||||
def media_pull_end(ok: bool, *, id: int | None = None) -> Frame:
|
||||
"""Terminal frame of a ``media.pull`` binary stream."""
|
||||
return Frame(type=TYPE_MEDIA_PULL_END, id=id, payload={"ok": ok})
|
||||
|
||||
|
||||
def media_upload_ack(ok: bool, media_ref: str, *, id: int | None = None) -> Frame:
|
||||
"""Response to ``media.upload.end``: the ref is cached and may be used in
|
||||
a ``message.send`` ``media_refs``. Failures use ``error`` frames instead."""
|
||||
"""Response to ``POST /v1/media``: the ref is cached and may be used in a
|
||||
``message.send`` ``media_refs``. Failures use ``error`` frames instead."""
|
||||
return Frame(
|
||||
type=TYPE_MEDIA_UPLOAD_ACK,
|
||||
id=id,
|
||||
@@ -783,10 +772,3 @@ def media_upload_ack(ok: bool, media_ref: str, *, id: int | None = None) -> Fram
|
||||
|
||||
def error(code: str, message: str, *, id: int | None = None) -> Frame:
|
||||
return Frame(type=TYPE_ERROR, id=id, payload={"code": code, "message": message})
|
||||
|
||||
|
||||
def pong(ts: int | None = None) -> Frame:
|
||||
payload: dict[str, Any] = {}
|
||||
if ts is not None:
|
||||
payload["ts"] = ts
|
||||
return Frame(type=TYPE_PONG, payload=payload)
|
||||
@@ -5,13 +5,13 @@ The plugin lives in the sibling ``iris_x_hermes`` checkout (installed into
|
||||
from the source tree directly so they never depend on that install.
|
||||
|
||||
Coverage (docs/13-testing.md §13.1, media bullets):
|
||||
* upload start -> binary chunks -> end reassembles + sha256 verified
|
||||
* over-limit (declared and mid-stream) -> ``media_too_large``
|
||||
* upload via ``POST /v1/media`` reassembles + sha256 verified
|
||||
* over-limit (Content-Length) -> ``media_too_large``
|
||||
* sha256 mismatch -> ``internal``
|
||||
* ``message.send`` with ``media_refs`` -> echo carries ``media[]`` and the
|
||||
``MessageEvent`` carries ``media_urls``/``media_types``
|
||||
* ``send_*`` -> ``media.offer`` (fields + message association)
|
||||
* ``media.pull`` serves only allowed paths (denied/unknown -> ``not_found``)
|
||||
* ``GET /v1/media/{id}`` serves only allowed paths (denied/unknown -> ``not_found``)
|
||||
* kind re-sniffing (don't trust the client)
|
||||
|
||||
Run via ``scripts/run_tests.sh tests/gateway/test_android.py``.
|
||||
@@ -26,6 +26,9 @@ import importlib.util
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
import socket
|
||||
import threading
|
||||
from http.client import HTTPConnection
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock
|
||||
@@ -47,8 +50,14 @@ def _plugin_dir() -> Path:
|
||||
env = os.environ.get("ANDROID_PLUGIN_DIR")
|
||||
if env:
|
||||
return Path(env)
|
||||
# hermes-agent/tests/gateway/test_android.py -> repo root is parents[3].
|
||||
return Path(__file__).resolve().parents[3] / "gateway-plugin"
|
||||
# Works from either copy of this file: gateway-plugin/tests/ (canonical,
|
||||
# plugin dir is parents[1]) or the hermes-agent/tests/gateway/ mirror
|
||||
# (repo root is parents[3]).
|
||||
here = Path(__file__).resolve()
|
||||
for candidate in (here.parents[1], here.parents[3] / "gateway-plugin"):
|
||||
if (candidate / "protocol.py").is_file():
|
||||
return candidate
|
||||
return here.parents[1]
|
||||
|
||||
|
||||
def _load_plugin():
|
||||
@@ -105,7 +114,7 @@ def adapter(plugin, monkeypatch):
|
||||
config = SimpleNamespace(
|
||||
extra={
|
||||
"host": "127.0.0.1",
|
||||
"port": 0, # ephemeral port
|
||||
"http_port": 0, # ephemeral HTTP port
|
||||
"max_upload_bytes": 1024 * 1024, # 1 MiB -- keeps over-limit tests fast
|
||||
},
|
||||
home_channel=None,
|
||||
@@ -122,48 +131,206 @@ def adapter(plugin, monkeypatch):
|
||||
pass
|
||||
|
||||
|
||||
async def _hello(ws) -> dict:
|
||||
await ws.send(
|
||||
json.dumps(
|
||||
{
|
||||
"v": 1,
|
||||
"type": "hello",
|
||||
"payload": {
|
||||
"token": TOKEN,
|
||||
"device_id": DEVICE_ID,
|
||||
"device_name": "Test Device",
|
||||
"caps": {},
|
||||
},
|
||||
}
|
||||
class HttpTestClient:
|
||||
"""Mimics the old WS client interface over the HTTP transport (docs/19).
|
||||
|
||||
``.send(json_str)`` -> ``POST /v1/frame``; ``.recv(timeout)`` -> the next
|
||||
frame from the SSE stream (a dict); ``.upload(...)`` -> ``POST /v1/media``
|
||||
returning the terminal frame (ack or error). The SSE stream is read in a
|
||||
background thread that feeds an asyncio queue, so ``recv`` is awaitable
|
||||
exactly like the old WS ``recv``.
|
||||
"""
|
||||
|
||||
def __init__(self, port: int, cursor: int = 0):
|
||||
self._port = port
|
||||
self._cursor = cursor
|
||||
self._queue: asyncio.Queue = asyncio.Queue()
|
||||
self._loop: asyncio.AbstractEventLoop | None = None
|
||||
self._thread: threading.Thread | None = None
|
||||
self._resp = None
|
||||
|
||||
async def start(self) -> dict:
|
||||
"""Open the SSE stream; returns the hello.ack frame. Any outbox
|
||||
catch-up frames replayed before the hello are buffered and re-enqueued
|
||||
so ``recv`` can still see them."""
|
||||
self._loop = asyncio.get_running_loop()
|
||||
self._thread = threading.Thread(target=self._sse_reader, daemon=True)
|
||||
self._thread.start()
|
||||
buffered: list = []
|
||||
hello = None
|
||||
while hello is None:
|
||||
frame = await asyncio.wait_for(self._queue.get(), timeout=5)
|
||||
if frame is None:
|
||||
raise AssertionError("SSE stream closed before hello")
|
||||
if frame.get("type") == "hello.ack":
|
||||
hello = frame
|
||||
else:
|
||||
buffered.append(frame)
|
||||
for f in buffered:
|
||||
self._queue.put_nowait(f)
|
||||
return hello
|
||||
|
||||
def _sse_reader(self) -> None:
|
||||
conn = HTTPConnection("127.0.0.1", self._port, timeout=30)
|
||||
conn.request(
|
||||
"GET",
|
||||
f"/v1/events?cursor={self._cursor}",
|
||||
headers={"Authorization": f"Bearer {TOKEN}", "X-Iris-Device": DEVICE_ID},
|
||||
)
|
||||
)
|
||||
raw = await asyncio.wait_for(ws.recv(), timeout=5)
|
||||
ack = json.loads(raw)
|
||||
assert ack["type"] == "hello.ack", f"expected hello.ack, got {ack}"
|
||||
return ack
|
||||
resp = conn.getresponse()
|
||||
self._resp = resp
|
||||
cur_data: list[str] = []
|
||||
loop = self._loop
|
||||
assert loop is not None
|
||||
while True:
|
||||
line = resp.fp.readline()
|
||||
if not line:
|
||||
break
|
||||
line = line.decode("utf-8").rstrip("\r\n")
|
||||
if line == "":
|
||||
if cur_data:
|
||||
try:
|
||||
frame = json.loads("\n".join(cur_data))
|
||||
except Exception:
|
||||
frame = None
|
||||
if frame is not None:
|
||||
loop.call_soon_threadsafe(self._queue.put_nowait, frame)
|
||||
cur_data = []
|
||||
elif line.startswith(":"):
|
||||
continue # heartbeat comment
|
||||
elif line.startswith("data:"):
|
||||
cur_data.append(line[5:].lstrip())
|
||||
# id:/event: fields are not needed for the test shim
|
||||
loop.call_soon_threadsafe(self._queue.put_nowait, None) # EOF sentinel
|
||||
|
||||
async def send(self, json_str: str) -> None:
|
||||
await asyncio.to_thread(self._post_frame, json_str)
|
||||
|
||||
def _post_frame(self, json_str: str) -> None:
|
||||
conn = HTTPConnection("127.0.0.1", self._port, timeout=30)
|
||||
conn.request(
|
||||
"POST",
|
||||
"/v1/frame",
|
||||
body=json_str.encode("utf-8"),
|
||||
headers={
|
||||
"Authorization": f"Bearer {TOKEN}",
|
||||
"X-Iris-Device": DEVICE_ID,
|
||||
"Content-Type": "application/json",
|
||||
},
|
||||
)
|
||||
resp = conn.getresponse()
|
||||
body = resp.read()
|
||||
conn.close()
|
||||
# Fast responses (validation, channel ops, …) come back on the POST
|
||||
# response body via the reply sink, NOT the SSE stream. Enqueue any
|
||||
# protocol frame so recv_until sees it (the plain {"ok":true} ack is
|
||||
# not a frame and is skipped).
|
||||
if body:
|
||||
try:
|
||||
obj = json.loads(body)
|
||||
except Exception:
|
||||
obj = None
|
||||
if isinstance(obj, dict) and "type" in obj:
|
||||
assert self._loop is not None
|
||||
self._loop.call_soon_threadsafe(self._queue.put_nowait, obj)
|
||||
|
||||
async def recv(self, timeout: float = 10.0) -> dict:
|
||||
frame = await asyncio.wait_for(self._queue.get(), timeout=timeout)
|
||||
if frame is None:
|
||||
raise ConnectionError("SSE stream closed")
|
||||
return frame
|
||||
|
||||
async def upload(
|
||||
self,
|
||||
media_ref: str,
|
||||
data: bytes,
|
||||
*,
|
||||
kind: str = "image",
|
||||
mime: str = "image/png",
|
||||
filename: str = "t.png",
|
||||
sha256: str | None = None,
|
||||
) -> dict:
|
||||
"""Drive a media upload via ``POST /v1/media``; returns the terminal
|
||||
frame (``media.upload.ack`` or ``error``)."""
|
||||
|
||||
def _do() -> dict:
|
||||
conn = HTTPConnection("127.0.0.1", self._port, timeout=60)
|
||||
conn.request(
|
||||
"POST",
|
||||
"/v1/media",
|
||||
body=data,
|
||||
headers={
|
||||
"Authorization": f"Bearer {TOKEN}",
|
||||
"X-Iris-Device": DEVICE_ID,
|
||||
"Content-Type": mime,
|
||||
"X-Iris-Media-Ref": media_ref,
|
||||
"X-Iris-Media-Kind": kind,
|
||||
"X-Iris-Media-Filename": filename,
|
||||
"X-Iris-Media-Sha256": sha256
|
||||
or hashlib.sha256(data).hexdigest(),
|
||||
},
|
||||
)
|
||||
resp = conn.getresponse()
|
||||
body = resp.read()
|
||||
conn.close()
|
||||
return json.loads(body)
|
||||
|
||||
return await asyncio.to_thread(_do)
|
||||
|
||||
async def pull(self, media_id: str) -> tuple[int, bytes]:
|
||||
"""Drive a media pull via ``GET /v1/media/{id}``; returns (status, body)."""
|
||||
|
||||
def _do() -> tuple[int, bytes]:
|
||||
conn = HTTPConnection("127.0.0.1", self._port, timeout=30)
|
||||
conn.request(
|
||||
"GET",
|
||||
f"/v1/media/{media_id}",
|
||||
headers={"Authorization": f"Bearer {TOKEN}", "X-Iris-Device": DEVICE_ID},
|
||||
)
|
||||
resp = conn.getresponse()
|
||||
body = resp.read()
|
||||
status = resp.status
|
||||
conn.close()
|
||||
return status, body
|
||||
|
||||
return await asyncio.to_thread(_do)
|
||||
|
||||
async def close(self) -> None:
|
||||
# Interrupt the reader thread's blocking readline() by shutting down
|
||||
# the socket first; otherwise resp.close() blocks until the in-flight
|
||||
# read returns (the file lock is held for the whole blocking read).
|
||||
if self._resp is not None:
|
||||
sock = getattr(self._resp.fp, "raw", None)
|
||||
sock = getattr(sock, "_sock", None) if sock is not None else None
|
||||
if sock is not None:
|
||||
try:
|
||||
sock.shutdown(socket.SHUT_RDWR)
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
self._resp.close()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
@pytest_asyncio.fixture
|
||||
async def ws_client(adapter):
|
||||
"""Connected + paired WS client; the adapter's server runs on an
|
||||
ephemeral port for the duration of the test."""
|
||||
from websockets.asyncio.client import connect
|
||||
|
||||
"""Connected + paired HTTP client; the adapter's server runs on an
|
||||
ephemeral port for the duration of the test. Yields ``(client, ack)`` to
|
||||
match the old WS fixture shape so test bodies need no signature change."""
|
||||
await adapter.connect()
|
||||
client = HttpTestClient(adapter._http_server.bound_port)
|
||||
try:
|
||||
port = adapter._ws_server._server.sockets[0].getsockname()[1]
|
||||
async with connect(
|
||||
f"ws://127.0.0.1:{port}/ws", max_size=4 * 1024 * 1024
|
||||
) as ws:
|
||||
ack = await _hello(ws)
|
||||
yield ws, ack
|
||||
ack = await client.start()
|
||||
yield client, ack
|
||||
finally:
|
||||
await client.close()
|
||||
await adapter.disconnect()
|
||||
|
||||
|
||||
async def recv_until(ws, predicate, timeout: float = 10.0) -> list:
|
||||
"""Collect frames (dicts; binary frames as ("binary", bytes)) until
|
||||
*predicate* matches a JSON frame. Returns all frames collected."""
|
||||
"""Collect frames (dicts) until *predicate* matches a frame. Returns all
|
||||
frames collected."""
|
||||
frames: list = []
|
||||
loop = asyncio.get_running_loop()
|
||||
deadline = loop.time() + timeout
|
||||
@@ -172,13 +339,9 @@ async def recv_until(ws, predicate, timeout: float = 10.0) -> list:
|
||||
if remaining <= 0:
|
||||
raise AssertionError(
|
||||
"timed out waiting for frame; got: "
|
||||
+ ", ".join(f.get("type", "?") if isinstance(f, dict) else "binary" for f in frames)
|
||||
+ ", ".join(f.get("type", "?") if isinstance(f, dict) else "?" for f in frames)
|
||||
)
|
||||
raw = await asyncio.wait_for(ws.recv(), timeout=remaining)
|
||||
if isinstance(raw, (bytes, bytearray)):
|
||||
frames.append(("binary", bytes(raw)))
|
||||
continue
|
||||
frame = json.loads(raw)
|
||||
frame = await ws.recv(timeout=remaining)
|
||||
frames.append(frame)
|
||||
if predicate(frame):
|
||||
return frames
|
||||
@@ -186,43 +349,42 @@ async def recv_until(ws, predicate, timeout: float = 10.0) -> list:
|
||||
|
||||
async def upload_file(ws, media_ref: str, data: bytes, *, kind: str = "image",
|
||||
mime: str = "image/png", filename: str = "t.png",
|
||||
request_id: int = 1) -> dict:
|
||||
"""Drive a full media.upload flow; returns the terminal frame (ack or error)."""
|
||||
await ws.send(
|
||||
json.dumps(
|
||||
{
|
||||
"v": 1,
|
||||
"id": request_id,
|
||||
"type": "media.upload.start",
|
||||
"payload": {
|
||||
"media_ref": media_ref,
|
||||
"kind": kind,
|
||||
"mime": mime,
|
||||
"size": len(data),
|
||||
"filename": filename,
|
||||
},
|
||||
}
|
||||
)
|
||||
request_id: int = 1, sha256: str | None = None) -> dict:
|
||||
"""Drive a media upload via the HTTP leg; returns the terminal frame
|
||||
(ack or error)."""
|
||||
return await ws.upload(
|
||||
media_ref, data, kind=kind, mime=mime, filename=filename, sha256=sha256
|
||||
)
|
||||
# Two chunks to exercise reassembly.
|
||||
half = len(data) // 2
|
||||
await ws.send(data[:half])
|
||||
await ws.send(data[half:])
|
||||
await ws.send(
|
||||
json.dumps(
|
||||
{
|
||||
"v": 1,
|
||||
"id": request_id + 1,
|
||||
"type": "media.upload.end",
|
||||
"payload": {
|
||||
"media_ref": media_ref,
|
||||
"sha256": hashlib.sha256(data).hexdigest(),
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
# ── Lifecycle ───────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_disconnect_broadcasts_status_restarting(adapter):
|
||||
"""Teardown broadcasts ``status{state=restarting}`` before closing the
|
||||
streams, so the app can distinguish a clean gateway teardown (restart/
|
||||
stop) from a plain network drop — it shows the "Gateway restarting" chat
|
||||
notice only when this frame was received (docs/04 §status)."""
|
||||
await adapter.connect()
|
||||
client = HttpTestClient(adapter._http_server.bound_port)
|
||||
try:
|
||||
await client.start()
|
||||
await adapter.disconnect()
|
||||
frames = []
|
||||
while True:
|
||||
try:
|
||||
frames.append(await client.recv(timeout=5))
|
||||
except (ConnectionError, asyncio.TimeoutError):
|
||||
break
|
||||
statuses = [f for f in frames if f.get("type") == "status"]
|
||||
assert any(f["payload"]["state"] == "restarting" for f in statuses), (
|
||||
f"expected status{{restarting}} before close, got: {frames}"
|
||||
)
|
||||
)
|
||||
frames = await recv_until(ws, lambda f: f.get("type") in ("media.upload.ack", "error"))
|
||||
return frames[-1]
|
||||
finally:
|
||||
await client.close()
|
||||
# Idempotent: the second call is a no-op on the already-stopped server.
|
||||
await adapter.disconnect()
|
||||
|
||||
|
||||
# ── Pure helpers ────────────────────────────────────────────────────────────
|
||||
@@ -375,7 +537,6 @@ async def test_upload_reassembles_verifies_and_caches(adapter, ws_client):
|
||||
|
||||
terminal = await upload_file(ws, "mu_t1", PNG_1X1)
|
||||
assert terminal["type"] == "media.upload.ack", terminal
|
||||
assert terminal["id"] == 2
|
||||
assert terminal["payload"]["ok"] is True
|
||||
assert terminal["payload"]["media_ref"] == "mu_t1"
|
||||
|
||||
@@ -396,100 +557,37 @@ async def test_upload_reassembles_verifies_and_caches(adapter, ws_client):
|
||||
async def test_upload_declared_over_limit_rejected(adapter, ws_client):
|
||||
ws, _ = ws_client
|
||||
limit = adapter.max_upload_bytes
|
||||
await ws.send(
|
||||
json.dumps(
|
||||
{
|
||||
"v": 1,
|
||||
"id": 1,
|
||||
"type": "media.upload.start",
|
||||
"payload": {
|
||||
"media_ref": "mu_big",
|
||||
"kind": "document",
|
||||
"mime": "application/pdf",
|
||||
"size": limit + 1,
|
||||
"filename": "big.pdf",
|
||||
},
|
||||
}
|
||||
)
|
||||
# Over HTTP the server checks Content-Length before reading the body.
|
||||
err = await upload_file(
|
||||
ws, "mu_big", b"x" * (limit + 1), kind="document", mime="application/pdf",
|
||||
filename="big.pdf",
|
||||
)
|
||||
frames = await recv_until(ws, lambda f: f.get("type") == "error")
|
||||
err = frames[-1]
|
||||
assert err["type"] == "error"
|
||||
assert err["payload"]["code"] == "media_too_large"
|
||||
assert err["id"] == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_upload_midstream_over_limit_rejected(adapter, ws_client):
|
||||
async def test_upload_over_limit_ref_not_consumed(adapter, ws_client):
|
||||
"""Over HTTP the over-limit check happens before the body is read, so a
|
||||
rejected upload must not consume its media_ref (a later valid upload with
|
||||
the same ref succeeds)."""
|
||||
ws, _ = ws_client
|
||||
limit = adapter.max_upload_bytes
|
||||
await ws.send(
|
||||
json.dumps(
|
||||
{
|
||||
"v": 1,
|
||||
"id": 1,
|
||||
"type": "media.upload.start",
|
||||
"payload": {
|
||||
"media_ref": "mu_mid",
|
||||
"kind": "document",
|
||||
"mime": "application/octet-stream",
|
||||
"size": limit,
|
||||
"filename": "mid.bin",
|
||||
},
|
||||
}
|
||||
)
|
||||
)
|
||||
# Declared size passes the start check; the stream exceeds it.
|
||||
await ws.send(b"x" * (limit // 2))
|
||||
await ws.send(b"x" * (limit // 2 + 1))
|
||||
frames = await recv_until(ws, lambda f: f.get("type") == "error")
|
||||
assert frames[-1]["payload"]["code"] == "media_too_large"
|
||||
# The session is discarded: a late end cannot complete it.
|
||||
await ws.send(
|
||||
json.dumps(
|
||||
{
|
||||
"v": 1,
|
||||
"id": 2,
|
||||
"type": "media.upload.end",
|
||||
"payload": {"media_ref": "mu_mid", "sha256": "0" * 64},
|
||||
}
|
||||
)
|
||||
)
|
||||
frames = await recv_until(ws, lambda f: f.get("type") == "error" and f.get("id") == 2)
|
||||
assert frames[-1]["payload"]["code"] == "not_found"
|
||||
err = await upload_file(ws, "mu_reuse", b"x" * (limit + 1), kind="document")
|
||||
assert err["type"] == "error"
|
||||
assert err["payload"]["code"] == "media_too_large"
|
||||
# The ref is free: a valid upload with the same ref now succeeds.
|
||||
ok = await upload_file(ws, "mu_reuse", PNG_1X1)
|
||||
assert ok["type"] == "media.upload.ack", ok
|
||||
assert ok["payload"]["ok"] is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_upload_sha256_mismatch_rejected(adapter, ws_client):
|
||||
ws, _ = ws_client
|
||||
await ws.send(
|
||||
json.dumps(
|
||||
{
|
||||
"v": 1,
|
||||
"id": 1,
|
||||
"type": "media.upload.start",
|
||||
"payload": {
|
||||
"media_ref": "mu_bad",
|
||||
"kind": "image",
|
||||
"mime": "image/png",
|
||||
"size": len(PNG_1X1),
|
||||
"filename": "bad.png",
|
||||
},
|
||||
}
|
||||
)
|
||||
)
|
||||
await ws.send(PNG_1X1)
|
||||
await ws.send(
|
||||
json.dumps(
|
||||
{
|
||||
"v": 1,
|
||||
"id": 2,
|
||||
"type": "media.upload.end",
|
||||
"payload": {"media_ref": "mu_bad", "sha256": "0" * 64},
|
||||
}
|
||||
)
|
||||
)
|
||||
frames = await recv_until(ws, lambda f: f.get("type") == "error")
|
||||
assert frames[-1]["payload"]["code"] == "internal"
|
||||
err = await upload_file(ws, "mu_bad", PNG_1X1, sha256="0" * 64)
|
||||
assert err["type"] == "error"
|
||||
assert err["payload"]["code"] == "internal"
|
||||
assert adapter._media.get_inbound("mu_bad") is None
|
||||
|
||||
|
||||
@@ -1223,47 +1321,26 @@ async def test_pull_serves_allowed_path(adapter, ws_client):
|
||||
str(img), "image", "image/png", "pull_test.png", len(PNG_1X1)
|
||||
)
|
||||
|
||||
await ws.send(
|
||||
json.dumps(
|
||||
{"v": 1, "id": 9, "type": "media.pull", "payload": {"media_id": entry.media_id}}
|
||||
)
|
||||
)
|
||||
chunks: list[bytes] = []
|
||||
terminal = None
|
||||
while terminal is None:
|
||||
raw = await asyncio.wait_for(ws.recv(), timeout=10)
|
||||
if isinstance(raw, (bytes, bytearray)):
|
||||
chunks.append(bytes(raw))
|
||||
continue
|
||||
frame = json.loads(raw)
|
||||
if frame.get("type") == "media.pull.end":
|
||||
terminal = frame
|
||||
assert terminal["id"] == 9
|
||||
assert terminal["payload"]["ok"] is True
|
||||
assert b"".join(chunks) == PNG_1X1
|
||||
status, body = await ws.pull(entry.media_id)
|
||||
assert status == 200
|
||||
assert body == PNG_1X1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pull_rejects_unknown_and_denied(adapter, ws_client):
|
||||
ws, _ = ws_client
|
||||
# Unknown media_id.
|
||||
await ws.send(
|
||||
json.dumps({"v": 1, "id": 1, "type": "media.pull", "payload": {"media_id": "md_nope"}})
|
||||
)
|
||||
frames = await recv_until(ws, lambda f: f.get("type") == "error" and f.get("id") == 1)
|
||||
assert frames[-1]["payload"]["code"] == "not_found"
|
||||
status, body = await ws.pull("md_nope")
|
||||
assert status == 404
|
||||
assert json.loads(body)["payload"]["code"] == "not_found"
|
||||
|
||||
# Known id, but the path fails delivery validation (denylist).
|
||||
entry = adapter._media.register_outbound(
|
||||
"/etc/passwd", "document", "text/plain", "passwd", 100
|
||||
)
|
||||
await ws.send(
|
||||
json.dumps(
|
||||
{"v": 1, "id": 2, "type": "media.pull", "payload": {"media_id": entry.media_id}}
|
||||
)
|
||||
)
|
||||
frames = await recv_until(ws, lambda f: f.get("type") == "error" and f.get("id") == 2)
|
||||
assert frames[-1]["payload"]["code"] == "not_found"
|
||||
status, body = await ws.pull(entry.media_id)
|
||||
assert status == 404
|
||||
assert json.loads(body)["payload"]["code"] == "not_found"
|
||||
|
||||
# Known id, file deleted since the offer.
|
||||
from gateway.platforms.base import get_image_cache_dir
|
||||
@@ -1274,13 +1351,9 @@ async def test_pull_rejects_unknown_and_denied(adapter, ws_client):
|
||||
str(img), "image", "image/png", "gone.png", len(PNG_1X1)
|
||||
)
|
||||
img.unlink()
|
||||
await ws.send(
|
||||
json.dumps(
|
||||
{"v": 1, "id": 3, "type": "media.pull", "payload": {"media_id": entry2.media_id}}
|
||||
)
|
||||
)
|
||||
frames = await recv_until(ws, lambda f: f.get("type") == "error" and f.get("id") == 3)
|
||||
assert frames[-1]["payload"]["code"] == "not_found"
|
||||
status, body = await ws.pull(entry2.media_id)
|
||||
assert status == 404
|
||||
assert json.loads(body)["payload"]["code"] == "not_found"
|
||||
|
||||
|
||||
# ── M5: push backends (pure) ───────────────────────────────────────────────
|
||||
@@ -1653,7 +1726,7 @@ async def test_push_skipped_when_backend_unconfigured(adapter):
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fcm_register_updates_registry_and_live_conn(adapter, ws_client):
|
||||
async def test_fcm_register_updates_registry(adapter, ws_client):
|
||||
ws, _ = ws_client
|
||||
await ws.send(
|
||||
json.dumps(
|
||||
@@ -1664,7 +1737,7 @@ async def test_fcm_register_updates_registry_and_live_conn(adapter, ws_client):
|
||||
}
|
||||
)
|
||||
)
|
||||
# Ordering barrier: WS frames are processed in order, so by the time the
|
||||
# Ordering barrier: frames are processed in order, so by the time the
|
||||
# channel.list response arrives, fcm.register has been handled.
|
||||
await ws.send(
|
||||
json.dumps({"v": 1, "id": 1, "type": "channel.list", "payload": {}})
|
||||
@@ -1673,16 +1746,13 @@ async def test_fcm_register_updates_registry_and_live_conn(adapter, ws_client):
|
||||
dev = adapter._devices.get(DEVICE_ID)
|
||||
assert dev["fcm_token"] == "rotated-token"
|
||||
assert dev["ntfy_topic"] == "dev-topic"
|
||||
conn = adapter._ws_server.connection(DEVICE_ID)
|
||||
assert conn.fcm_token == "rotated-token"
|
||||
assert conn.ntfy_topic == "dev-topic"
|
||||
|
||||
# The next offline push targets the rotated token.
|
||||
# The next offline push targets the rotated token. Close the stream and
|
||||
# force-remove the (possibly stale) subscriber so the push isn't counted
|
||||
# as delivered to a dead connection.
|
||||
await ws.close()
|
||||
for _ in range(100):
|
||||
if not adapter._ws_server.has_devices():
|
||||
break
|
||||
await asyncio.sleep(0.01)
|
||||
with adapter._http_server._subs_lock:
|
||||
adapter._http_server._subs.clear()
|
||||
fake = _FakePush()
|
||||
adapter._push = fake
|
||||
await adapter.send("android:default", "after rotation", metadata={"notify": True})
|
||||
@@ -1898,32 +1968,32 @@ async def test_sync_replays_parked_frames_and_done_cursor(adapter):
|
||||
await adapter.send("android:default", "two", metadata={"notify": True})
|
||||
assert adapter._outbox.latest_cursor() == 2
|
||||
|
||||
from websockets.asyncio.client import connect
|
||||
|
||||
await adapter.connect()
|
||||
# Open at the latest cursor so the SSE catch-up doesn't replay the parked
|
||||
# frames (the sync request below is what we're testing).
|
||||
ws = HttpTestClient(adapter._http_server.bound_port, cursor=adapter._outbox.latest_cursor())
|
||||
try:
|
||||
port = adapter._ws_server._server.sockets[0].getsockname()[1]
|
||||
async with connect(f"ws://127.0.0.1:{port}/ws") as ws:
|
||||
ack = await _hello(ws)
|
||||
assert ack["payload"]["sync_cursor"] == 2
|
||||
await ws.send(
|
||||
json.dumps({"v": 1, "id": 10, "type": "sync", "payload": {"cursor": 0}})
|
||||
)
|
||||
frames = await recv_until(ws, lambda f: f.get("type") == "sync.done")
|
||||
texts = [f["payload"]["text"] for f in frames if f.get("type") == "message"]
|
||||
assert texts == ["one", "two"]
|
||||
done = frames[-1]
|
||||
assert done["id"] == 10
|
||||
assert done["payload"]["cursor"] == 2
|
||||
# A sync at the current cursor replays nothing.
|
||||
await ws.send(
|
||||
json.dumps({"v": 1, "id": 11, "type": "sync", "payload": {"cursor": 2}})
|
||||
)
|
||||
frames = await recv_until(
|
||||
ws, lambda f: f.get("type") == "sync.done" and f.get("id") == 11
|
||||
)
|
||||
assert len(frames) == 1
|
||||
ack = await ws.start()
|
||||
assert ack["payload"]["sync_cursor"] == 2
|
||||
await ws.send(
|
||||
json.dumps({"v": 1, "id": 10, "type": "sync", "payload": {"cursor": 0}})
|
||||
)
|
||||
frames = await recv_until(ws, lambda f: f.get("type") == "sync.done")
|
||||
texts = [f["payload"]["text"] for f in frames if f.get("type") == "message"]
|
||||
assert texts == ["one", "two"]
|
||||
done = frames[-1]
|
||||
assert done["id"] == 10
|
||||
assert done["payload"]["cursor"] == 2
|
||||
# A sync at the current cursor replays nothing.
|
||||
await ws.send(
|
||||
json.dumps({"v": 1, "id": 11, "type": "sync", "payload": {"cursor": 2}})
|
||||
)
|
||||
frames = await recv_until(
|
||||
ws, lambda f: f.get("type") == "sync.done" and f.get("id") == 11
|
||||
)
|
||||
assert len(frames) == 1
|
||||
finally:
|
||||
await ws.close()
|
||||
await adapter.disconnect()
|
||||
|
||||
|
||||
@@ -1931,15 +2001,13 @@ async def test_sync_replays_parked_frames_and_done_cursor(adapter):
|
||||
async def test_hello_ack_last_pushed_cursor_default_zero(adapter):
|
||||
"""A device that never received a push reports last_pushed_cursor=0 in
|
||||
hello.ack (docs/08 §8.7 dedupe watermark)."""
|
||||
from websockets.asyncio.client import connect
|
||||
|
||||
await adapter.connect()
|
||||
ws = HttpTestClient(adapter._http_server.bound_port)
|
||||
try:
|
||||
port = adapter._ws_server._server.sockets[0].getsockname()[1]
|
||||
async with connect(f"ws://127.0.0.1:{port}/ws") as ws:
|
||||
ack = await _hello(ws)
|
||||
assert ack["payload"]["last_pushed_cursor"] == 0
|
||||
ack = await ws.start()
|
||||
assert ack["payload"]["last_pushed_cursor"] == 0
|
||||
finally:
|
||||
await ws.close()
|
||||
await adapter.disconnect()
|
||||
|
||||
|
||||
@@ -1949,8 +2017,6 @@ async def test_push_success_advances_last_pushed_cursor(adapter):
|
||||
next hello.ack reports it — the app uses it to skip re-notifying
|
||||
sync-replayed frames (docs/08 §8.7). Back-to-back frames for the same
|
||||
chat coalesce into one push; a failed push does not advance the cursor."""
|
||||
from websockets.asyncio.client import connect
|
||||
|
||||
fake = _FakePush()
|
||||
adapter._push = fake
|
||||
adapter._devices.upsert(DEVICE_ID, "Test", {}, fcm_token="tok-1")
|
||||
@@ -1978,12 +2044,12 @@ async def test_push_success_advances_last_pushed_cursor(adapter):
|
||||
assert adapter._devices.last_pushed_cursor(DEVICE_ID) == 3
|
||||
|
||||
await adapter.connect()
|
||||
ws = HttpTestClient(adapter._http_server.bound_port, cursor=adapter._outbox.latest_cursor())
|
||||
try:
|
||||
port = adapter._ws_server._server.sockets[0].getsockname()[1]
|
||||
async with connect(f"ws://127.0.0.1:{port}/ws") as ws:
|
||||
ack = await _hello(ws)
|
||||
assert ack["payload"]["last_pushed_cursor"] == 3
|
||||
ack = await ws.start()
|
||||
assert ack["payload"]["last_pushed_cursor"] == 3
|
||||
finally:
|
||||
await ws.close()
|
||||
await adapter.disconnect()
|
||||
|
||||
|
||||
@@ -1992,25 +2058,23 @@ async def test_sync_replay_frames_carry_outbox_cursor(adapter):
|
||||
"""Frames replayed by sync carry their outbox cursor in the envelope so
|
||||
the app can compare it against last_pushed_cursor (docs/08 §8.7). Live
|
||||
frames carry no cursor."""
|
||||
from websockets.asyncio.client import connect
|
||||
|
||||
await adapter.send("android:default", "one", metadata={"notify": True})
|
||||
await adapter.send("android:default", "two", metadata={"notify": True})
|
||||
|
||||
await adapter.connect()
|
||||
ws = HttpTestClient(adapter._http_server.bound_port, cursor=adapter._outbox.latest_cursor())
|
||||
try:
|
||||
port = adapter._ws_server._server.sockets[0].getsockname()[1]
|
||||
async with connect(f"ws://127.0.0.1:{port}/ws") as ws:
|
||||
await _hello(ws)
|
||||
await ws.send(
|
||||
json.dumps({"v": 1, "id": 20, "type": "sync", "payload": {"cursor": 0}})
|
||||
)
|
||||
frames = await recv_until(ws, lambda f: f.get("type") == "sync.done")
|
||||
msgs = [f for f in frames if f.get("type") == "message"]
|
||||
assert [f["cursor"] for f in msgs] == [1, 2]
|
||||
# sync.done itself carries no envelope cursor.
|
||||
assert "cursor" not in frames[-1]
|
||||
await ws.start()
|
||||
await ws.send(
|
||||
json.dumps({"v": 1, "id": 20, "type": "sync", "payload": {"cursor": 0}})
|
||||
)
|
||||
frames = await recv_until(ws, lambda f: f.get("type") == "sync.done")
|
||||
msgs = [f for f in frames if f.get("type") == "message"]
|
||||
assert [f["cursor"] for f in msgs] == [1, 2]
|
||||
# sync.done itself carries no envelope cursor.
|
||||
assert "cursor" not in frames[-1]
|
||||
finally:
|
||||
await ws.close()
|
||||
await adapter.disconnect()
|
||||
|
||||
|
||||
@@ -2129,30 +2193,27 @@ def test_channels_delete_hard_deletes_row_and_child_threads(plugin, tmp_path):
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_wrong_token_rejected(adapter):
|
||||
from websockets.asyncio.client import connect
|
||||
|
||||
"""A request with a wrong Bearer token is rejected with 401 (the HTTP
|
||||
equivalent of the WS hello auth rejection)."""
|
||||
await adapter.connect()
|
||||
try:
|
||||
port = adapter._ws_server._server.sockets[0].getsockname()[1]
|
||||
async with connect(f"ws://127.0.0.1:{port}/ws") as ws:
|
||||
await ws.send(
|
||||
json.dumps(
|
||||
{
|
||||
"v": 1,
|
||||
"type": "hello",
|
||||
"payload": {
|
||||
"token": "wrong-token",
|
||||
"device_id": DEVICE_ID,
|
||||
"device_name": "Bad",
|
||||
"caps": {},
|
||||
},
|
||||
}
|
||||
)
|
||||
port = adapter._http_server.bound_port
|
||||
|
||||
def _req() -> int:
|
||||
conn = HTTPConnection("127.0.0.1", port, timeout=5)
|
||||
conn.request(
|
||||
"GET",
|
||||
"/v1/events",
|
||||
headers={"Authorization": "Bearer wrong-token", "X-Iris-Device": DEVICE_ID},
|
||||
)
|
||||
raw = await asyncio.wait_for(ws.recv(), timeout=5)
|
||||
err = json.loads(raw)
|
||||
assert err["type"] == "error"
|
||||
assert err["payload"]["code"] == "auth"
|
||||
resp = conn.getresponse()
|
||||
resp.read()
|
||||
status = resp.status
|
||||
conn.close()
|
||||
return status
|
||||
|
||||
status = await asyncio.to_thread(_req)
|
||||
assert status == 401
|
||||
finally:
|
||||
await adapter.disconnect()
|
||||
|
||||
|
||||
@@ -0,0 +1,826 @@
|
||||
"""Tests for the android plugin's HTTP fallback transport (docs/19).
|
||||
|
||||
The plugin lives in the sibling ``iris_x_hermes`` checkout; tests load it
|
||||
from the source tree directly (same pattern as ``test_android.py``).
|
||||
|
||||
Coverage (docs/19 §19.12):
|
||||
* auth: bad/missing token -> 401; missing device header -> 401;
|
||||
allowlist rejection -> 401
|
||||
* ``POST /v1/frame``: valid ``message.send`` dispatches (202 + echo on
|
||||
the SSE stream); empty text -> 400 error frame; automation channel ->
|
||||
400; bad JSON -> 400; wrong content-type -> 400; oversize body -> 413;
|
||||
media frames -> 400 (WS-only in v1); rate limit -> 429
|
||||
* SSE: catch-up rows carry correct ``id``s + cursor envelope;
|
||||
``event: hello`` present; a live frame appended after connect arrives
|
||||
on the stream; ``Last-Event-ID`` resume replays exactly the delta;
|
||||
heartbeat observed
|
||||
* long-poll: returns on new frame; empty 200 at timeout with advanced
|
||||
cursor
|
||||
* **delivery counting (docs/19 §19.8)**: a frame with only an SSE
|
||||
subscriber is ``delivered >= 1`` -> NO push fired (the critical
|
||||
regression test)
|
||||
|
||||
Run via ``scripts/run_tests.sh tests/gateway/test_android_http.py``.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import base64
|
||||
import contextlib
|
||||
import importlib.util
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
from http.client import HTTPConnection
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
import pytest
|
||||
import pytest_asyncio
|
||||
|
||||
# Test-only token (not a credential; the adapter is built with it via
|
||||
# monkeypatch in the fixture below).
|
||||
# pi-lens-ignore: S105
|
||||
TOKEN = "test-android-http-token-0123456789"
|
||||
DEVICE_ID = "test-http-device"
|
||||
CHAT_ID = "android:default"
|
||||
|
||||
# 1x1 PNG (same fixture as test_android.py).
|
||||
PNG_1X1 = base64.b64decode(
|
||||
"iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJ"
|
||||
"AAAADUlEQVR42mNkYPhfDwAChwGA60e6kgAAAABJRU5ErkJggg=="
|
||||
)
|
||||
|
||||
|
||||
def _plugin_dir() -> Path:
|
||||
env = os.environ.get("ANDROID_PLUGIN_DIR")
|
||||
if env:
|
||||
return Path(env)
|
||||
# Works from either copy of this file: gateway-plugin/tests/ (canonical,
|
||||
# plugin dir is parents[1]) or the hermes-agent/tests/gateway/ mirror
|
||||
# (repo root is parents[3]).
|
||||
here = Path(__file__).resolve()
|
||||
for candidate in (here.parents[1], here.parents[3] / "gateway-plugin"):
|
||||
if (candidate / "protocol.py").is_file():
|
||||
return candidate
|
||||
return here.parents[1]
|
||||
|
||||
|
||||
def _load_plugin():
|
||||
"""Load the gateway-plugin package under a unique module name (same
|
||||
pattern as test_android.py)."""
|
||||
name = "android_plugin_http_under_test"
|
||||
cached = sys.modules.get(name)
|
||||
if cached is not None:
|
||||
return cached
|
||||
pkg_dir = _plugin_dir()
|
||||
if not (pkg_dir / "__init__.py").is_file():
|
||||
pytest.fail(f"android plugin not found at {pkg_dir}")
|
||||
spec = importlib.util.spec_from_file_location(
|
||||
name, pkg_dir / "__init__.py", submodule_search_locations=[str(pkg_dir)]
|
||||
)
|
||||
if spec is None or spec.loader is None:
|
||||
pytest.fail(f"could not build import spec for {pkg_dir}")
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
sys.modules[name] = module
|
||||
try:
|
||||
spec.loader.exec_module(module)
|
||||
except Exception:
|
||||
sys.modules.pop(name, None)
|
||||
raise
|
||||
return module
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def plugin():
|
||||
return _load_plugin()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def adapter(plugin, monkeypatch):
|
||||
"""A live AndroidAdapter with an isolated HERMES_HOME (conftest)."""
|
||||
monkeypatch.setenv("ANDROID_TOKEN", TOKEN)
|
||||
from gateway.platform_registry import PlatformEntry, platform_registry
|
||||
|
||||
if not platform_registry.is_registered("android"):
|
||||
platform_registry.register(
|
||||
PlatformEntry(
|
||||
name="android",
|
||||
label="Android",
|
||||
adapter_factory=lambda cfg: None,
|
||||
check_fn=lambda: True,
|
||||
)
|
||||
)
|
||||
|
||||
config = SimpleNamespace(
|
||||
extra={
|
||||
"host": "127.0.0.1",
|
||||
"port": 0, # ephemeral WS port
|
||||
"http_port": 0, # ephemeral HTTP port
|
||||
"max_upload_bytes": 1024 * 1024,
|
||||
},
|
||||
home_channel=None,
|
||||
)
|
||||
a = plugin.adapter.AndroidAdapter(config)
|
||||
yield a
|
||||
with contextlib.suppress(Exception):
|
||||
a._devices.close()
|
||||
with contextlib.suppress(Exception):
|
||||
a._outbox.close()
|
||||
|
||||
|
||||
@pytest_asyncio.fixture
|
||||
async def gw(adapter):
|
||||
"""Connected adapter (WS + HTTP legs up); the HTTP port is ephemeral."""
|
||||
await adapter.connect()
|
||||
try:
|
||||
yield adapter
|
||||
finally:
|
||||
await adapter.disconnect()
|
||||
|
||||
|
||||
def http_port(adapter) -> int:
|
||||
assert adapter._http_server.enabled, "HTTP leg should be enabled after connect()"
|
||||
return adapter._http_server.bound_port
|
||||
|
||||
|
||||
# ── Blocking HTTP helpers (run via asyncio.to_thread) ──────────────────────
|
||||
|
||||
|
||||
def _request(
|
||||
port: int,
|
||||
method: str,
|
||||
path: str,
|
||||
*,
|
||||
token: str | None = TOKEN,
|
||||
device: str | None = DEVICE_ID,
|
||||
body: bytes | str | None = None,
|
||||
content_type: str = "application/json",
|
||||
timeout: float = 10.0,
|
||||
extra_headers: dict | None = None,
|
||||
) -> tuple[int, bytes]:
|
||||
conn = HTTPConnection("127.0.0.1", port, timeout=timeout)
|
||||
headers = {}
|
||||
if token is not None:
|
||||
headers["Authorization"] = f"Bearer {token}"
|
||||
if device is not None:
|
||||
headers["X-Iris-Device"] = device
|
||||
if extra_headers:
|
||||
headers.update(extra_headers)
|
||||
if body is not None:
|
||||
data = body if isinstance(body, bytes) else body.encode("utf-8")
|
||||
headers["Content-Type"] = content_type
|
||||
conn.request(method, path, body=data, headers=headers)
|
||||
else:
|
||||
conn.request(method, path, headers=headers)
|
||||
resp = conn.getresponse()
|
||||
payload = resp.read()
|
||||
status = resp.status
|
||||
conn.close()
|
||||
return status, payload
|
||||
|
||||
|
||||
def _post_frame(port: int, frame: dict, **kw) -> tuple[int, dict]:
|
||||
status, payload = _request(port, "POST", "/v1/frame", body=json.dumps(frame), **kw)
|
||||
return status, json.loads(payload)
|
||||
|
||||
|
||||
def _frame_json(frame: dict) -> dict:
|
||||
return {"v": 1, **frame}
|
||||
|
||||
|
||||
def _parse_sse(lines: list[str]) -> tuple[list[tuple[str | None, str | None, str]], int]:
|
||||
"""Parse raw SSE lines into ``[(event, id, data), ...]`` + comment count."""
|
||||
events: list[tuple[str | None, str | None, str]] = []
|
||||
comments = 0
|
||||
cur_event: str | None = None
|
||||
cur_id: str | None = None
|
||||
cur_data: list[str] = []
|
||||
for raw in lines:
|
||||
line = raw.rstrip("\r\n")
|
||||
if line == "":
|
||||
if cur_data:
|
||||
events.append((cur_event, cur_id, "\n".join(cur_data)))
|
||||
cur_event, cur_id, cur_data = None, None, []
|
||||
elif line.startswith(":"):
|
||||
comments += 1
|
||||
else:
|
||||
field, _, value = line.partition(":")
|
||||
if value.startswith(" "):
|
||||
value = value[1:]
|
||||
if field == "event":
|
||||
cur_event = value
|
||||
elif field == "id":
|
||||
cur_id = value
|
||||
elif field == "data":
|
||||
cur_data.append(value)
|
||||
return events, comments
|
||||
|
||||
|
||||
def _sse_open(port: int, *, cursor: int | None = None, last_event_id: str | None = None):
|
||||
"""Open an SSE connection (blocking); returns the HTTPResponse (read
|
||||
lines via ``_sse_read_lines``; close with ``resp.close()``)."""
|
||||
conn = HTTPConnection("127.0.0.1", port, timeout=30)
|
||||
path = "/v1/events" + (f"?cursor={cursor}" if cursor is not None else "")
|
||||
headers = {
|
||||
"Authorization": f"Bearer {TOKEN}",
|
||||
"X-Iris-Device": DEVICE_ID,
|
||||
}
|
||||
if last_event_id is not None:
|
||||
headers["Last-Event-ID"] = last_event_id
|
||||
conn.request("GET", path, headers=headers)
|
||||
resp = conn.getresponse()
|
||||
assert resp.status == 200, f"SSE open failed: {resp.status}"
|
||||
assert resp.getheader("Content-Type", "").startswith("text/event-stream")
|
||||
return resp
|
||||
|
||||
|
||||
def _sse_read_lines(resp, n: int, timeout: float = 10.0) -> list[str]:
|
||||
"""Read up to n lines from the SSE stream (blocking)."""
|
||||
raw = resp.fp.raw
|
||||
sock = getattr(raw, "_sock", None)
|
||||
if sock is not None:
|
||||
sock.settimeout(timeout)
|
||||
lines: list[str] = []
|
||||
while len(lines) < n:
|
||||
line = resp.fp.readline()
|
||||
if not line:
|
||||
break
|
||||
lines.append(line.decode("utf-8"))
|
||||
return lines
|
||||
|
||||
|
||||
# ── /v1/health ──────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_health_no_auth(gw):
|
||||
status, payload = await asyncio.to_thread(
|
||||
_request, http_port(gw), "GET", "/v1/health", token=None, device=None
|
||||
)
|
||||
assert status == 200
|
||||
assert json.loads(payload) == {"ok": True}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unknown_path_404(gw):
|
||||
status, _ = await asyncio.to_thread(_request, http_port(gw), "GET", "/v1/nope")
|
||||
assert status == 404
|
||||
|
||||
|
||||
# ── Auth ────────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_post_bad_token_401(gw):
|
||||
status, _ = await asyncio.to_thread(
|
||||
_request,
|
||||
http_port(gw),
|
||||
"POST",
|
||||
"/v1/frame",
|
||||
token="wrong-token",
|
||||
body=json.dumps(_frame_json({"type": "ping", "payload": {}})),
|
||||
)
|
||||
assert status == 401
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_post_missing_token_401(gw):
|
||||
status, _ = await asyncio.to_thread(
|
||||
_request,
|
||||
http_port(gw),
|
||||
"POST",
|
||||
"/v1/frame",
|
||||
token=None,
|
||||
body=json.dumps(_frame_json({"type": "ping", "payload": {}})),
|
||||
)
|
||||
assert status == 401
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_post_missing_device_401(gw):
|
||||
status, _ = await asyncio.to_thread(
|
||||
_request,
|
||||
http_port(gw),
|
||||
"POST",
|
||||
"/v1/frame",
|
||||
device=None,
|
||||
body=json.dumps(_frame_json({"type": "ping", "payload": {}})),
|
||||
)
|
||||
assert status == 401
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_post_allowlist_rejection_401(gw):
|
||||
gw.allowed_users = ["some-other-device"]
|
||||
gw.allow_all = False
|
||||
status, _ = await asyncio.to_thread(
|
||||
_request,
|
||||
http_port(gw),
|
||||
"POST",
|
||||
"/v1/frame",
|
||||
body=json.dumps(_frame_json({"type": "ping", "payload": {}})),
|
||||
)
|
||||
assert status == 401
|
||||
|
||||
|
||||
# ── POST /v1/frame: validation ─────────────────────────────────────────────
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_post_bad_json_400(gw):
|
||||
status, payload = await asyncio.to_thread(
|
||||
_request, http_port(gw), "POST", "/v1/frame", body=b"not json"
|
||||
)
|
||||
assert status == 400
|
||||
frame = json.loads(payload)
|
||||
assert frame["type"] == "error"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_post_wrong_content_type_400(gw):
|
||||
status, _ = await asyncio.to_thread(
|
||||
_request,
|
||||
http_port(gw),
|
||||
"POST",
|
||||
"/v1/frame",
|
||||
body=json.dumps(_frame_json({"type": "ping", "payload": {}})),
|
||||
content_type="text/plain",
|
||||
)
|
||||
assert status == 400
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_post_oversize_body_413(gw):
|
||||
big = json.dumps(_frame_json({"type": "ping", "payload": {"pad": "x" * (1024 * 1024 + 1)}}))
|
||||
status, _ = await asyncio.to_thread(_request, http_port(gw), "POST", "/v1/frame", body=big)
|
||||
assert status == 413
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_post_empty_message_400(gw):
|
||||
gw.handle_message = AsyncMock()
|
||||
status, payload = await asyncio.to_thread(
|
||||
_post_frame,
|
||||
http_port(gw),
|
||||
_frame_json(
|
||||
{"id": 7, "type": "message.send", "chat_id": CHAT_ID, "payload": {"text": " "}}
|
||||
),
|
||||
)
|
||||
assert status == 400
|
||||
frame = payload
|
||||
assert frame["type"] == "error"
|
||||
assert frame["id"] == 7
|
||||
assert frame["payload"]["code"] == "unsupported"
|
||||
gw.handle_message.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_post_automation_channel_400(gw):
|
||||
gw.handle_message = AsyncMock()
|
||||
entry = gw._channels.create(name="Cron")
|
||||
gw._channels.set_automation(entry["chat_id"], True)
|
||||
status, payload = await asyncio.to_thread(
|
||||
_post_frame,
|
||||
http_port(gw),
|
||||
_frame_json(
|
||||
{
|
||||
"id": 8,
|
||||
"type": "message.send",
|
||||
"chat_id": entry["chat_id"],
|
||||
"payload": {"text": "hi"},
|
||||
}
|
||||
),
|
||||
)
|
||||
assert status == 400
|
||||
assert payload["type"] == "error"
|
||||
gw.handle_message.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_post_rate_limit_429(gw):
|
||||
gw.handle_message = AsyncMock()
|
||||
port = http_port(gw)
|
||||
# Exhaust the per-device bucket (INBOUND_BURST = 40) then expect 429.
|
||||
got_429 = False
|
||||
for i in range(60):
|
||||
status, _ = await asyncio.to_thread(
|
||||
_post_frame,
|
||||
port,
|
||||
_frame_json({"id": i, "type": "ping", "payload": {}}),
|
||||
)
|
||||
if status == 429:
|
||||
got_429 = True
|
||||
break
|
||||
assert got_429, "expected a 429 within 60 rapid frames"
|
||||
|
||||
|
||||
# ── POST /v1/frame: dispatch ────────────────────────────────────────────────
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_post_message_send_dispatches(gw):
|
||||
"""202 ack; the user echo + read receipt arrive on the SSE stream; the
|
||||
agent turn fires (docs/19 §19.7: async responses on the event stream)."""
|
||||
gw.handle_message = AsyncMock()
|
||||
port = http_port(gw)
|
||||
conn = _sse_open(port)
|
||||
try:
|
||||
# Consume the open sequence (hello + status = 6 lines) first.
|
||||
lines = _sse_read_lines(conn, 6, timeout=5)
|
||||
events, _ = _parse_sse(lines)
|
||||
assert events[0][0] == "hello"
|
||||
|
||||
status, payload = await asyncio.to_thread(
|
||||
_post_frame,
|
||||
port,
|
||||
_frame_json(
|
||||
{
|
||||
"id": 42,
|
||||
"type": "message.send",
|
||||
"chat_id": CHAT_ID,
|
||||
"payload": {"text": "hi there"},
|
||||
}
|
||||
),
|
||||
)
|
||||
# The read receipt (sent when the turn is handed to the agent) is
|
||||
# the handler's single point-to-point reply -> 200 with the frame
|
||||
# as the body (a plain 202 {"ok": true} is also valid when no
|
||||
# synchronous reply exists).
|
||||
assert status in (200, 202)
|
||||
if status == 200:
|
||||
assert payload["type"] == "read.receipt"
|
||||
else:
|
||||
assert payload == {"ok": True}
|
||||
|
||||
# The echo must arrive on the stream, tagged with its outbox
|
||||
# cursor as the SSE id (live frames carry no cursor in the
|
||||
# envelope, same as the WS path).
|
||||
deadline = time.monotonic() + 10
|
||||
echo = None
|
||||
while time.monotonic() < deadline and echo is None:
|
||||
lines = _sse_read_lines(conn, 4, timeout=5)
|
||||
for _event, sse_id, data in _parse_sse(lines)[0]:
|
||||
frame = json.loads(data)
|
||||
if (
|
||||
frame.get("type") == "message"
|
||||
and frame.get("payload", {}).get("text") == "hi there"
|
||||
):
|
||||
echo = (frame, sse_id)
|
||||
assert echo is not None, "user echo did not arrive on the SSE stream"
|
||||
assert echo[1] is not None # SSE id = outbox cursor
|
||||
await asyncio.sleep(0.2)
|
||||
gw.handle_message.assert_called_once()
|
||||
finally:
|
||||
conn.close()
|
||||
|
||||
|
||||
# ── SSE: catch-up, hello, live, resume, heartbeat ──────────────────────────
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sse_catchup_and_hello(gw):
|
||||
"""Catch-up rows carry correct ids + cursor envelope; hello present."""
|
||||
port = http_port(gw)
|
||||
# Park two frames through the real outbound path (no live devices).
|
||||
push_calls: list = []
|
||||
gw._maybe_push = AsyncMock(side_effect=lambda *a, **k: push_calls.append(a))
|
||||
for text in ("one", "two"):
|
||||
await _park_frame(gw, text)
|
||||
cursors = [1, 2]
|
||||
|
||||
conn = _sse_open(port, cursor=0)
|
||||
try:
|
||||
# 2 replayed frames (id + event + data + blank = 4 lines each) +
|
||||
# hello (3 lines) + status (3 lines) = 14 lines.
|
||||
lines = _sse_read_lines(conn, 14, timeout=5)
|
||||
events, _ = _parse_sse(lines)
|
||||
assert events[0][0] == "frame"
|
||||
assert events[0][1] == str(cursors[0])
|
||||
f0 = json.loads(events[0][2])
|
||||
assert f0["payload"]["text"] == "one"
|
||||
assert f0["cursor"] == cursors[0]
|
||||
assert events[1][0] == "frame"
|
||||
assert events[1][1] == str(cursors[1])
|
||||
assert json.loads(events[1][2])["payload"]["text"] == "two"
|
||||
assert events[2][0] == "hello"
|
||||
hello = json.loads(events[2][2])
|
||||
assert hello["type"] == "hello.ack"
|
||||
assert hello["payload"]["sync_cursor"] == 2
|
||||
assert events[3][0] == "frame"
|
||||
assert json.loads(events[3][2])["type"] == "status"
|
||||
finally:
|
||||
conn.close()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sse_live_frame_after_connect(gw):
|
||||
port = http_port(gw)
|
||||
conn = _sse_open(port)
|
||||
try:
|
||||
# Consume the open sequence (hello + status = 6 lines).
|
||||
_sse_read_lines(conn, 6, timeout=5)
|
||||
await _park_frame(gw, "live!")
|
||||
deadline = time.monotonic() + 10
|
||||
got = None
|
||||
while time.monotonic() < deadline and got is None:
|
||||
lines = _sse_read_lines(conn, 4, timeout=5)
|
||||
for _event, sse_id, data in _parse_sse(lines)[0]:
|
||||
frame = json.loads(data)
|
||||
if frame.get("payload", {}).get("text") == "live!":
|
||||
got = (frame, sse_id)
|
||||
assert got is not None, "live frame did not arrive on the SSE stream"
|
||||
assert got[1] is not None # SSE id = outbox cursor
|
||||
finally:
|
||||
conn.close()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sse_last_event_id_resume(gw):
|
||||
"""Resume with Last-Event-ID replays exactly the delta."""
|
||||
port = http_port(gw)
|
||||
for text in ("a", "b", "c"):
|
||||
await _park_frame(gw, text)
|
||||
|
||||
conn = _sse_open(port, last_event_id="1")
|
||||
try:
|
||||
# Frames 2 and 3 replayed (8 lines) + hello (3) + status (3) = 14.
|
||||
lines = _sse_read_lines(conn, 14, timeout=5)
|
||||
events, _ = _parse_sse(lines)
|
||||
replayed = [e for e in events if e[0] == "frame" and e[1] is not None]
|
||||
assert [e[1] for e in replayed] == ["2", "3"]
|
||||
assert json.loads(replayed[0][2])["payload"]["text"] == "b"
|
||||
assert json.loads(replayed[1][2])["payload"]["text"] == "c"
|
||||
finally:
|
||||
conn.close()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sse_heartbeat(gw, monkeypatch, plugin):
|
||||
"""A comment heartbeat is written when the stream is idle."""
|
||||
monkeypatch.setattr(plugin.http_server, "SSE_HEARTBEAT_S", 1.0)
|
||||
port = http_port(gw)
|
||||
conn = _sse_open(port)
|
||||
try:
|
||||
# Consume the open sequence (6 lines), then wait for the heartbeat.
|
||||
_sse_read_lines(conn, 6, timeout=5)
|
||||
lines = _sse_read_lines(conn, 2, timeout=5)
|
||||
events, comments = _parse_sse(lines)
|
||||
assert comments >= 1, f"no heartbeat comment in {lines!r}"
|
||||
assert events == []
|
||||
finally:
|
||||
conn.close()
|
||||
|
||||
|
||||
# ── Long-poll ───────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_poll_returns_on_new_frame(gw):
|
||||
port = http_port(gw)
|
||||
await _park_frame(gw, "pre") # cursor 1: returned immediately (catch-up)
|
||||
|
||||
# A second poll at the high-water mark blocks until a new frame lands.
|
||||
def poll_and_broadcast():
|
||||
status, payload = _request(port, "GET", "/v1/poll?cursor=1", timeout=30)
|
||||
return status, json.loads(payload)
|
||||
|
||||
async def late_frame():
|
||||
await asyncio.sleep(0.5)
|
||||
await _park_frame(gw, "late")
|
||||
|
||||
poll_task = asyncio.create_task(asyncio.to_thread(poll_and_broadcast))
|
||||
late_task = asyncio.create_task(late_frame())
|
||||
status, body = await asyncio.wait_for(poll_task, timeout=15)
|
||||
await late_task
|
||||
assert status == 200
|
||||
assert body["cursor"] >= 2
|
||||
assert len(body["frames"]) == 1
|
||||
assert json.loads(body["frames"][0])["payload"]["text"] == "late"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_poll_timeout_empty(gw, monkeypatch, plugin):
|
||||
monkeypatch.setattr(plugin.http_server, "POLL_TIMEOUT_S", 1.0)
|
||||
port = http_port(gw)
|
||||
await _park_frame(gw, "x")
|
||||
hwm = gw._outbox.latest_cursor()
|
||||
status, payload = await asyncio.to_thread(
|
||||
_request, port, "GET", f"/v1/poll?cursor={hwm}", timeout=15
|
||||
)
|
||||
body = json.loads(payload)
|
||||
assert status == 200
|
||||
assert body["frames"] == []
|
||||
assert body["cursor"] == hwm
|
||||
|
||||
|
||||
# ── Delivery counting (docs/19 §19.8 — the critical regression) ────────────
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sse_subscriber_counts_as_delivered_no_push(gw):
|
||||
"""A frame with only an SSE subscriber is delivered >= 1 -> NO push."""
|
||||
port = http_port(gw)
|
||||
conn = _sse_open(port)
|
||||
try:
|
||||
_sse_read_lines(conn, 6, timeout=5) # open sequence
|
||||
push = AsyncMock()
|
||||
gw._maybe_push = push
|
||||
await _park_frame(gw, "no push for me")
|
||||
await asyncio.sleep(0.2)
|
||||
push.assert_not_called()
|
||||
finally:
|
||||
conn.close()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_no_subscribers_still_pushes(gw):
|
||||
"""Control: with no live devices at all, the push path still fires."""
|
||||
push = AsyncMock()
|
||||
gw._maybe_push = push
|
||||
await _park_frame(gw, "wake me up")
|
||||
await asyncio.sleep(0.2)
|
||||
push.assert_called_once()
|
||||
|
||||
|
||||
# ── Media over HTTP (docs/19 §19.15, v2) ──────────────────────────────────
|
||||
|
||||
|
||||
def _upload(
|
||||
port: int,
|
||||
data: bytes,
|
||||
*,
|
||||
media_ref: str = "mu_http1",
|
||||
kind: str = "image",
|
||||
mime: str = "image/png",
|
||||
filename: str = "t.png",
|
||||
sha256: str | None = None,
|
||||
**kw,
|
||||
) -> tuple[int, dict]:
|
||||
import hashlib
|
||||
|
||||
headers = {
|
||||
"X-Iris-Media-Ref": media_ref,
|
||||
"X-Iris-Media-Kind": kind,
|
||||
"X-Iris-Media-Filename": filename,
|
||||
"X-Iris-Media-Sha256": sha256 if sha256 is not None else hashlib.sha256(data).hexdigest(),
|
||||
}
|
||||
status, payload = _request(
|
||||
port,
|
||||
"POST",
|
||||
"/v1/media",
|
||||
body=data,
|
||||
content_type=mime,
|
||||
extra_headers=headers,
|
||||
**kw,
|
||||
)
|
||||
return status, json.loads(payload)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_media_upload_ok(gw):
|
||||
port = http_port(gw)
|
||||
status, body = await asyncio.to_thread(_upload, port, PNG_1X1)
|
||||
assert status == 201, body
|
||||
assert body["type"] == "media.upload.ack"
|
||||
assert body["payload"]["ok"] is True
|
||||
assert body["payload"]["media_ref"] == "mu_http1"
|
||||
entry = gw._media.get_inbound("mu_http1")
|
||||
assert entry is not None
|
||||
assert entry.kind == "image"
|
||||
assert entry.size == len(PNG_1X1)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_media_upload_sha_mismatch(gw):
|
||||
port = http_port(gw)
|
||||
status, body = await asyncio.to_thread(
|
||||
_upload, port, PNG_1X1, media_ref="mu_badsha", sha256="0" * 64
|
||||
)
|
||||
assert status == 500, body # internal: digest mismatch
|
||||
assert body["type"] == "error"
|
||||
assert body["payload"]["code"] == "internal"
|
||||
assert gw._media.get_inbound("mu_badsha") is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_media_upload_oversize_413(gw):
|
||||
port = http_port(gw)
|
||||
oversize = b"x" * (gw.max_upload_bytes + 1)
|
||||
status, body = await asyncio.to_thread(_upload, port, oversize, media_ref="mu_big")
|
||||
assert status == 413, body
|
||||
assert body["payload"]["code"] == "media_too_large"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_media_upload_missing_ref_400(gw):
|
||||
port = http_port(gw)
|
||||
status, payload = await asyncio.to_thread(
|
||||
_request,
|
||||
port,
|
||||
"POST",
|
||||
"/v1/media",
|
||||
body=PNG_1X1,
|
||||
content_type="image/png",
|
||||
extra_headers={"X-Iris-Media-Kind": "image"},
|
||||
)
|
||||
body = json.loads(payload)
|
||||
assert status == 400
|
||||
assert body["payload"]["code"] == "unsupported"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_media_upload_bad_kind_400(gw):
|
||||
port = http_port(gw)
|
||||
status, body = await asyncio.to_thread(_upload, port, PNG_1X1, kind="hologram")
|
||||
assert status == 400
|
||||
assert body["payload"]["code"] == "unsupported"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_media_upload_auth_401(gw):
|
||||
port = http_port(gw)
|
||||
status, _ = await asyncio.to_thread(_upload, port, PNG_1X1, token="wrong-token")
|
||||
assert status == 401
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_media_upload_liar_reclassified(gw):
|
||||
"""Lies about being a PNG: magic-byte re-sniff keeps it out of the image
|
||||
cache (lands as a document) — same contract as the WS path."""
|
||||
port = http_port(gw)
|
||||
payload = b"<html>not an image</html>"
|
||||
status, body = await asyncio.to_thread(
|
||||
_upload, port, payload, media_ref="mu_liar", filename="liar.html"
|
||||
)
|
||||
assert status == 201, body
|
||||
entry = gw._media.get_inbound("mu_liar")
|
||||
assert entry is not None
|
||||
assert entry.kind == "document"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_media_pull_ok(gw):
|
||||
from gateway.platforms.base import get_image_cache_dir
|
||||
|
||||
img = get_image_cache_dir() / "http_pull_test.png"
|
||||
img.write_bytes(PNG_1X1)
|
||||
entry = gw._media.register_outbound(
|
||||
str(img), "image", "image/png", "http_pull_test.png", len(PNG_1X1)
|
||||
)
|
||||
port = http_port(gw)
|
||||
status, payload = await asyncio.to_thread(_request, port, "GET", f"/v1/media/{entry.media_id}")
|
||||
assert status == 200
|
||||
assert payload == PNG_1X1
|
||||
conn = HTTPConnection("127.0.0.1", port, timeout=10)
|
||||
conn.request(
|
||||
"GET",
|
||||
f"/v1/media/{entry.media_id}",
|
||||
headers={"Authorization": f"Bearer {TOKEN}", "X-Iris-Device": DEVICE_ID},
|
||||
)
|
||||
resp = conn.getresponse()
|
||||
resp.read()
|
||||
assert resp.getheader("Content-Type") == "image/png"
|
||||
assert resp.getheader("Content-Length") == str(len(PNG_1X1))
|
||||
conn.close()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_media_pull_unknown_404(gw):
|
||||
port = http_port(gw)
|
||||
status, payload = await asyncio.to_thread(_request, port, "GET", "/v1/media/md_nope")
|
||||
body = json.loads(payload)
|
||||
assert status == 404
|
||||
assert body["payload"]["code"] == "not_found"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_media_pull_denied_path_404(gw):
|
||||
"""Known id, but the path fails delivery validation (denylist) — same
|
||||
re-check at pull time as the WS path."""
|
||||
entry = gw._media.register_outbound("/etc/passwd", "document", "text/plain", "passwd", 100)
|
||||
port = http_port(gw)
|
||||
status, payload = await asyncio.to_thread(_request, port, "GET", f"/v1/media/{entry.media_id}")
|
||||
body = json.loads(payload)
|
||||
assert status == 404
|
||||
assert body["payload"]["code"] == "not_found"
|
||||
|
||||
|
||||
# ── Helpers ─────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
async def _park_frame(adapter, text: str) -> int:
|
||||
"""Emit one message frame through ``_broadcast_or_log`` (the real
|
||||
outbound path); returns the outbox cursor."""
|
||||
plugin = _load_plugin()
|
||||
frame = plugin.protocol.message(
|
||||
chat_id=CHAT_ID,
|
||||
message_id=f"m_{abs(hash(text)) % 10**8:08x}",
|
||||
role="assistant",
|
||||
text=text,
|
||||
)
|
||||
await adapter._broadcast_or_log(CHAT_ID, frame)
|
||||
return adapter._outbox.latest_cursor()
|
||||
+389
-514
File diff suppressed because it is too large.
Load diff
@@ -1,444 +0,0 @@
|
||||
"""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 contextlib
|
||||
import logging
|
||||
import ssl
|
||||
import time
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
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
|
||||
|
||||
# Max length of a client-supplied device_id.
|
||||
MAX_DEVICE_ID_LEN = 128
|
||||
|
||||
|
||||
async def dispatch_frame(adapter: Any, frame: protocol.Frame, device_id: str) -> None:
|
||||
"""Shared inbound frame dispatch for the WS and HTTP transports
|
||||
(docs/19 §19.4). Transport-specific frames (``ping``/``pong``, binary
|
||||
media chunks) are handled by their own server before this is called;
|
||||
unknown types are ignored (forward-compat)."""
|
||||
if frame.type == protocol.TYPE_MESSAGE_SEND:
|
||||
await adapter.on_message_send(frame, device_id)
|
||||
elif frame.type == protocol.TYPE_CHANNEL_CREATE:
|
||||
await adapter.on_channel_create(frame, device_id)
|
||||
elif frame.type == protocol.TYPE_CHANNEL_RENAME:
|
||||
await adapter.on_channel_rename(frame, device_id)
|
||||
elif frame.type == protocol.TYPE_CHANNEL_SET_DEFAULT:
|
||||
await adapter.on_channel_set_default(frame, device_id)
|
||||
elif frame.type == protocol.TYPE_CHANNEL_FAVORITE:
|
||||
await adapter.on_channel_favorite(frame, device_id)
|
||||
elif frame.type == protocol.TYPE_CHANNEL_ICON:
|
||||
await adapter.on_channel_icon(frame, device_id)
|
||||
elif frame.type == protocol.TYPE_CHANNEL_SET_AUTOMATION:
|
||||
await adapter.on_channel_set_automation(frame, device_id)
|
||||
elif frame.type == protocol.TYPE_CHANNEL_DELETE:
|
||||
await adapter.on_channel_delete(frame, device_id)
|
||||
elif frame.type == protocol.TYPE_CHANNEL_LIST:
|
||||
await adapter.on_channel_list(frame, device_id)
|
||||
elif frame.type == protocol.TYPE_COMMANDS_CATALOG:
|
||||
await adapter.on_commands_catalog(frame, device_id)
|
||||
elif frame.type == protocol.TYPE_SEARCH:
|
||||
await adapter.on_search(frame, device_id)
|
||||
elif frame.type == protocol.TYPE_SYNC:
|
||||
await adapter.on_sync(frame, device_id)
|
||||
elif frame.type == protocol.TYPE_HISTORY:
|
||||
await adapter.on_history(frame, device_id)
|
||||
elif frame.type == protocol.TYPE_MESSAGE_DELETE:
|
||||
await adapter.on_message_delete(frame, device_id)
|
||||
elif frame.type == protocol.TYPE_MEDIA_UPLOAD_START:
|
||||
await adapter.on_media_upload_start(frame, device_id)
|
||||
elif frame.type == protocol.TYPE_MEDIA_UPLOAD_END:
|
||||
await adapter.on_media_upload_end(frame, device_id)
|
||||
elif frame.type == protocol.TYPE_MEDIA_PULL:
|
||||
await adapter.on_media_pull(frame, device_id)
|
||||
elif frame.type == protocol.TYPE_FCM_REGISTER:
|
||||
await adapter.on_fcm_register(frame, device_id)
|
||||
# Unknown types are ignored (forward-compat).
|
||||
|
||||
|
||||
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: str | None = None
|
||||
ntfy_topic: str | None = 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: Any | None = 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: ssl.SSLContext | None = 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()
|
||||
# Best-effort: the server is already closing; a failure here is
|
||||
# not actionable (nothing left to clean up besides the registry).
|
||||
with contextlib.suppress(Exception):
|
||||
await self._server.wait_closed()
|
||||
self._server = None
|
||||
for conn in list(self._connections.values()):
|
||||
# Best-effort: a socket that is already gone needs no handling.
|
||||
with contextlib.suppress(Exception):
|
||||
await conn.ws.close(code=CLOSE_SHUTDOWN, reason="gateway shutting down")
|
||||
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) -> DeviceConnection | None:
|
||||
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()):
|
||||
# Best-effort: a dead or stalled socket is skipped (it is
|
||||
# deregistered on its own close); one slow peer must not starve
|
||||
# the rest of the broadcast.
|
||||
with contextlib.suppress(Exception):
|
||||
await asyncio.wait_for(conn.ws.send(data), timeout=SEND_TIMEOUT_S)
|
||||
sent += 1
|
||||
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, ConnectionClosed) as e:
|
||||
if isinstance(e, asyncio.TimeoutError):
|
||||
logger.warning("android: dropping socket with no hello (timeout)")
|
||||
await self._close_quiet(ws, 1000, "no hello")
|
||||
# A peer that vanished before hello needs no further handling.
|
||||
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) > MAX_DEVICE_ID_LEN:
|
||||
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.
|
||||
# Best-effort close of the superseded socket.
|
||||
with contextlib.suppress(Exception):
|
||||
await old.ws.close(code=CLOSE_REPLACED, reason="replaced by newer connection")
|
||||
|
||||
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 Exception as e:
|
||||
# A clean disconnect (ConnectionClosed) is the normal path and is
|
||||
# not worth a warning; anything else is unexpected.
|
||||
if not isinstance(e, ConnectionClosed):
|
||||
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))
|
||||
return True
|
||||
await dispatch_frame(self._adapter, frame, device_id)
|
||||
return True
|
||||
|
||||
# ── Helpers ───────────────────────────────────────────────────────────
|
||||
|
||||
def _connection_for(self, ws: ServerConnection) -> DeviceConnection | None:
|
||||
"""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:
|
||||
# "Quiet" by contract: the caller does not care whether the peer was
|
||||
# still there (e.g. an error frame right before the close).
|
||||
with contextlib.suppress(Exception):
|
||||
await ws.send(frame.to_json())
|
||||
|
||||
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:
|
||||
# "Quiet" by contract: closing an already-closed socket is a no-op.
|
||||
with contextlib.suppress(Exception):
|
||||
await ws.close(code=code, reason=reason)
|
||||
Reference in new issue
Block a user