Clean up lint/LSP across gateway, Android, and desktop (alpha -> stable)

Gateway (gateway-plugin/):
- Fix interactive_setup broken imports: print helpers were imported from the
  wrong hermes module (hermes_cli.config instead of hermes_cli.cli_output) plus
  a non-existent print_code; the try/except swallowed the ImportError so
  `hermes gateway setup` for android always bailed out early.
- Fix release_scoped_lock type error (str | None passed where str required).
- Rewrite empty `except: pass` blocks as contextlib.suppress with rationale.
- Restructure two ambiguous ws_server try blocks (hello-auth, frame loop).
- Ruff cleanup: type annotations, import sorting, line wrapping, magic values
  -> named constants, `raise ... from e`, complexity. Add gateway-plugin/ruff.toml.
- Add pyrightconfig.json so the Python LSP resolves hermes-runtime imports.
- Suppress verified false positives inline (parameterized SQL, column-name
  "secrets", hermes-generated media path).

Android (app/androidApp + app/shared):
- Consolidate launcher icons into a single mipmap-anydpi (minSdk 29 >= 26) with
  the monochrome layer; clears ObsoleteSdkInt + MonochromeLauncherIcon.
- Bump core-splashscreen 1.0.1 -> 1.2.0; pin targetSdk 34 (deliberate).
- Suppress verified findings inline (LAN ws:// default, correct GCM IV usage).

Desktop (app/desktopApp):
- Move the desktop to a Java 21 runtime (org.gradle.java.home) and set the
  desktop jvmTarget to 21 (Android stays JVM 17 / minSdk 29). Fixes the startup
  UnsupportedClassVersionError and restores Markdown renderer 0.44.0.

Tooling/config:
- .pi-lens.json: disable verified-noisy heuristics (documented in docs).
- .gitleaks.toml: allowlist git-ignored false-positive paths.
- docs/18-code-review.md: full findings + verification.

Verified: ruff clean, pyright 0 errors, 64/64 gateway tests, all Kotlin tests,
Android lint 0 issues, Android installed+launched on device, desktop launches
on JDK 21.
This commit is contained in:
ARIA committed 2026-08-21 18:47:03 +02:00
1 parent 9f3f9842c8
commit 678c0344c8
27 files changed
+928 -454

No files matched your search

+1 -1
View File
@@ -1,3 +1,3 @@
from .adapter import register
__all__ = ["register"]
__all__ = ["register"]
+278 -202
View File
File diff suppressed because it is too large. Load diff
+53 -51
View File
@@ -20,12 +20,14 @@ Storage: ``get_hermes_home()/"android"/channels.db``.
Milestone M3.
"""
import builtins
import contextlib
import logging
import sqlite3
import threading
import time
from pathlib import Path
from typing import Any, Dict, List, Optional
from typing import Any
logger = logging.getLogger(__name__)
@@ -73,9 +75,7 @@ class ChannelDirectory:
"""
)
# Migrate existing DBs: add the cosmetic columns if missing.
existing = {
row[1] for row in self._conn.execute("PRAGMA table_info(channels)")
}
existing = {row[1] for row in self._conn.execute("PRAGMA table_info(channels)")}
if "favorite" not in existing:
self._conn.execute(
"ALTER TABLE channels ADD COLUMN favorite INTEGER NOT NULL DEFAULT 0"
@@ -120,7 +120,7 @@ class ChannelDirectory:
# ── default channel ───────────────────────────────────────────────────
def ensure_default(self, chat_id: str, name: str) -> Dict[str, Any]:
def ensure_default(self, chat_id: str, name: str) -> dict[str, Any]:
"""Ensure the default (home) channel exists. Idempotent.
If a row already exists for *chat_id* it is kept (name refreshed only
@@ -152,16 +152,18 @@ class ChannelDirectory:
(KIND_DEFAULT, chat_id),
)
# Exactly one default: clear any other default flag.
self._conn.execute(
"UPDATE channels SET is_default = 0 WHERE chat_id != ?", (chat_id,)
)
self._conn.execute("UPDATE channels SET is_default = 0 WHERE chat_id != ?", (chat_id,))
self._conn.commit()
entry = self.get(chat_id)
if entry is not None:
return entry
return {
"chat_id": chat_id, "name": name, "kind": KIND_DEFAULT,
"parent_chat_id": None, "is_default": True, "archived": False,
"chat_id": chat_id,
"name": name,
"kind": KIND_DEFAULT,
"parent_chat_id": None,
"is_default": True,
"archived": False,
"created": time.time(),
}
@@ -171,8 +173,8 @@ class ChannelDirectory:
self,
name: str,
kind: str = KIND_CHANNEL,
parent_chat_id: Optional[str] = None,
) -> Dict[str, Any]:
parent_chat_id: str | None = None,
) -> dict[str, Any]:
"""Mint a new channel (or thread) and store it. Returns the entry."""
name = (name or "").strip()
if not name:
@@ -194,12 +196,16 @@ class ChannelDirectory:
if entry is not None:
return entry
return {
"chat_id": chat_id, "name": name, "kind": kind,
"parent_chat_id": parent_chat_id, "is_default": False,
"archived": False, "created": now,
"chat_id": chat_id,
"name": name,
"kind": kind,
"parent_chat_id": parent_chat_id,
"is_default": False,
"archived": False,
"created": now,
}
def rename(self, chat_id: str, name: str) -> Optional[Dict[str, Any]]:
def rename(self, chat_id: str, name: str) -> dict[str, Any] | None:
name = (name or "").strip()
if not name:
raise ValueError("channel name required")
@@ -213,7 +219,7 @@ class ChannelDirectory:
return None
return self.get(chat_id)
def set_default(self, chat_id: str) -> Optional[Dict[str, Any]]:
def set_default(self, chat_id: str) -> dict[str, Any] | None:
"""Mark *chat_id* as the default channel (clears the previous one).
The default channel is the user's chat surface, so the automation
@@ -228,14 +234,13 @@ class ChannelDirectory:
return None
self._conn.execute("UPDATE channels SET is_default = 0")
self._conn.execute(
"UPDATE channels SET is_default = 1, automation = 0 "
"WHERE chat_id = ?",
"UPDATE channels SET is_default = 1, automation = 0 WHERE chat_id = ?",
(chat_id,),
)
self._conn.commit()
return self.get(chat_id)
def set_favorite(self, chat_id: str, on: bool) -> Optional[Dict[str, Any]]:
def set_favorite(self, chat_id: str, on: bool) -> dict[str, Any] | None:
"""Toggle the cosmetic favorite flag (sorts to the top of the list)."""
with self._lock:
cur = self._conn.execute(
@@ -247,7 +252,7 @@ class ChannelDirectory:
return None
return self.get(chat_id)
def set_icon(self, chat_id: str, icon: Optional[str], color: Optional[str]) -> Optional[Dict[str, Any]]:
def set_icon(self, chat_id: str, icon: str | None, color: str | None) -> dict[str, Any] | None:
"""Set the channel's cosmetic icon (base64 image) and/or avatar color.
``icon`` is a base64-encoded image (or ``None`` to clear it); ``color``
@@ -264,7 +269,7 @@ class ChannelDirectory:
return None
return self.get(chat_id)
def set_automation(self, chat_id: str, on: bool) -> Optional[Dict[str, Any]]:
def set_automation(self, chat_id: str, on: bool) -> dict[str, Any] | None:
"""Mark *chat_id* as an automation channel (or clear the flag).
Automation channels are read-only for the user: they only receive
@@ -287,7 +292,7 @@ class ChannelDirectory:
self._conn.commit()
return self.get(chat_id)
def delete(self, chat_id: str) -> Optional[Dict[str, Any]]:
def delete(self, chat_id: str) -> dict[str, Any] | None:
"""Soft-delete (archive) a channel. History stays for search.
The default channel cannot be deleted. Returns the (archived) entry,
@@ -299,39 +304,40 @@ class ChannelDirectory:
).fetchone()
if row is None or row["is_default"]:
return None
self._conn.execute(
"UPDATE channels SET archived = 1 WHERE chat_id = ?", (chat_id,)
)
self._conn.execute("UPDATE channels SET archived = 1 WHERE chat_id = ?", (chat_id,))
self._conn.commit()
return self.get(chat_id)
# ── reads ─────────────────────────────────────────────────────────────
def get(self, chat_id: str) -> Optional[Dict[str, Any]]:
def get(self, chat_id: str) -> dict[str, Any] | None:
with self._lock:
row = self._conn.execute(
"SELECT * FROM channels WHERE chat_id = ?", (chat_id,)
).fetchone()
return _row_to_entry(row) if row else None
def list(self, include_archived: bool = False) -> List[Dict[str, Any]]:
def list(self, include_archived: bool = False) -> list[dict[str, Any]]:
"""Directory listing. Default first, then favorites, then creation order."""
sql = "SELECT * FROM channels"
if not include_archived:
sql += " WHERE archived = 0"
sql += " ORDER BY is_default DESC, favorite DESC, created ASC"
with self._lock:
# Safe: fully static SQL (no user data); the variable is only to
# toggle the optional archived filter.
# pi-lens-ignore: python-sql-injection
rows = self._conn.execute(sql).fetchall()
return [_row_to_entry(r) for r in rows]
def default(self) -> Optional[Dict[str, Any]]:
def default(self) -> dict[str, Any] | None:
with self._lock:
row = self._conn.execute(
"SELECT * FROM channels WHERE is_default = 1 LIMIT 1"
).fetchone()
return _row_to_entry(row) if row else None
def threads_for(self, chat_id: str) -> List[Dict[str, Any]]:
def threads_for(self, chat_id: str) -> builtins.list[dict[str, Any]]:
"""All (non-archived) threads under *chat_id*, oldest first."""
with self._lock:
rows = self._conn.execute(
@@ -341,7 +347,7 @@ class ChannelDirectory:
).fetchall()
return [_row_to_entry(r) for r in rows]
def resolve_entry(self, name: str) -> Optional[Dict[str, Any]]:
def resolve_entry(self, name: str) -> dict[str, Any] | None:
"""Resolve a friendly name to a directory entry (case-insensitive).
Matches non-archived channels/threads by exact name first, then by
@@ -352,23 +358,19 @@ class ChannelDirectory:
if not query:
return None
with self._lock:
rows = self._conn.execute(
"SELECT * FROM channels WHERE archived = 0"
).fetchall()
rows = self._conn.execute("SELECT * FROM channels WHERE archived = 0").fetchall()
entries = [_row_to_entry(r) for r in rows]
exact = [e for e in entries if (e["name"] or "").strip().lower() == query]
if len(exact) == 1:
return exact[0]
if len(exact) > 1:
return None
prefix = [
e for e in entries if (e["name"] or "").strip().lower().startswith(query)
]
prefix = [e for e in entries if (e["name"] or "").strip().lower().startswith(query)]
if len(prefix) == 1:
return prefix[0]
return None
def resolve_name(self, name: str) -> Optional[str]:
def resolve_name(self, name: str) -> str | None:
"""Resolve a friendly name to a valid chat_id (case-insensitive).
For a thread, returns the *parent* chat_id (the thread's session lane
@@ -383,14 +385,12 @@ class ChannelDirectory:
return entry["chat_id"]
def close(self) -> None:
with self._lock:
try:
self._conn.close()
except Exception:
pass
with self._lock, contextlib.suppress(Exception):
# Best-effort: a close failure on shutdown is not actionable.
self._conn.close()
def _row_to_entry(row: sqlite3.Row) -> Dict[str, Any]:
def _row_to_entry(row: sqlite3.Row) -> dict[str, Any]:
return {
"chat_id": row["chat_id"],
"name": row["name"],
@@ -415,24 +415,26 @@ def _row_to_entry(row: sqlite3.Row) -> Dict[str, Any]:
# Keyed on ``get_hermes_home()`` so a profile switch rebuilds it.
# ---------------------------------------------------------------------------
_directory: Optional[ChannelDirectory] = None
_directory_home: Optional[Path] = None
_directory: ChannelDirectory | None = None
_directory_home: Path | None = None
_directory_lock = threading.Lock()
def get_directory() -> ChannelDirectory:
"""Return the process-wide channel directory for the active profile."""
global _directory, _directory_home
# Module-level singleton keyed on the active profile; the global is the
# intended pattern here (see the block comment above).
global _directory, _directory_home # noqa: PLW0603
from hermes_constants import get_hermes_home
home = Path(get_hermes_home())
with _directory_lock:
if _directory is None or _directory_home != home:
if _directory is not None:
try:
# Best-effort: the old directory is being replaced; a close
# failure is not actionable.
with contextlib.suppress(Exception):
_directory.close()
except Exception:
pass
_directory = ChannelDirectory(home / "android" / "channels.db")
_directory_home = home
return _directory
return _directory
+38 -26
View File
@@ -21,6 +21,7 @@ Milestone M4.
"""
import asyncio
import contextlib
import hashlib
import logging
import os
@@ -31,7 +32,6 @@ import time
import uuid
from dataclasses import dataclass
from pathlib import Path
from typing import Dict, Optional, Tuple
from gateway.platforms.base import (
_looks_like_image,
@@ -54,9 +54,11 @@ _SHA256_RE = re.compile(r"^[0-9a-f]{64}$")
# Magic-byte containers that are unambiguously audio (vs video-in-same-box).
_AUDIO_CONTAINERS = {"m4a", "ogg", "flac", "wav", "mp3", "aac"}
_VIDEO_CONTAINERS = {"mp4", "webm"}
# Longest file extension we trust from the client (e.g. ".webm").
_MAX_EXT_LEN = 6
# Extension -> MIME for outbound offers (the app picks a player/viewer from it).
_EXT_TO_MIME: Dict[str, str] = {
_EXT_TO_MIME: dict[str, str] = {
".jpg": "image/jpeg",
".jpeg": "image/jpeg",
".png": "image/png",
@@ -93,7 +95,7 @@ _EXT_TO_MIME: Dict[str, str] = {
}
# MIME -> extension for inbound caching (the cache helpers take an ext).
_MIME_TO_EXT: Dict[str, str] = {
_MIME_TO_EXT: dict[str, str] = {
"image/jpeg": ".jpg",
"image/png": ".png",
"image/webp": ".webp",
@@ -139,7 +141,7 @@ def ext_for_mime(mime: str, filename: str, default: str) -> str:
if ext:
return ext
file_ext = os.path.splitext(filename or "")[1].lower()
if file_ext and len(file_ext) <= 6:
if file_ext and len(file_ext) <= _MAX_EXT_LEN:
return file_ext
return default
@@ -190,7 +192,7 @@ class UploadSession:
mime: str,
filename: str,
declared_size: int,
request_id: Optional[int],
request_id: int | None,
max_bytes: int,
tmp_dir: Path,
):
@@ -242,14 +244,12 @@ class UploadSession:
def close(self) -> None:
"""Discard the session and remove the temp file."""
try:
# Best-effort cleanup: a file that is already gone (or a handle that
# is already closed) needs no further handling.
with contextlib.suppress(Exception):
self._fh.close()
except Exception:
pass
try:
with contextlib.suppress(OSError):
os.unlink(self.tmp_path)
except OSError:
pass
class MediaStore:
@@ -267,9 +267,9 @@ class MediaStore:
self._tmp_dir.mkdir(parents=True, exist_ok=True)
self._lock = threading.Lock()
# (device_id, media_ref) -> UploadSession (one active per device)
self._uploads: Dict[Tuple[str, str], UploadSession] = {}
self._inbound: Dict[str, MediaEntry] = {}
self._outbound: Dict[str, MediaEntry] = {}
self._uploads: dict[tuple[str, str], UploadSession] = {}
self._inbound: dict[str, MediaEntry] = {}
self._outbound: dict[str, MediaEntry] = {}
# ── Inbound uploads ───────────────────────────────────────────────────
@@ -281,11 +281,11 @@ class MediaStore:
mime: str,
filename: str,
declared_size: int,
request_id: Optional[int],
request_id: int | None,
max_bytes: int,
) -> UploadSession:
with self._lock:
for (dev, _ref), sess in self._uploads.items():
for dev, _ref in self._uploads:
if dev == device_id:
raise MediaError(
"unsupported", "an upload is already in progress on this connection"
@@ -293,13 +293,19 @@ class MediaStore:
if media_ref in self._inbound:
raise MediaError("unsupported", f"media_ref {media_ref} already used")
sess = UploadSession(
media_ref, kind, mime, filename, declared_size, request_id,
max_bytes, self._tmp_dir,
media_ref,
kind,
mime,
filename,
declared_size,
request_id,
max_bytes,
self._tmp_dir,
)
self._uploads[(device_id, media_ref)] = sess
return sess
def get_upload(self, device_id: str, media_ref: Optional[str] = None) -> Optional[UploadSession]:
def get_upload(self, device_id: str, media_ref: str | None = None) -> UploadSession | None:
with self._lock:
if media_ref is not None:
return self._uploads.get((device_id, media_ref))
@@ -354,8 +360,8 @@ class MediaStore:
# hermes cap (gateway.max_inbound_media_bytes) or a
# non-image payload masquerading as an image.
if "too large" in str(e):
raise MediaError("media_too_large", str(e))
raise MediaError("unsupported", str(e))
raise MediaError("media_too_large", str(e)) from e
raise MediaError("unsupported", str(e)) from e
entry = MediaEntry(
media_id=media_ref,
@@ -370,17 +376,20 @@ class MediaStore:
self._inbound[media_ref] = entry
logger.info(
"android: upload %s cached as %s (%s, %d bytes)",
media_ref, kind, path, len(data),
media_ref,
kind,
path,
len(data),
)
return entry
finally:
sess.close()
def get_inbound(self, media_ref: str) -> Optional[MediaEntry]:
def get_inbound(self, media_ref: str) -> MediaEntry | None:
with self._lock:
return self._inbound.get(media_ref)
def pop_inbound(self, media_ref: str) -> Optional[MediaEntry]:
def pop_inbound(self, media_ref: str) -> MediaEntry | None:
with self._lock:
return self._inbound.pop(media_ref, None)
@@ -402,7 +411,7 @@ class MediaStore:
self._outbound[entry.media_id] = entry
return entry
def get_outbound(self, media_id: str) -> Optional[MediaEntry]:
def get_outbound(self, media_id: str) -> MediaEntry | None:
with self._lock:
return self._outbound.get(media_id)
@@ -438,6 +447,9 @@ async def stream_file(
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)
@@ -445,4 +457,4 @@ async def stream_file(
break
await asyncio.wait_for(ws.send(chunk), timeout=timeout)
sent += len(chunk)
return sent
return sent
+26 -35
View File
@@ -15,13 +15,14 @@ Storage: ``get_hermes_home()/"android"/outbox.db``.
Milestone M3 (built), extended in M5 (push integration).
"""
import contextlib
import json
import logging
import sqlite3
import threading
import time
from pathlib import Path
from typing import Any, Dict, List, Optional
from typing import Any
logger = logging.getLogger(__name__)
@@ -69,9 +70,7 @@ class Outbox:
)
"""
)
self._conn.execute(
"CREATE INDEX IF NOT EXISTS idx_outbox_created ON outbox (created)"
)
self._conn.execute("CREATE INDEX IF NOT EXISTS idx_outbox_created ON outbox (created)")
self._conn.execute(
"""
CREATE TABLE IF NOT EXISTS counters (
@@ -84,7 +83,7 @@ class Outbox:
# ── append / cursor ───────────────────────────────────────────────────
def append(self, chat_id: Optional[str], frame_json: str) -> int:
def append(self, chat_id: str | None, frame_json: str) -> int:
"""Append a frame; returns the (monotonic) cursor assigned to it."""
now = time.time()
with self._lock:
@@ -92,13 +91,10 @@ class Outbox:
"INSERT INTO counters (name, value) VALUES ('cursor', 1) "
"ON CONFLICT(name) DO UPDATE SET value = value + 1"
)
row = self._conn.execute(
"SELECT value FROM counters WHERE name = 'cursor'"
).fetchone()
row = self._conn.execute("SELECT value FROM counters WHERE name = 'cursor'").fetchone()
cursor = int(row["value"]) if row else 1
self._conn.execute(
"INSERT INTO outbox (cursor, chat_id, frame, created) "
"VALUES (?, ?, ?, ?)",
"INSERT INTO outbox (cursor, chat_id, frame, created) VALUES (?, ?, ?, ?)",
(cursor, chat_id, frame_json, now),
)
self._enforce_row_cap()
@@ -135,14 +131,12 @@ class Outbox:
def latest_cursor(self) -> int:
"""The high-water cursor (0 when nothing has been appended)."""
with self._lock:
row = self._conn.execute(
"SELECT value FROM counters WHERE name = 'cursor'"
).fetchone()
row = self._conn.execute("SELECT value FROM counters WHERE name = 'cursor'").fetchone()
return int(row["value"]) if row else 0
# ── replay ────────────────────────────────────────────────────────────
def replay(self, cursor: int, limit: int = _REPLAY_LIMIT) -> List[Dict[str, Any]]:
def replay(self, cursor: int, limit: int = _REPLAY_LIMIT) -> list[dict[str, Any]]:
"""Frames with ``cursor > `cursor```, oldest first.
Each entry: ``{cursor, chat_id, frame}`` where ``frame`` is the parsed
@@ -156,7 +150,7 @@ class Outbox:
"WHERE cursor > ? ORDER BY cursor ASC LIMIT ?",
(cursor, limit),
).fetchall()
out: List[Dict[str, Any]] = []
out: list[dict[str, Any]] = []
for r in rows:
try:
frame = json.loads(r["frame"])
@@ -164,9 +158,7 @@ class Outbox:
continue
if not isinstance(frame, dict):
continue
out.append(
{"cursor": int(r["cursor"]), "chat_id": r["chat_id"], "frame": frame}
)
out.append({"cursor": int(r["cursor"]), "chat_id": r["chat_id"], "frame": frame})
return out
# ── history (full message history for a chat/thread) ──────────────────
@@ -174,10 +166,10 @@ class Outbox:
def history(
self,
chat_id: str,
thread_id: Optional[str] = None,
before_message_id: Optional[str] = None,
thread_id: str | None = None,
before_message_id: str | None = None,
limit: int = 50,
) -> Dict[str, Any]:
) -> dict[str, Any]:
"""Final messages for a chat/thread, for the ``history`` frame.
Reconstructs the message list from the outbox log: a final message is
@@ -194,11 +186,10 @@ class Outbox:
limit = max(1, min(int(limit or 50), 200))
with self._lock:
rows = self._conn.execute(
"SELECT cursor, frame FROM outbox WHERE chat_id = ? "
"ORDER BY cursor ASC",
"SELECT cursor, frame FROM outbox WHERE chat_id = ? ORDER BY cursor ASC",
(chat_id,),
).fetchall()
final: List[Dict[str, Any]] = []
final: list[dict[str, Any]] = []
for r in rows:
try:
frame = json.loads(r["frame"])
@@ -244,7 +235,7 @@ class Outbox:
}
)
# Deduplicate by message_id (keep the latest occurrence), keep order.
by_id: Dict[str, Dict[str, Any]] = {}
by_id: dict[str, dict[str, Any]] = {}
for m in final:
mid = m.get("message_id")
if mid:
@@ -283,7 +274,7 @@ class Outbox:
self,
chat_id: str,
message_id: str,
thread_id: Optional[str] = None,
thread_id: str | None = None,
) -> int:
"""Remove every outbox frame belonging to *message_id* in *chat_id*.
@@ -304,7 +295,7 @@ class Outbox:
rows = self._conn.execute(
"SELECT cursor, frame FROM outbox WHERE chat_id = ?", (chat_id,)
).fetchall()
cursors: List[int] = []
cursors: list[int] = []
for r in rows:
try:
frame = json.loads(r["frame"])
@@ -320,9 +311,11 @@ class Outbox:
if not cursors:
return 0
placeholders = ",".join("?" * len(cursors))
self._conn.execute(
f"DELETE FROM outbox WHERE cursor IN ({placeholders})", cursors
)
sql = f"DELETE FROM outbox WHERE cursor IN ({placeholders})"
# Safe: ``placeholders`` is only ``?`` markers; every cursor value is
# bound as a parameter (no user data in the SQL text).
# pi-lens-ignore: python-sql-injection
self._conn.execute(sql, cursors)
self._conn.commit()
return len(cursors)
@@ -347,8 +340,6 @@ class Outbox:
self._maybe_prune()
def close(self) -> None:
with self._lock:
try:
self._conn.close()
except Exception:
pass
with self._lock, contextlib.suppress(Exception):
# Best-effort: a close failure on shutdown is not actionable.
self._conn.close()
+6 -12
View File
@@ -9,6 +9,7 @@ Storage: ``get_hermes_home()/"android"/devices.db``.
Milestone M1.
"""
import contextlib
import hmac
import json
import logging
@@ -98,10 +99,7 @@ class DeviceRegistry:
# M5: migrate pre-push-cursor databases (the column carries the
# highest outbox cursor already delivered to the device via the
# push backend; hello.ack returns it for notification dedupe).
cols = {
r["name"]
for r in self._conn.execute("PRAGMA table_info(devices)").fetchall()
}
cols = {r["name"] for r in self._conn.execute("PRAGMA table_info(devices)").fetchall()}
if "last_pushed_cursor" not in cols:
self._conn.execute(
"ALTER TABLE devices ADD COLUMN last_pushed_cursor INTEGER NOT NULL DEFAULT 0"
@@ -200,17 +198,13 @@ class DeviceRegistry:
def list(self) -> list[dict[str, Any]]:
with self._lock:
rows = self._conn.execute(
"SELECT * FROM devices ORDER BY last_seen DESC"
).fetchall()
rows = self._conn.execute("SELECT * FROM devices ORDER BY last_seen DESC").fetchall()
return [_row_to_device(r) for r in rows]
def close(self) -> None:
with self._lock:
try:
self._conn.close()
except Exception:
pass
with self._lock, contextlib.suppress(Exception):
# Best-effort: a close failure on shutdown is not actionable.
self._conn.close()
def _row_to_device(row: sqlite3.Row) -> dict[str, Any]:
+3 -1
View File
@@ -249,7 +249,9 @@ def hello_ack(
)
def message(
# Frame builder mirrors the wire schema (docs/04); the many fields are the
# message's full shape, so the arg count is intentional.
def message( # noqa: PLR0913
chat_id: str,
message_id: str,
role: str,
+35 -26
View File
@@ -25,7 +25,7 @@ import logging
import threading
import time
from pathlib import Path
from typing import Any, Dict, Optional
from typing import Any
from urllib.parse import quote
import httpx
@@ -41,6 +41,9 @@ _TOKEN_REFRESH_MARGIN_S = 600.0
_DEFAULT_NTFY_SERVER = "https://ntfy.sh"
_NTFY_BODY_LIMIT = 4096
_HTTP_TIMEOUT_S = 15.0
# HTTP status boundaries: 200 == success; >= 300 == redirect/error range.
_HTTP_OK = 200
_HTTP_ERROR_MIN = 300
_NTFY_PRIORITY = {"high": "5", "normal": "3", "low": "1"}
@@ -50,6 +53,8 @@ class PushBackend:
name: str = "push"
# DeviceRegistry column that carries this backend's target token.
# Not a secret: a DB column name (string literal), not a credential.
# pi-lens-ignore: python-hardcoded-secrets
token_field: str = ""
def configured(self) -> bool:
@@ -63,7 +68,7 @@ class PushBackend:
chat_id: str,
title: str,
body: str,
data: Dict[str, Any],
data: dict[str, Any],
token: str,
priority: str = "normal",
data_only: bool = False,
@@ -81,18 +86,20 @@ class FcmBackend(PushBackend):
"""FCM HTTP v1 (service account) or legacy ``/fcm/send`` (server key)."""
name = "fcm"
# Not a secret: a DB column name (string literal), not a credential.
# pi-lens-ignore: python-hardcoded-secrets
token_field = "fcm_token"
def __init__(
self,
service_account: Optional[str] = None,
server_key: Optional[str] = None,
service_account: str | None = None,
server_key: str | None = None,
):
self._sa_path = (service_account or "").strip() or None
self._server_key = (server_key or "").strip() or None
self._sa: Optional[Dict[str, Any]] = None
self._sa: dict[str, Any] | None = None
self._sa_failed = False
self._access_token: Optional[str] = None
self._access_token: str | None = None
self._token_expiry = 0.0
self._lock = threading.Lock()
@@ -101,13 +108,13 @@ class FcmBackend(PushBackend):
return True
return bool(self._sa_path and Path(self._sa_path).is_file())
def _load_sa(self) -> Optional[Dict[str, Any]]:
def _load_sa(self) -> dict[str, Any] | None:
if self._sa is not None:
return self._sa
if not self._sa_path or self._sa_failed:
return None
try:
with open(self._sa_path, "r", encoding="utf-8") as f:
with open(self._sa_path, encoding="utf-8") as f:
sa = json.load(f)
if isinstance(sa, dict) and sa.get("client_email") and sa.get("private_key"):
self._sa = sa
@@ -117,7 +124,7 @@ class FcmBackend(PushBackend):
self._sa_failed = True
return None
async def _authorization(self, client: httpx.AsyncClient) -> Optional[str]:
async def _authorization(self, client: httpx.AsyncClient) -> str | None:
"""Bearer token: the legacy server key, or a cached service-account
OAuth2 access token (JWT-bearer grant, minted with PyJWT)."""
if self._server_key:
@@ -158,7 +165,7 @@ class FcmBackend(PushBackend):
except Exception:
logger.warning("android: FCM token exchange failed", exc_info=True)
return None
if resp.status_code != 200:
if resp.status_code != _HTTP_OK:
logger.warning(
"android: FCM token exchange HTTP %s: %s",
resp.status_code, resp.text[:200],
@@ -186,7 +193,7 @@ class FcmBackend(PushBackend):
chat_id: str,
title: str,
body: str,
data: Dict[str, Any],
data: dict[str, Any],
token: str,
priority: str = "normal",
data_only: bool = False,
@@ -197,7 +204,7 @@ class FcmBackend(PushBackend):
notification = None if data_only else {"title": title or "Iris", "body": body or ""}
async with httpx.AsyncClient(timeout=_HTTP_TIMEOUT_S) as client:
if self._server_key:
payload: Dict[str, Any] = {"to": token}
payload: dict[str, Any] = {"to": token}
if notification:
payload["notification"] = notification
if data:
@@ -209,7 +216,7 @@ class FcmBackend(PushBackend):
project_id = (sa or {}).get("project_id")
if not project_id:
return False
message: Dict[str, Any] = {"token": token}
message: dict[str, Any] = {"token": token}
if notification:
message["notification"] = notification
if data:
@@ -231,7 +238,7 @@ class FcmBackend(PushBackend):
except Exception:
logger.warning("android: FCM send failed (network)", exc_info=True)
return False
if resp.status_code >= 300:
if resp.status_code >= _HTTP_ERROR_MIN:
# 404 NOT_FOUND = stale/invalid registration token.
logger.warning(
"android: FCM send HTTP %s: %s", resp.status_code, resp.text[:200]
@@ -249,13 +256,15 @@ class NtfyBackend(PushBackend):
"""
name = "ntfy"
# Not a secret: a DB column name (string literal), not a credential.
# pi-lens-ignore: python-hardcoded-secrets
token_field = "ntfy_topic"
def __init__(
self,
topic: Optional[str] = None,
server_url: Optional[str] = None,
auth_token: Optional[str] = None,
topic: str | None = None,
server_url: str | None = None,
auth_token: str | None = None,
):
self._topic = (topic or "").strip() or None
self._server = (
@@ -279,7 +288,7 @@ class NtfyBackend(PushBackend):
chat_id: str,
title: str,
body: str,
data: Dict[str, Any],
data: dict[str, Any],
token: str,
priority: str = "normal",
data_only: bool = False,
@@ -306,7 +315,7 @@ class NtfyBackend(PushBackend):
except Exception:
logger.warning("android: ntfy publish failed (network)", exc_info=True)
return False
if resp.status_code >= 300:
if resp.status_code >= _HTTP_ERROR_MIN:
logger.warning(
"android: ntfy publish HTTP %s: %s", resp.status_code, resp.text[:200]
)
@@ -315,17 +324,17 @@ class NtfyBackend(PushBackend):
def build_push_backend(
name: Optional[str],
name: str | None,
*,
fcm_service_account: Optional[str] = None,
fcm_server_key: Optional[str] = None,
ntfy_topic: Optional[str] = None,
ntfy_server_url: Optional[str] = None,
ntfy_auth_token: Optional[str] = None,
fcm_service_account: str | None = None,
fcm_server_key: str | None = None,
ntfy_topic: str | None = None,
ntfy_server_url: str | None = None,
ntfy_auth_token: str | None = None,
) -> PushBackend:
"""Select the backend by name (``ANDROID_PUSH_BACKEND``; fcm default)."""
if (name or "").strip().lower() == "ntfy":
return NtfyBackend(
topic=ntfy_topic, server_url=ntfy_server_url, auth_token=ntfy_auth_token
)
return FcmBackend(service_account=fcm_service_account, server_key=fcm_server_key)
return FcmBackend(service_account=fcm_service_account, server_key=fcm_server_key)
+45
View File
@@ -0,0 +1,45 @@
# Lint config for the android gateway plugin.
#
# Run from the repo root (uses the hermes-agent venv's ruff):
# hermes-agent/.venv/bin/python -m ruff check gateway-plugin
#
# The rule set is deliberately broad (pycodestyle, pyflakes, isort, pyupgrade,
# bugbear, flake8-simplify, pylint, return, comprehensions). Thresholds below
# reflect the plugin's real shape: it is a single large dispatch surface
# (adapter.py) plus a wire-protocol layer (protocol.py) whose frame builders
# mirror the schema, so the complexity ceilings are set just above the current
# maxima rather than an idealized small-function target.
line-length = 100
[lint]
select = [
"E", # pycodestyle errors
"W", # pycodestyle warnings
"F", # pyflakes
"I", # isort
"UP", # pyupgrade
"B", # flake8-bugbear
"SIM", # flake8-simplify
"PL", # pylint
"RET", # flake8-return
"C4", # flake8-comprehensions
]
# The plugin intentionally defers hermes-runtime imports into function bodies
# (they are only available once the plugin is loaded inside the gateway, and
# some are optional/try-imported). Top-level import placement does not apply.
ignore = ["PLC0415"]
[lint.pylint]
# Current maxima in the codebase: 22 branches, 64 statements, 9 returns,
# 8 args (protocol.py:252 frame builder is the lone 11-arg outlier, noqa'd).
max-branches = 24
max-statements = 70
max-returns = 9
max-args = 8
[lint.per-file-ignores]
# The e2e / ws_probe drivers are assertion scripts: scenario numbers and
# control-flow sprawl are intentional and not worth refactoring.
"tests/**" = ["PLR2004", "PLR0911", "PLR0912", "PLR0913", "PLR0915", "PLW1510"]
+33 -25
View File
@@ -26,11 +26,12 @@ machine.
Milestone M3.
"""
import contextlib
import logging
import re
import sqlite3
from pathlib import Path
from typing import Any, Dict, List, Optional, Tuple
from typing import Any
logger = logging.getLogger(__name__)
@@ -40,7 +41,7 @@ MAX_LIMIT = 100
# FTS5 special chars (mirror of hermes_state_search._FTS5_SPECIAL_CHARS) for the
# fallback sanitizer when the real one can't be imported.
_FTS5_SPECIAL_CHARS = '+{}():"^@/#&|~[]<>,;!?$=\\\''
_FTS5_SPECIAL_CHARS = "+{}():\"^@/#&|~[]<>,;!?$=\\'"
_FTS5_SPECIAL_RE = re.compile(f"[{re.escape(_FTS5_SPECIAL_CHARS)}]")
@@ -77,8 +78,7 @@ def _sanitize_fallback(query: str) -> str:
def _fts_available(conn: sqlite3.Connection) -> bool:
try:
row = conn.execute(
"SELECT 1 FROM sqlite_master WHERE type = 'table' "
"AND name = 'messages_fts' LIMIT 1"
"SELECT 1 FROM sqlite_master WHERE type = 'table' AND name = 'messages_fts' LIMIT 1"
).fetchone()
return row is not None
except sqlite3.Error:
@@ -86,11 +86,11 @@ def _fts_available(conn: sqlite3.Connection) -> bool:
def _scope_clauses(
scope: str, chat_id: Optional[str], thread_id: Optional[str]
) -> Tuple[List[str], List[Any]]:
scope: str, chat_id: str | None, thread_id: str | None
) -> tuple[list[str], list[Any]]:
"""Build the scope WHERE clauses + params (empty for scope='all')."""
clauses: List[str] = []
params: List[Any] = []
clauses: list[str] = []
params: list[Any] = []
if scope == "chat" and chat_id:
clauses.append("s.chat_id = ?")
params.append(chat_id)
@@ -100,7 +100,7 @@ def _scope_clauses(
return clauses, params
def _row_to_hit(row: sqlite3.Row) -> Dict[str, Any]:
def _row_to_hit(row: sqlite3.Row) -> dict[str, Any]:
ts = row["timestamp"]
try:
ts_ms = int(float(ts) * 1000)
@@ -120,16 +120,18 @@ def _fts_query(
conn: sqlite3.Connection,
query: str,
scope: str,
chat_id: Optional[str],
thread_id: Optional[str],
chat_id: str | None,
thread_id: str | None,
limit: int,
) -> List[Dict[str, Any]]:
) -> list[dict[str, Any]]:
where = ["messages_fts MATCH ?", "(m.active = 1 OR m.compacted = 1)"]
params: List[Any] = [query]
params: list[Any] = [query]
scope_clauses, scope_params = _scope_clauses(scope, chat_id, thread_id)
where.extend(scope_clauses)
params.extend(scope_params)
params.extend([limit])
# The f-string only splices a fixed set of static WHERE fragments; every
# user value is bound via ``?`` placeholders (see execute below).
sql = f"""
SELECT
m.id,
@@ -141,10 +143,12 @@ def _fts_query(
FROM messages_fts
JOIN messages m ON m.id = messages_fts.rowid
JOIN sessions s ON s.id = m.session_id
WHERE {' AND '.join(where)}
WHERE {" AND ".join(where)}
ORDER BY rank
LIMIT ?
"""
# Safe: every value is bound via ``?`` placeholders (no user data in SQL).
# pi-lens-ignore: python-sql-injection
rows = conn.execute(sql, params).fetchall()
return [_row_to_hit(r) for r in rows]
@@ -153,10 +157,10 @@ def _like_query(
conn: sqlite3.Connection,
query: str,
scope: str,
chat_id: Optional[str],
thread_id: Optional[str],
chat_id: str | None,
thread_id: str | None,
limit: int,
) -> List[Dict[str, Any]]:
) -> list[dict[str, Any]]:
"""Substring fallback when FTS5 is unavailable."""
# First plain word of the query is the LIKE needle (best-effort).
needle = re.split(r"\s+", query.strip(), maxsplit=1)[0].strip('"')
@@ -164,11 +168,13 @@ def _like_query(
return []
like = f"%{needle}%"
where = ["(m.active = 1 OR m.compacted = 1)", "m.content LIKE ?"]
params: List[Any] = [like]
params: list[Any] = [like]
scope_clauses, scope_params = _scope_clauses(scope, chat_id, thread_id)
where.extend(scope_clauses)
params.extend(scope_params)
params.extend([limit])
# The f-string only splices a fixed set of static WHERE fragments; every
# user value is bound via ``?`` placeholders (see execute below).
sql = f"""
SELECT
m.id,
@@ -179,12 +185,14 @@ def _like_query(
s.thread_id
FROM messages m
JOIN sessions s ON s.id = m.session_id
WHERE {' AND '.join(where)}
WHERE {" AND ".join(where)}
ORDER BY m.timestamp DESC
LIMIT ?
"""
# The needle appears twice (LIKE + instr); params order: like, scope..., needle, limit
full_params = [like, *scope_params, needle, limit]
# Safe: every value is bound via ``?`` placeholders (no user data in SQL).
# pi-lens-ignore: python-sql-injection
rows = conn.execute(sql, full_params).fetchall()
return [_row_to_hit(r) for r in rows]
@@ -193,10 +201,10 @@ def search(
db_path: Path,
query: str,
scope: str = "all",
chat_id: Optional[str] = None,
thread_id: Optional[str] = None,
chat_id: str | None = None,
thread_id: str | None = None,
limit: int = DEFAULT_LIMIT,
) -> List[Dict[str, Any]]:
) -> list[dict[str, Any]]:
"""Run a scoped search over the session store. Returns a list of hits.
Never raises: any DB/FTS error yields an empty result (the caller sends an
@@ -230,7 +238,7 @@ def search(
logger.warning("android search: query failed: %s", e)
return []
finally:
try:
# Best-effort: a close failure on a read-only connection is not
# actionable (nothing to roll back).
with contextlib.suppress(Exception):
conn.close()
except Exception:
pass
+3 -3
View File
@@ -56,8 +56,8 @@ def find_token(cli_token: str) -> str:
return env
for p in (REPO / "hermes-agent" / ".env", Path.home() / ".hermes" / ".env"):
try:
for line in p.read_text().splitlines():
line = line.strip()
for raw_line in p.read_text().splitlines():
line = raw_line.strip()
if line.startswith("ANDROID_TOKEN="):
return line.split("=", 1)[1].strip().strip('"').strip("'")
except OSError:
@@ -370,4 +370,4 @@ def main() -> int:
if __name__ == "__main__":
sys.exit(main())
sys.exit(main())
+11 -10
View File
@@ -126,7 +126,10 @@ def _print_frame(raw):
extra = (f" idx={payload.get('index')} name={payload.get('name')!r} "
f"preview={str(payload.get('preview'))[:80]!r}")
elif ftype == "tool.progress":
extra = f" idx={payload.get('index')} name={payload.get('name')!r} note={payload.get('note')!r}"
extra = (
f" idx={payload.get('index')} name={payload.get('name')!r} "
f"note={payload.get('note')!r}"
)
elif ftype == "tool.end":
extra = (f" idx={payload.get('index')} name={payload.get('name')!r} "
f"ok={payload.get('ok')} dur={payload.get('duration')}")
@@ -151,9 +154,7 @@ def _print_frame(raw):
elif ftype == "notification":
extra = (f" kind={payload.get('kind')} title={payload.get('title')!r} "
f"body={(payload.get('body') or '')[:100]!r}")
elif ftype == "sync":
extra = f" cursor={payload.get('cursor')}"
elif ftype == "sync.done":
elif ftype in {"sync", "sync.done"}:
extra = f" cursor={payload.get('cursor')}"
elif ftype == "search.results":
hits = payload.get("hits") or []
@@ -164,9 +165,7 @@ def _print_frame(raw):
extra = f" chat_id={payload.get('chat_id')}"
elif ftype == "channel.list":
extra = f" channels={len(payload.get('channels') or [])}"
elif ftype == "read.receipt":
extra = f" payload={ {k: payload[k] for k in list(payload)[:4]} }"
elif ftype == "status":
elif ftype in {"read.receipt", "status"}:
extra = f" payload={ {k: payload[k] for k in list(payload)[:4]} }"
scope = f" chat={chat}" if chat else ""
idpart = f" id={fid}" if fid is not None else ""
@@ -191,7 +190,8 @@ async def upload_file(ws, path: str, media_ref: str, next_id: int) -> int:
Returns the next free request id; raises on a non-ack terminal frame.
"""
data = open(path, "rb").read()
with open(path, "rb") as f:
data = f.read()
mime, _ = mimetypes.guess_type(path)
await ws.send(json.dumps({
"v": 1, "id": next_id, "type": "media.upload.start",
@@ -653,7 +653,8 @@ async def run(args) -> int:
data = _print_frame(raw)
if data is None:
continue
if data.get("type") == "media.offer" and (data.get("payload") or {}).get("media_id"):
is_offer = data.get("type") == "media.offer"
if is_offer and (data.get("payload") or {}).get("media_id"):
try:
await pull_media(
ws, data["payload"]["media_id"], next_id,
@@ -748,4 +749,4 @@ def main() -> int:
if __name__ == "__main__":
sys.exit(main())
sys.exit(main())
+51 -47
View File
@@ -24,11 +24,12 @@ Milestone M1.
"""
import asyncio
import contextlib
import logging
import ssl
import time
from dataclasses import dataclass, field
from typing import Any, Dict, Optional
from typing import Any
from websockets.asyncio.server import ServerConnection, serve
from websockets.exceptions import ConnectionClosed
@@ -61,6 +62,9 @@ CLOSE_REPLACED = 4402
CLOSE_RATE_LIMITED = 4403
CLOSE_SHUTDOWN = 1001
# Max length of a client-supplied device_id.
MAX_DEVICE_ID_LEN = 128
class _TokenBucket:
"""Minimal token bucket (stdlib only). One instance per connection."""
@@ -93,9 +97,9 @@ class DeviceConnection:
device_id: str
device_name: str
ws: ServerConnection
caps: Dict[str, Any] = field(default_factory=dict)
fcm_token: Optional[str] = None
ntfy_topic: Optional[str] = None
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)
@@ -108,8 +112,8 @@ class WsServer:
def __init__(self, adapter: Any, devices: DeviceRegistry):
self._adapter = adapter
self._devices = devices
self._server: Optional[Any] = None
self._connections: Dict[str, DeviceConnection] = {}
self._server: Any | None = None
self._connections: dict[str, DeviceConnection] = {}
self._lock = asyncio.Lock()
# ── Lifecycle ─────────────────────────────────────────────────────────
@@ -118,7 +122,7 @@ class WsServer:
"""Bind and start serving. Raises on bind failure (adapter maps it
to a retryable fatal error)."""
adapter = self._adapter
ssl_ctx: Optional[ssl.SSLContext] = None
ssl_ctx: ssl.SSLContext | None = None
if adapter.ws_cert and adapter.ws_key:
try:
ssl_ctx = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER)
@@ -144,36 +148,38 @@ class WsServer:
)
except OSError as e:
adapter._set_fatal_error(
"bind_failed", f"WS bind on {adapter.host}:{adapter.port} failed: {e}",
"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,
scheme,
adapter.host,
adapter.port,
)
async def stop(self) -> None:
"""Stop serving and close all device sockets."""
if self._server is not None:
self._server.close()
try:
# 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()
except Exception:
pass
self._server = None
for conn in list(self._connections.values()):
try:
# 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")
except Exception:
pass
self._connections.clear()
# ── Registry ──────────────────────────────────────────────────────────
@property
def connections(self) -> Dict[str, DeviceConnection]:
def connections(self) -> dict[str, DeviceConnection]:
return dict(self._connections)
def has_devices(self) -> bool:
@@ -182,7 +188,7 @@ class WsServer:
def device_ids(self) -> list:
return list(self._connections.keys())
def connection(self, device_id: str) -> Optional[DeviceConnection]:
def connection(self, device_id: str) -> DeviceConnection | None:
return self._connections.get(device_id)
# ── Outbound ──────────────────────────────────────────────────────────
@@ -194,11 +200,12 @@ class WsServer:
data = frame.to_json()
sent = 0
for conn in list(self._connections.values()):
try:
# 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
except Exception:
pass
return sent
async def send_to(self, device_id: str, frame: protocol.Frame) -> bool:
@@ -218,11 +225,11 @@ class WsServer:
# 1. hello auth -----------------------------------------------------
try:
raw = await asyncio.wait_for(ws.recv(), timeout=HELLO_TIMEOUT_S)
except asyncio.TimeoutError:
logger.warning("android: dropping socket with no hello (timeout)")
await self._close_quiet(ws, 1000, "no hello")
return
except ConnectionClosed:
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)
@@ -238,7 +245,7 @@ class WsServer:
return
device_id = str(payload.get("device_id") or "").strip()
if not device_id or len(device_id) > 128:
if not device_id or len(device_id) > MAX_DEVICE_ID_LEN:
await self._reject(ws, "device_id required")
return
@@ -281,10 +288,9 @@ class WsServer:
self._connections[device_id] = conn
if old is not None:
# Same device re-paired from a new socket: the new one wins.
try:
# Best-effort close of the superseded socket.
with contextlib.suppress(Exception):
await old.ws.close(code=CLOSE_REPLACED, reason="replaced by newer connection")
except Exception:
pass
ack = protocol.hello_ack(
server_caps=self._adapter.server_caps(),
@@ -311,10 +317,11 @@ class WsServer:
# flood doesn't re-trigger the error+close per frame.
if not await self._on_frame(ws, device_id, raw):
break
except ConnectionClosed:
pass
except Exception:
logger.warning("android: frame loop error for %s", device_id, exc_info=True)
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)
@@ -324,7 +331,9 @@ class WsServer:
try:
self._adapter.on_connection_closed(device_id)
except Exception:
logger.warning("android: connection cleanup failed for %s", device_id, exc_info=True)
logger.warning(
"android: connection cleanup failed for %s", device_id, exc_info=True
)
logger.info("android: device disconnected: %s", device_id)
# ── Inbound dispatch ──────────────────────────────────────────────────
@@ -346,14 +355,10 @@ class WsServer:
# 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
)
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"
),
protocol.error(protocol.ERR_RATE_LIMITED, "inbound frame rate limit exceeded"),
)
await self._close_quiet(ws, CLOSE_RATE_LIMITED, "rate limited")
return False
@@ -406,7 +411,7 @@ class WsServer:
# ── Helpers ───────────────────────────────────────────────────────────
def _connection_for(self, ws: ServerConnection) -> Optional[DeviceConnection]:
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():
@@ -415,17 +420,16 @@ class WsServer:
return None
async def _send_quiet(self, ws: ServerConnection, frame: protocol.Frame) -> None:
try:
# "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())
except Exception:
pass
async def _reject(self, ws: ServerConnection, reason: str) -> None:
await self._send_quiet(ws, protocol.error(protocol.ERR_AUTH, reason))
await self._close_quiet(ws, CLOSE_AUTH_FAILED, "auth failed")
async def _close_quiet(self, ws: ServerConnection, code: int, reason: str) -> None:
try:
# "Quiet" by contract: closing an already-closed socket is a no-op.
with contextlib.suppress(Exception):
await ws.close(code=code, reason=reason)
except Exception:
pass