a42038b825
Coolify container logs aren't accessible from the dev sandbox, so the only
log surface we can read is the Pi journal. The relay now forwards every
upstream event type to the device as {type:'debug', message:'relay evt=<T>'}
and every mid-session upstream error event as {type:'error',code:'upstream_<C>'}.
The client's Session.dispatch handles 'debug' as an INFO log line. Updates
the wake-flow integration test's receive-loop to skip over the new debug
envelopes via a ReceiveUntil helper.
236 lines
8.4 KiB
Python
236 lines
8.4 KiB
Python
"""Backend WebSocket client (asyncio)."""
|
|
import asyncio
|
|
import json
|
|
import random
|
|
from typing import Callable, Protocol
|
|
|
|
from client.log import get_logger
|
|
|
|
_log = get_logger("session")
|
|
|
|
CLIENT_VERSION = "0.1.0"
|
|
PING_INTERVAL_S = 15.0
|
|
RECONNECT_MIN_S = 1.0
|
|
RECONNECT_MAX_S = 30.0
|
|
|
|
|
|
class Callbacks(Protocol):
|
|
def on_session_started(self, conversation_id: str) -> None: ...
|
|
def on_assistant_audio(self, pcm16: bytes) -> None: ...
|
|
def on_assistant_done(self) -> None: ...
|
|
def on_session_ended(self, reason: str) -> None: ...
|
|
# Called once per backend WS lifecycle when the connection drops, so the
|
|
# owner can reset state (the conversation is dead).
|
|
def on_disconnected(self) -> None: ...
|
|
|
|
|
|
class Session:
|
|
def __init__(
|
|
self,
|
|
backend_ws: str,
|
|
device_token: str,
|
|
device_id: str,
|
|
callbacks: Callbacks,
|
|
) -> None:
|
|
self._backend_ws = backend_ws
|
|
self._device_token = device_token
|
|
self._device_id = device_id
|
|
self._cb = callbacks
|
|
self._tools: dict[str, Callable[[dict], dict]] = {}
|
|
self._ws = None
|
|
self._loop: asyncio.AbstractEventLoop | None = None
|
|
|
|
def register_tool(self, name: str, handler: Callable[[dict], dict]) -> None:
|
|
self._tools[name] = handler
|
|
|
|
async def run_forever(self) -> None:
|
|
"""Connect, pump, reconnect with exponential backoff. Never returns
|
|
cleanly except on cancellation."""
|
|
# Lazy import so the module is importable without `websockets`.
|
|
import websockets
|
|
|
|
self._loop = asyncio.get_running_loop()
|
|
backoff = RECONNECT_MIN_S
|
|
while True:
|
|
try:
|
|
_log.info("connecting to %s", self._backend_ws)
|
|
async with websockets.connect(
|
|
self._backend_ws,
|
|
additional_headers={"Authorization": f"Bearer {self._device_token}"},
|
|
max_size=4 * 1024 * 1024,
|
|
) as ws:
|
|
self._ws = ws
|
|
backoff = RECONNECT_MIN_S
|
|
ping_task = asyncio.create_task(self._ping_loop(ws))
|
|
try:
|
|
await self._pump(ws)
|
|
finally:
|
|
ping_task.cancel()
|
|
self._ws = None
|
|
self._notify_disconnected()
|
|
except asyncio.CancelledError:
|
|
raise
|
|
except Exception as exc: # noqa: BLE001
|
|
_log.warning("session error: %s; reconnecting in %.1fs", exc, backoff)
|
|
self._notify_disconnected()
|
|
jitter = random.uniform(0, backoff * 0.1)
|
|
await asyncio.sleep(backoff + jitter)
|
|
backoff = min(backoff * 2.0, RECONNECT_MAX_S)
|
|
|
|
async def _pump(self, ws) -> None:
|
|
# 1. Send hello.
|
|
await ws.send(json.dumps({
|
|
"type": "hello",
|
|
"device_id": self._device_id,
|
|
"client_version": CLIENT_VERSION,
|
|
}))
|
|
# Track ws so the threadsafe shims work when _pump is exercised directly
|
|
# (i.e. without going through run_forever, as in unit tests).
|
|
self._ws = ws
|
|
try:
|
|
self._loop = asyncio.get_running_loop()
|
|
except RuntimeError:
|
|
pass
|
|
# 2. Pump.
|
|
while True:
|
|
frame = await ws.recv()
|
|
if isinstance(frame, (bytes, bytearray, memoryview)):
|
|
self._cb.on_assistant_audio(bytes(frame))
|
|
continue
|
|
try:
|
|
env = json.loads(frame)
|
|
except json.JSONDecodeError:
|
|
_log.warning("non-JSON text frame; dropping")
|
|
continue
|
|
await self._dispatch(env, ws)
|
|
|
|
async def _dispatch(self, env: dict, ws) -> None:
|
|
t = env.get("type")
|
|
if t == "hello_ack":
|
|
_log.info("hello_ack config=%s", env.get("config"))
|
|
elif t == "pong":
|
|
pass
|
|
elif t == "session_started":
|
|
self._cb.on_session_started(str(env.get("conversation_id", "")))
|
|
elif t == "assistant_done":
|
|
self._cb.on_assistant_done()
|
|
elif t == "session_ended":
|
|
self._cb.on_session_ended(str(env.get("reason", "unknown")))
|
|
elif t == "tool_call":
|
|
await self._handle_tool_call(env, ws)
|
|
elif t == "error":
|
|
_log.warning(
|
|
"backend error code=%s message=%s fatal=%s",
|
|
env.get("code"), env.get("message"), env.get("fatal"),
|
|
)
|
|
elif t == "config_updated":
|
|
_log.info("config_updated: %s", env.get("config"))
|
|
elif t == "debug":
|
|
_log.info("backend %s", env.get("message"))
|
|
else:
|
|
_log.info("unhandled envelope type=%s", t)
|
|
|
|
async def _handle_tool_call(self, env: dict, ws) -> None:
|
|
call_id = env.get("call_id", "")
|
|
name = env.get("name", "")
|
|
args = env.get("arguments") or {}
|
|
if isinstance(args, str):
|
|
try:
|
|
args = json.loads(args)
|
|
except json.JSONDecodeError:
|
|
args = {}
|
|
handler = self._tools.get(name)
|
|
if handler is None:
|
|
await ws.send(json.dumps({
|
|
"type": "tool_result",
|
|
"call_id": call_id,
|
|
"ok": False,
|
|
"error": f"unknown tool: {name}",
|
|
}))
|
|
return
|
|
try:
|
|
result = handler(args)
|
|
except Exception as exc: # noqa: BLE001
|
|
await ws.send(json.dumps({
|
|
"type": "tool_result",
|
|
"call_id": call_id,
|
|
"ok": False,
|
|
"error": str(exc),
|
|
}))
|
|
return
|
|
await ws.send(json.dumps({
|
|
"type": "tool_result",
|
|
"call_id": call_id,
|
|
"ok": True,
|
|
"result": result,
|
|
}))
|
|
|
|
def _notify_disconnected(self) -> None:
|
|
# Tolerant: Callbacks Protocol marks on_disconnected as part of the
|
|
# surface, but real callers (and tests) may omit it.
|
|
fn = getattr(self._cb, "on_disconnected", None)
|
|
if fn is None:
|
|
return
|
|
try:
|
|
fn()
|
|
except Exception as exc: # noqa: BLE001
|
|
_log.warning("on_disconnected raised: %s", exc)
|
|
|
|
async def _ping_loop(self, ws) -> None:
|
|
try:
|
|
while True:
|
|
await asyncio.sleep(PING_INTERVAL_S)
|
|
await ws.send(json.dumps({"type": "ping"}))
|
|
except asyncio.CancelledError:
|
|
pass
|
|
|
|
# ---- Public helpers used by main.py from the asyncio loop ----
|
|
#
|
|
# All of these are called fire-and-forget via call_soon_threadsafe, so they
|
|
# must NEVER raise — the asyncio loop has nowhere to surface the exception
|
|
# except as an unhandled-task-exception traceback in the journal.
|
|
|
|
async def _safe_send(self, payload) -> None:
|
|
if self._ws is None:
|
|
return
|
|
try:
|
|
await self._ws.send(payload)
|
|
except Exception as exc: # noqa: BLE001
|
|
# ConnectionClosed / ConnectionClosedOK / ConnectionClosedError all
|
|
# land here. The reconnect loop already handles WS lifecycle.
|
|
_log.debug("send dropped: %s", exc)
|
|
|
|
async def send_wake(self) -> None:
|
|
await self._safe_send(json.dumps({"type": "wake"}))
|
|
|
|
async def send_binary(self, pcm16: bytes) -> None:
|
|
await self._safe_send(pcm16)
|
|
|
|
async def send_playback_done(self, label: str) -> None:
|
|
await self._safe_send(json.dumps({"type": "playback_done", "for": label}))
|
|
|
|
async def send_cancel(self) -> None:
|
|
await self._safe_send(json.dumps({"type": "cancel"}))
|
|
|
|
def call_soon_threadsafe_send_binary(self, pcm16: bytes) -> None:
|
|
"""Thread-safe shim used by the bridge thread to queue an uplink send."""
|
|
if self._loop is None or self._ws is None:
|
|
return
|
|
self._loop.call_soon_threadsafe(
|
|
lambda: asyncio.create_task(self.send_binary(pcm16))
|
|
)
|
|
|
|
def call_soon_threadsafe_send_wake(self) -> None:
|
|
if self._loop is None or self._ws is None:
|
|
return
|
|
self._loop.call_soon_threadsafe(
|
|
lambda: asyncio.create_task(self.send_wake())
|
|
)
|
|
|
|
def call_soon_threadsafe_send_playback_done(self, label: str) -> None:
|
|
if self._loop is None or self._ws is None:
|
|
return
|
|
self._loop.call_soon_threadsafe(
|
|
lambda: asyncio.create_task(self.send_playback_done(label))
|
|
)
|