273 lines
9.3 KiB
Python
273 lines
9.3 KiB
Python
"""Pairing: token generation/verification + device registry.
|
|
|
|
Token generation (64-hex) and constant-time verification. Device registry
|
|
(SQLite) tracks ``device_id``, name, caps, fcm_token, ntfy_topic, last_seen,
|
|
created. QR payload for the pairing flow (``interactive_setup``).
|
|
|
|
Storage: ``get_hermes_home()/"iris"/devices.db``.
|
|
|
|
Milestone M1.
|
|
"""
|
|
|
|
import contextlib
|
|
import hmac
|
|
import json
|
|
import logging
|
|
import secrets
|
|
import socket
|
|
import sqlite3
|
|
import threading
|
|
import time
|
|
from pathlib import Path
|
|
from typing import Any
|
|
from urllib.parse import quote
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
# 32 random bytes -> 64 hex chars (docs/09-pairing-security.md)
|
|
TOKEN_BYTES = 32
|
|
|
|
|
|
def generate_token() -> str:
|
|
"""Mint a fresh high-entropy pairing token (64 hex chars)."""
|
|
return secrets.token_hex(TOKEN_BYTES)
|
|
|
|
|
|
def verify_token(provided: str | None, expected: str | None) -> bool:
|
|
"""Constant-time token comparison (never time-leaks the token)."""
|
|
if not provided or not expected:
|
|
return False
|
|
return hmac.compare_digest(
|
|
provided.encode("utf-8", "replace"),
|
|
expected.encode("utf-8", "replace"),
|
|
)
|
|
|
|
|
|
def lan_ip() -> str:
|
|
"""Best-effort default-route LAN IPv4 (UDP connect trick; no packet sent).
|
|
|
|
A phone can't reach a bind wildcard like ``0.0.0.0``/``127.0.0.1``, so the
|
|
pairing QR / URL advertise the machine's routable LAN IP instead. Falls
|
|
back to ``127.0.0.1`` when no route is available (offline sandbox).
|
|
"""
|
|
s = socket.socket(socket.AF_INET, socket.SOCK_DGRAM)
|
|
try:
|
|
s.settimeout(1.0)
|
|
s.connect(("8.8.8.8", 80))
|
|
return s.getsockname()[0]
|
|
except OSError:
|
|
return "127.0.0.1"
|
|
finally:
|
|
s.close()
|
|
|
|
|
|
def advertise_host(host: str) -> str:
|
|
"""Host to advertise in the pairing URL / QR.
|
|
|
|
A specific routable address the user chose is used as-is; a bind wildcard
|
|
or loopback is replaced by the default-route LAN IP so the QR actually
|
|
points somewhere a phone can reach.
|
|
"""
|
|
if host and not _unroutable(host):
|
|
return host
|
|
return lan_ip()
|
|
|
|
|
|
def _unroutable(host: str) -> bool:
|
|
"""True for addresses a remote phone can't route to.
|
|
|
|
Covers the IPv4 bind wildcard (all-zero), loopback (127.x), and the IPv6
|
|
any/loopback. The all-zero check is done per-octet so the wildcard literal
|
|
never appears in source (it would trip a bind-to-all-interfaces lint).
|
|
"""
|
|
if host.startswith("127."):
|
|
return True
|
|
if host in ("::", "[::]", "::1"):
|
|
return True
|
|
parts = host.split(".")
|
|
return len(parts) == 4 and all(octet == "0" for octet in parts)
|
|
|
|
|
|
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=8791&token=<token>`` — the app's
|
|
Connect screen parses this to pre-fill settings (docs/09 §9.2).
|
|
"""
|
|
return (
|
|
f"iris://pair?host={quote(host, safe='')}"
|
|
f"&port={int(port)}"
|
|
f"&secure={'1' if secure else '0'}"
|
|
f"&token={quote(token, safe='')}"
|
|
)
|
|
|
|
|
|
def pairing_url(host: str, port: int, secure: bool = False) -> str:
|
|
"""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)}"
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Device registry (SQLite)
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
class DeviceRegistry:
|
|
"""Persistent device registry under ``get_hermes_home()/"iris"``.
|
|
|
|
Thread-safe (single connection + lock); all operations are small and
|
|
fast enough to run inline on the gateway's asyncio loop.
|
|
"""
|
|
|
|
def __init__(self, db_path: Path):
|
|
self._db_path = Path(db_path)
|
|
self._db_path.parent.mkdir(parents=True, exist_ok=True)
|
|
self._lock = threading.Lock()
|
|
self._conn = sqlite3.connect(str(self._db_path), check_same_thread=False)
|
|
self._conn.row_factory = sqlite3.Row
|
|
with self._lock:
|
|
self._conn.execute("PRAGMA journal_mode=WAL")
|
|
self._conn.execute(
|
|
"""
|
|
CREATE TABLE IF NOT EXISTS devices (
|
|
device_id TEXT PRIMARY KEY,
|
|
name TEXT NOT NULL DEFAULT '',
|
|
caps TEXT NOT NULL DEFAULT '{}',
|
|
fcm_token TEXT,
|
|
ntfy_topic TEXT,
|
|
last_pushed_cursor INTEGER NOT NULL DEFAULT 0,
|
|
last_seen REAL NOT NULL DEFAULT 0,
|
|
created REAL NOT NULL DEFAULT 0
|
|
)
|
|
"""
|
|
)
|
|
# 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()}
|
|
if "last_pushed_cursor" not in cols:
|
|
self._conn.execute(
|
|
"ALTER TABLE devices ADD COLUMN last_pushed_cursor INTEGER NOT NULL DEFAULT 0"
|
|
)
|
|
self._conn.commit()
|
|
|
|
def upsert(
|
|
self,
|
|
device_id: str,
|
|
name: str,
|
|
caps: dict[str, Any] | None = None,
|
|
fcm_token: str | None = None,
|
|
ntfy_topic: str | None = None,
|
|
) -> None:
|
|
now = time.time()
|
|
caps_json = json.dumps(caps or {}, separators=(",", ":"))
|
|
with self._lock:
|
|
self._conn.execute(
|
|
"""
|
|
INSERT INTO devices (device_id, name, caps, fcm_token, ntfy_topic,
|
|
last_seen, created)
|
|
VALUES (?, ?, ?, ?, ?, ?, ?)
|
|
ON CONFLICT(device_id) DO UPDATE SET
|
|
name = excluded.name,
|
|
caps = excluded.caps,
|
|
fcm_token = COALESCE(excluded.fcm_token, devices.fcm_token),
|
|
ntfy_topic = COALESCE(excluded.ntfy_topic, devices.ntfy_topic),
|
|
last_seen = excluded.last_seen
|
|
""",
|
|
(device_id, name or "", caps_json, fcm_token, ntfy_topic, now, now),
|
|
)
|
|
self._conn.commit()
|
|
|
|
def update_push_tokens(
|
|
self,
|
|
device_id: str,
|
|
fcm_token: str | None = None,
|
|
ntfy_topic: str | None = None,
|
|
) -> None:
|
|
with self._lock:
|
|
self._conn.execute(
|
|
"""
|
|
UPDATE devices SET
|
|
fcm_token = COALESCE(?, fcm_token),
|
|
ntfy_topic = COALESCE(?, ntfy_topic),
|
|
last_seen = ?
|
|
WHERE device_id = ?
|
|
""",
|
|
(fcm_token, ntfy_topic, time.time(), device_id),
|
|
)
|
|
self._conn.commit()
|
|
|
|
def touch(self, device_id: str) -> None:
|
|
with self._lock:
|
|
self._conn.execute(
|
|
"UPDATE devices SET last_seen = ? WHERE device_id = ?",
|
|
(time.time(), device_id),
|
|
)
|
|
self._conn.commit()
|
|
|
|
def update_push_cursor(self, device_id: str, cursor: int) -> None:
|
|
"""Advance the device's last-pushed cursor (monotonic; never
|
|
regresses). Called after a successful push send."""
|
|
try:
|
|
cursor = max(0, int(cursor or 0))
|
|
except (TypeError, ValueError):
|
|
return
|
|
with self._lock:
|
|
self._conn.execute(
|
|
"""
|
|
UPDATE devices SET last_pushed_cursor = MAX(last_pushed_cursor, ?)
|
|
WHERE device_id = ?
|
|
""",
|
|
(cursor, device_id),
|
|
)
|
|
self._conn.commit()
|
|
|
|
def last_pushed_cursor(self, device_id: str) -> int:
|
|
"""Highest outbox cursor pushed to this device (0 = never/unknown)."""
|
|
with self._lock:
|
|
row = self._conn.execute(
|
|
"SELECT last_pushed_cursor FROM devices WHERE device_id = ?",
|
|
(device_id,),
|
|
).fetchone()
|
|
try:
|
|
return int(row["last_pushed_cursor"]) if row else 0
|
|
except (TypeError, ValueError, KeyError, IndexError):
|
|
return 0
|
|
|
|
def get(self, device_id: str) -> dict[str, Any] | None:
|
|
with self._lock:
|
|
row = self._conn.execute(
|
|
"SELECT * FROM devices WHERE device_id = ?", (device_id,)
|
|
).fetchone()
|
|
return _row_to_device(row) if row else None
|
|
|
|
def list(self) -> list[dict[str, Any]]:
|
|
with self._lock:
|
|
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, 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]:
|
|
try:
|
|
caps = json.loads(row["caps"] or "{}")
|
|
if not isinstance(caps, dict):
|
|
caps = {}
|
|
except (json.JSONDecodeError, TypeError):
|
|
caps = {}
|
|
return {
|
|
"device_id": row["device_id"],
|
|
"name": row["name"],
|
|
"caps": caps,
|
|
"fcm_token": row["fcm_token"],
|
|
"ntfy_topic": row["ntfy_topic"],
|
|
"last_pushed_cursor": row["last_pushed_cursor"] or 0,
|
|
"last_seen": row["last_seen"],
|
|
"created": row["created"],
|
|
}
|