realtime.py 13 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399
  1. """Chatto realtime WebSocket client.
  2. Chatto exposes a binary-protobuf realtime channel at ``/api/realtime`` (see
  3. ``proto/chatto/realtime/v1/realtime.proto``). This module speaks that
  4. protocol: it opens the WebSocket, exchanges ``hello`` frames, subscribes to
  5. the caller's authorized live-event stream, and yields decoded events.
  6. Usage::
  7. async with ChattoClient(token="cht_...") as client:
  8. async for event in stream_events(client):
  9. print(event.kind, event.payload)
  10. Requires the ``chattolib[realtime]`` extra (which pulls in ``websockets`` and
  11. ``protobuf``).
  12. """
  13. from __future__ import annotations
  14. from collections.abc import AsyncIterator
  15. from dataclasses import dataclass
  16. from datetime import datetime
  17. from typing import TYPE_CHECKING, Any
  18. from chattolib import _pb # noqa: F401 — installs pb import path
  19. from chattolib.exceptions import ChattoError
  20. from chattolib.types import parse_datetime
  21. if TYPE_CHECKING:
  22. from chattolib.client import ChattoClient
  23. REALTIME_PATH = "/api/realtime"
  24. REALTIME_PROTOCOL_VERSION = 2
  25. class ChattoRealtimeError(ChattoError):
  26. """Server returned a protocol error over the realtime WebSocket."""
  27. def __init__(self, code: str, message: str, *, fatal: bool = False) -> None:
  28. self.code = code
  29. self.message = message
  30. self.fatal = fatal
  31. super().__init__(f"{code}: {message}")
  32. class ChattoRealtimeCloseError(ChattoError):
  33. """Server sent a close frame."""
  34. def __init__(
  35. self,
  36. code: str,
  37. message: str,
  38. *,
  39. reconnect: bool = False,
  40. retry_after_ms: int = 0,
  41. ) -> None:
  42. self.code = code
  43. self.message = message
  44. self.reconnect = reconnect
  45. self.retry_after_ms = retry_after_ms
  46. super().__init__(f"{code}: {message}")
  47. def realtime_url(base_url: str) -> str:
  48. """Convert an HTTP(S) base URL to the realtime WebSocket URL."""
  49. base = base_url.rstrip("/")
  50. if base.startswith("https://"):
  51. return "wss://" + base[len("https://") :] + REALTIME_PATH
  52. if base.startswith("http://"):
  53. return "ws://" + base[len("http://") :] + REALTIME_PATH
  54. return "wss://" + base + REALTIME_PATH
  55. @dataclass
  56. class ServerHello:
  57. """Server's response to the initial hello frame."""
  58. protocol_version: int
  59. server_version: str
  60. heartbeat_interval_seconds: int
  61. capabilities: list[str]
  62. @dataclass
  63. class RealtimeProjectionOperation:
  64. """One durable state change inside a :class:`RealtimeProjectionEvent`.
  65. ``operation`` names the ``oneof operation`` case set on the operation
  66. (``reset``, ``server_upsert``, ``viewer_upsert``, ``user_upsert``,
  67. ``user_remove``, ``room_upsert``, ``room_remove``,
  68. ``room_groups_replace``, ``room_timeline_replace``,
  69. ``room_timeline_event_upsert``, ``server_state_upsert``,
  70. ``room_timeline_event_remove``, ``room_viewer_state_replace``,
  71. ``active_calls_replace``, ``presences_replace``,
  72. ``thread_viewer_states_replace``, ``room_activity``,
  73. ``notification_occurrences_replace``). ``payload`` is the concrete
  74. protobuf sub-message.
  75. """
  76. operation: str
  77. payload: Any
  78. raw: Any # the full RealtimeProjectionOperation
  79. @dataclass
  80. class RealtimeProjectionEvent:
  81. """A server-projection frame: the durable state the server wants clients
  82. to converge on, as a batch of :class:`RealtimeProjectionOperation` s.
  83. """
  84. id: str
  85. created_at: datetime | None
  86. actor_id: str | None
  87. resume_cursor: str | None
  88. operations: list[RealtimeProjectionOperation]
  89. raw: Any # the full RealtimeProjectionEvent
  90. @dataclass
  91. class RealtimeEvent:
  92. """One live event delivered over the realtime WebSocket.
  93. ``kind`` names the ``oneof event`` case set on the envelope
  94. (``message_posted``, ``reaction_added``, ``presence_changed``, …).
  95. ``payload`` is the concrete protobuf sub-message; access its fields
  96. directly (e.g. ``event.payload.room_id``). Callers that want to hydrate
  97. the referenced resource should follow the hydration hints documented on
  98. each event message in ``realtime.proto``.
  99. """
  100. id: str
  101. created_at: datetime | None
  102. actor_id: str | None
  103. kind: str
  104. payload: Any
  105. raw: Any # the full RealtimeEventEnvelope
  106. class RealtimeConnection:
  107. """Live realtime WebSocket session.
  108. Prefer :func:`stream_events` for the common case; use this class directly
  109. when you also need to send client pings, close cleanly, or inspect the
  110. negotiated :class:`ServerHello`.
  111. """
  112. def __init__(
  113. self,
  114. client: ChattoClient,
  115. *,
  116. protocol_version: int = REALTIME_PROTOCOL_VERSION,
  117. resume_cursor: str | None = None,
  118. retained_room_ids: list[str] | None = None,
  119. ) -> None:
  120. self._client = client
  121. self._protocol_version = protocol_version
  122. self._resume_cursor = resume_cursor
  123. self._retained_room_ids = list(retained_room_ids or [])
  124. self._ws: Any = None
  125. self._server_hello: ServerHello | None = None
  126. @property
  127. def server_hello(self) -> ServerHello | None:
  128. return self._server_hello
  129. async def __aenter__(self) -> RealtimeConnection:
  130. await self.connect()
  131. return self
  132. async def __aexit__(self, *exc: Any) -> None:
  133. await self.close()
  134. async def connect(self) -> None:
  135. try:
  136. import websockets
  137. except ImportError as exc: # pragma: no cover
  138. raise ChattoError(
  139. "The realtime channel requires the 'websockets' package. "
  140. "Install with `pip install chattolib[realtime]`."
  141. ) from exc
  142. from chattolib._pb.chatto.realtime.v1 import realtime_pb2
  143. url = realtime_url(self._client.base_url)
  144. headers: dict[str, str] = {}
  145. if self._client.session_cookie:
  146. headers["Cookie"] = f"chatto_session={self._client.session_cookie}"
  147. self._ws = await websockets.connect(url, additional_headers=headers)
  148. client_hello = realtime_pb2.RealtimeClientFrame()
  149. client_hello.hello.protocol_version = self._protocol_version
  150. if self._client.token:
  151. client_hello.hello.bearer_token = self._client.token
  152. await self._ws.send(client_hello.SerializeToString())
  153. first = await self._ws.recv()
  154. first_frame = realtime_pb2.RealtimeServerFrame()
  155. first_frame.ParseFromString(first)
  156. if first_frame.WhichOneof("frame") != "hello":
  157. _raise_for_control(first_frame)
  158. raise ChattoRealtimeError(
  159. "unexpected_frame",
  160. f"expected server hello, got {first_frame.WhichOneof('frame')!r}",
  161. )
  162. hello = first_frame.hello
  163. self._server_hello = ServerHello(
  164. protocol_version=hello.protocol_version,
  165. server_version=hello.server_version,
  166. heartbeat_interval_seconds=hello.heartbeat_interval_seconds,
  167. capabilities=list(hello.capabilities),
  168. )
  169. subscribe = realtime_pb2.RealtimeClientFrame()
  170. subscribe.subscribe_events.SetInParent()
  171. if self._resume_cursor:
  172. subscribe.subscribe_events.resume_cursor = self._resume_cursor
  173. subscribe.subscribe_events.retained_room_ids.extend(self._retained_room_ids)
  174. await self._ws.send(subscribe.SerializeToString())
  175. async def close(self) -> None:
  176. if self._ws is None:
  177. return
  178. try:
  179. await self._ws.close()
  180. finally:
  181. self._ws = None
  182. async def ping(self, nonce: str = "") -> None:
  183. """Send a client ping. The server replies with a matching pong."""
  184. if self._ws is None:
  185. raise ChattoError("realtime connection is closed")
  186. from chattolib._pb.chatto.realtime.v1 import realtime_pb2
  187. frame = realtime_pb2.RealtimeClientFrame()
  188. frame.ping.nonce = nonce
  189. await self._ws.send(frame.SerializeToString())
  190. async def hydrate_room(self, room_id: str) -> None:
  191. """Ask the server to re-send the projection state for a room."""
  192. if self._ws is None:
  193. raise ChattoError("realtime connection is closed")
  194. from chattolib._pb.chatto.realtime.v1 import realtime_pb2
  195. frame = realtime_pb2.RealtimeClientFrame()
  196. frame.hydrate_room.room_id = room_id
  197. await self._ws.send(frame.SerializeToString())
  198. async def events(
  199. self,
  200. ) -> AsyncIterator[RealtimeEvent | RealtimeProjectionEvent]:
  201. """Yield decoded live events and projection events until the
  202. connection closes."""
  203. if self._ws is None:
  204. raise ChattoError("realtime connection is not open")
  205. from chattolib._pb.chatto.realtime.v1 import realtime_pb2
  206. async for raw in self._ws:
  207. if isinstance(raw, str):
  208. # The protocol is binary; a text frame indicates a protocol violation.
  209. raise ChattoRealtimeError(
  210. "unexpected_text_frame",
  211. "server sent a text frame; realtime protocol expects binary",
  212. )
  213. frame = realtime_pb2.RealtimeServerFrame()
  214. frame.ParseFromString(raw)
  215. case = frame.WhichOneof("frame")
  216. if case == "event":
  217. yield _wrap_event(frame.event)
  218. elif case in ("heartbeat", "pong", "subscribed", "caught_up"):
  219. continue
  220. elif case == "projection_event":
  221. yield _wrap_projection_event(frame.projection_event)
  222. elif case == "error":
  223. err = frame.error
  224. exc = ChattoRealtimeError(err.code, err.message, fatal=err.fatal)
  225. if err.fatal:
  226. raise exc
  227. # Non-fatal errors are surfaced but the stream continues.
  228. # Callers can still receive them if desired; for now we log
  229. # by re-raising on fatal only.
  230. continue
  231. elif case == "close":
  232. close = frame.close
  233. raise ChattoRealtimeCloseError(
  234. close.code,
  235. close.message,
  236. reconnect=close.reconnect,
  237. retry_after_ms=close.retry_after_ms,
  238. )
  239. else:
  240. raise ChattoRealtimeError(
  241. "unexpected_frame",
  242. f"unknown server frame: {case!r}",
  243. )
  244. def _raise_for_control(frame: Any) -> None:
  245. """If ``frame`` is an error or close frame, translate it and raise."""
  246. case = frame.WhichOneof("frame")
  247. if case == "error":
  248. err = frame.error
  249. raise ChattoRealtimeError(err.code, err.message, fatal=err.fatal)
  250. if case == "close":
  251. close = frame.close
  252. raise ChattoRealtimeCloseError(
  253. close.code,
  254. close.message,
  255. reconnect=close.reconnect,
  256. retry_after_ms=close.retry_after_ms,
  257. )
  258. def _wrap_projection_event(envelope: Any) -> RealtimeProjectionEvent:
  259. created_at = None
  260. if envelope.HasField("created_at"):
  261. created_at = parse_datetime(envelope.created_at.ToJsonString())
  262. actor_id: str | None = None
  263. if envelope.HasField("actor_id"):
  264. actor_id = envelope.actor_id
  265. resume_cursor: str | None = None
  266. if envelope.HasField("resume_cursor"):
  267. resume_cursor = envelope.resume_cursor
  268. operations = [_wrap_projection_operation(op) for op in envelope.operations]
  269. return RealtimeProjectionEvent(
  270. id=envelope.id,
  271. created_at=created_at,
  272. actor_id=actor_id,
  273. resume_cursor=resume_cursor,
  274. operations=operations,
  275. raw=envelope,
  276. )
  277. def _wrap_projection_operation(op: Any) -> RealtimeProjectionOperation:
  278. kind = op.WhichOneof("operation") or ""
  279. payload = getattr(op, kind, None) if kind else None
  280. return RealtimeProjectionOperation(operation=kind, payload=payload, raw=op)
  281. def _wrap_event(envelope: Any) -> RealtimeEvent:
  282. kind = envelope.WhichOneof("event") or ""
  283. payload = getattr(envelope, kind, None) if kind else None
  284. created_at = None
  285. if envelope.HasField("created_at"):
  286. created_at = parse_datetime(envelope.created_at.ToJsonString())
  287. actor_id: str | None = None
  288. if envelope.HasField("actor_id"):
  289. actor_id = envelope.actor_id
  290. return RealtimeEvent(
  291. id=envelope.id,
  292. created_at=created_at,
  293. actor_id=actor_id,
  294. kind=kind,
  295. payload=payload,
  296. raw=envelope,
  297. )
  298. async def stream_events(
  299. client: ChattoClient,
  300. *,
  301. protocol_version: int = REALTIME_PROTOCOL_VERSION,
  302. resume_cursor: str | None = None,
  303. retained_room_ids: list[str] | None = None,
  304. ) -> AsyncIterator[RealtimeEvent | RealtimeProjectionEvent]:
  305. """Open a realtime connection and yield events until the server closes.
  306. Raises :class:`ChattoRealtimeCloseError` when the server sends a close frame,
  307. :class:`ChattoRealtimeError` on fatal protocol errors, or
  308. :class:`ChattoConnectError` if the initial HTTP handshake fails.
  309. """
  310. conn = RealtimeConnection(
  311. client,
  312. protocol_version=protocol_version,
  313. resume_cursor=resume_cursor,
  314. retained_room_ids=retained_room_ids,
  315. )
  316. try:
  317. await conn.connect()
  318. async for event in conn.events():
  319. yield event
  320. finally:
  321. await conn.close()
  322. __all__ = [
  323. "REALTIME_PATH",
  324. "REALTIME_PROTOCOL_VERSION",
  325. "ChattoRealtimeCloseError",
  326. "ChattoRealtimeError",
  327. "RealtimeConnection",
  328. "RealtimeEvent",
  329. "RealtimeProjectionEvent",
  330. "RealtimeProjectionOperation",
  331. "ServerHello",
  332. "realtime_url",
  333. "stream_events",
  334. ]