realtime.py 15 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433
  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 await ChattoClient.login(...) 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, Literal, overload
  18. from chattolib import _pb # noqa: F401 — installs pb import path
  19. from chattolib.exceptions import ChattoError
  20. from chattolib.realtime_types import RealtimePayload, parse_payload
  21. from chattolib.types import parse_datetime
  22. if TYPE_CHECKING:
  23. from chattolib.client import ChattoClient
  24. from chattolib.realtime_types import (
  25. AssetDeletedPayload,
  26. AssetProcessingPayload,
  27. CallEventPayload,
  28. MentionNotificationPayload,
  29. MessageEditedPayload,
  30. MessagePostedPayload,
  31. MessageRetractedPayload,
  32. NewDirectMessageNotificationPayload,
  33. NotificationCreatedPayload,
  34. NotificationDismissedPayload,
  35. NotificationLevelChangedPayload,
  36. PresenceChangedPayload,
  37. ReactionPayload,
  38. RoomEventPayload,
  39. RoomGroupsUpdatedPayload,
  40. RoomMarkedAsReadPayload,
  41. RoomUniversalChangedPayload,
  42. ServerMemberDeletedPayload,
  43. ServerUpdatedPayload,
  44. ServerUserPreferencesUpdatedPayload,
  45. SessionTerminatedPayload,
  46. ThreadCreatedPayload,
  47. ThreadFollowChangedPayload,
  48. TypingPayload,
  49. UserCustomStatusClearedPayload,
  50. UserCustomStatusSetPayload,
  51. UserProfileUpdatedPayload,
  52. )
  53. REALTIME_PATH = "/api/realtime"
  54. REALTIME_PROTOCOL_VERSION = 1
  55. class ChattoRealtimeError(ChattoError):
  56. """Server returned a protocol error over the realtime WebSocket."""
  57. def __init__(self, code: str, message: str, *, fatal: bool = False) -> None:
  58. self.code = code
  59. self.message = message
  60. self.fatal = fatal
  61. super().__init__(f"{code}: {message}")
  62. class ChattoRealtimeCloseError(ChattoError):
  63. """Server sent a close frame."""
  64. def __init__(
  65. self,
  66. code: str,
  67. message: str,
  68. *,
  69. reconnect: bool = False,
  70. retry_after_ms: int = 0,
  71. ) -> None:
  72. self.code = code
  73. self.message = message
  74. self.reconnect = reconnect
  75. self.retry_after_ms = retry_after_ms
  76. super().__init__(f"{code}: {message}")
  77. def realtime_url(base_url: str) -> str:
  78. """Convert an HTTP(S) base URL to the realtime WebSocket URL."""
  79. base = base_url.rstrip("/")
  80. if base.startswith("https://"):
  81. return "wss://" + base[len("https://") :] + REALTIME_PATH
  82. if base.startswith("http://"):
  83. return "ws://" + base[len("http://") :] + REALTIME_PATH
  84. return "wss://" + base + REALTIME_PATH
  85. @dataclass
  86. class ServerHello:
  87. """Server's response to the initial hello frame."""
  88. protocol_version: int
  89. server_version: str
  90. heartbeat_interval_seconds: int
  91. capabilities: list[str]
  92. @dataclass
  93. class RealtimeEvent:
  94. """One live event delivered over the realtime WebSocket.
  95. ``kind`` names the ``oneof event`` case set on the envelope
  96. (``message_posted``, ``reaction_added``, ``presence_changed``, …).
  97. ``payload`` is the typed payload dataclass for that kind (see
  98. :mod:`chattolib.realtime_types`), or ``None`` for unknown/future kinds.
  99. ``raw`` exposes the original protobuf envelope as an escape hatch.
  100. """
  101. id: str
  102. created_at: datetime | None
  103. actor_id: str | None
  104. kind: str
  105. payload: RealtimePayload | None
  106. raw: Any # the full RealtimeEventEnvelope
  107. @overload
  108. def get(self, kind: Literal["message_posted"]) -> MessagePostedPayload | None: ...
  109. @overload
  110. def get(self, kind: Literal["message_edited"]) -> MessageEditedPayload | None: ...
  111. @overload
  112. def get(self, kind: Literal["message_retracted"]) -> MessageRetractedPayload | None: ...
  113. @overload
  114. def get(
  115. self, kind: Literal["reaction_added", "reaction_removed"]
  116. ) -> ReactionPayload | None: ...
  117. @overload
  118. def get(self, kind: Literal["user_typing"]) -> TypingPayload | None: ...
  119. @overload
  120. def get(self, kind: Literal["presence_changed"]) -> PresenceChangedPayload | None: ...
  121. @overload
  122. def get(
  123. self,
  124. kind: Literal[
  125. "room_created",
  126. "room_updated",
  127. "room_deleted",
  128. "room_archived",
  129. "room_unarchived",
  130. "user_joined_room",
  131. "user_left_room",
  132. ],
  133. ) -> RoomEventPayload | None: ...
  134. @overload
  135. def get(
  136. self, kind: Literal["room_universal_changed"]
  137. ) -> RoomUniversalChangedPayload | None: ...
  138. @overload
  139. def get(self, kind: Literal["notification_created"]) -> NotificationCreatedPayload | None: ...
  140. @overload
  141. def get(
  142. self, kind: Literal["notification_dismissed"]
  143. ) -> NotificationDismissedPayload | None: ...
  144. @overload
  145. def get(
  146. self, kind: Literal["notification_level_changed"]
  147. ) -> NotificationLevelChangedPayload | None: ...
  148. @overload
  149. def get(self, kind: Literal["thread_follow_changed"]) -> ThreadFollowChangedPayload | None: ...
  150. @overload
  151. def get(self, kind: Literal["thread_created"]) -> ThreadCreatedPayload | None: ...
  152. @overload
  153. def get(self, kind: Literal["room_marked_as_read"]) -> RoomMarkedAsReadPayload | None: ...
  154. @overload
  155. def get(self, kind: Literal["server_updated"]) -> ServerUpdatedPayload | None: ...
  156. @overload
  157. def get(self, kind: Literal["user_profile_updated"]) -> UserProfileUpdatedPayload | None: ...
  158. @overload
  159. def get(self, kind: Literal["user_custom_status_set"]) -> UserCustomStatusSetPayload | None: ...
  160. @overload
  161. def get(
  162. self, kind: Literal["user_custom_status_cleared"]
  163. ) -> UserCustomStatusClearedPayload | None: ...
  164. @overload
  165. def get(
  166. self, kind: Literal["server_user_preferences_updated"]
  167. ) -> ServerUserPreferencesUpdatedPayload | None: ...
  168. @overload
  169. def get(self, kind: Literal["room_groups_updated"]) -> RoomGroupsUpdatedPayload | None: ...
  170. @overload
  171. def get(self, kind: Literal["server_member_deleted"]) -> ServerMemberDeletedPayload | None: ...
  172. @overload
  173. def get(
  174. self,
  175. kind: Literal[
  176. "asset_processing_started", "asset_processing_succeeded", "asset_processing_failed"
  177. ],
  178. ) -> AssetProcessingPayload | None: ...
  179. @overload
  180. def get(self, kind: Literal["asset_deleted"]) -> AssetDeletedPayload | None: ...
  181. @overload
  182. def get(
  183. self,
  184. kind: Literal[
  185. "call_started", "call_participant_joined", "call_participant_left", "call_ended"
  186. ],
  187. ) -> CallEventPayload | None: ...
  188. @overload
  189. def get(self, kind: Literal["mention_notification"]) -> MentionNotificationPayload | None: ...
  190. @overload
  191. def get(
  192. self, kind: Literal["new_direct_message_notification"]
  193. ) -> NewDirectMessageNotificationPayload | None: ...
  194. @overload
  195. def get(self, kind: Literal["session_terminated"]) -> SessionTerminatedPayload | None: ...
  196. @overload
  197. def get(self, kind: str) -> RealtimePayload | None: ...
  198. def get(self, kind: str) -> RealtimePayload | None:
  199. """Return ``payload`` if this event is of ``kind``, else ``None``."""
  200. return self.payload if self.kind == kind else None
  201. class RealtimeConnection:
  202. """Live realtime WebSocket session.
  203. Prefer :func:`stream_events` for the common case; use this class directly
  204. when you also need to send client pings, close cleanly, or inspect the
  205. negotiated :class:`ServerHello`.
  206. """
  207. def __init__(
  208. self,
  209. client: ChattoClient,
  210. *,
  211. protocol_version: int = REALTIME_PROTOCOL_VERSION,
  212. ) -> None:
  213. self._client = client
  214. self._protocol_version = protocol_version
  215. self._ws: Any = None
  216. self._server_hello: ServerHello | None = None
  217. @property
  218. def server_hello(self) -> ServerHello | None:
  219. return self._server_hello
  220. async def __aenter__(self) -> RealtimeConnection:
  221. await self.connect()
  222. return self
  223. async def __aexit__(self, *exc: Any) -> None:
  224. await self.close()
  225. async def connect(self) -> None:
  226. try:
  227. import websockets
  228. except ImportError as exc: # pragma: no cover
  229. raise ChattoError(
  230. "The realtime channel requires the 'websockets' package. "
  231. "Install with `pip install chattolib[realtime]`."
  232. ) from exc
  233. from chattolib._pb.chatto.realtime.v1 import realtime_pb2
  234. url = realtime_url(self._client.base_url)
  235. headers: dict[str, str] = {}
  236. if self._client.session_cookie:
  237. headers["Cookie"] = f"chatto_session={self._client.session_cookie}"
  238. self._ws = await websockets.connect(url, additional_headers=headers)
  239. client_hello = realtime_pb2.RealtimeClientFrame()
  240. client_hello.hello.protocol_version = self._protocol_version
  241. if self._client.token:
  242. client_hello.hello.bearer_token = self._client.token
  243. await self._ws.send(client_hello.SerializeToString())
  244. first = await self._ws.recv()
  245. first_frame = realtime_pb2.RealtimeServerFrame()
  246. first_frame.ParseFromString(first)
  247. if first_frame.WhichOneof("frame") != "hello":
  248. _raise_for_control(first_frame)
  249. raise ChattoRealtimeError(
  250. "unexpected_frame",
  251. f"expected server hello, got {first_frame.WhichOneof('frame')!r}",
  252. )
  253. hello = first_frame.hello
  254. self._server_hello = ServerHello(
  255. protocol_version=hello.protocol_version,
  256. server_version=hello.server_version,
  257. heartbeat_interval_seconds=hello.heartbeat_interval_seconds,
  258. capabilities=list(hello.capabilities),
  259. )
  260. subscribe = realtime_pb2.RealtimeClientFrame()
  261. subscribe.subscribe_events.SetInParent()
  262. await self._ws.send(subscribe.SerializeToString())
  263. async def close(self) -> None:
  264. if self._ws is None:
  265. return
  266. try:
  267. await self._ws.close()
  268. finally:
  269. self._ws = None
  270. async def ping(self, nonce: str = "") -> None:
  271. """Send a client ping. The server replies with a matching pong."""
  272. if self._ws is None:
  273. raise ChattoError("realtime connection is closed")
  274. from chattolib._pb.chatto.realtime.v1 import realtime_pb2
  275. frame = realtime_pb2.RealtimeClientFrame()
  276. frame.ping.nonce = nonce
  277. await self._ws.send(frame.SerializeToString())
  278. async def events(self) -> AsyncIterator[RealtimeEvent]:
  279. """Yield decoded live events until the connection closes."""
  280. if self._ws is None:
  281. raise ChattoError("realtime connection is not open")
  282. from chattolib._pb.chatto.realtime.v1 import realtime_pb2
  283. async for raw in self._ws:
  284. if isinstance(raw, str):
  285. # The protocol is binary; a text frame indicates a protocol violation.
  286. raise ChattoRealtimeError(
  287. "unexpected_text_frame",
  288. "server sent a text frame; realtime protocol expects binary",
  289. )
  290. frame = realtime_pb2.RealtimeServerFrame()
  291. frame.ParseFromString(raw)
  292. case = frame.WhichOneof("frame")
  293. if case == "event":
  294. yield _wrap_event(frame.event)
  295. elif case in ("heartbeat", "pong", "subscribed"):
  296. continue
  297. elif case == "error":
  298. err = frame.error
  299. exc = ChattoRealtimeError(err.code, err.message, fatal=err.fatal)
  300. if err.fatal:
  301. raise exc
  302. # Non-fatal errors are surfaced but the stream continues.
  303. # Callers can still receive them if desired; for now we log
  304. # by re-raising on fatal only.
  305. continue
  306. elif case == "close":
  307. close = frame.close
  308. raise ChattoRealtimeCloseError(
  309. close.code,
  310. close.message,
  311. reconnect=close.reconnect,
  312. retry_after_ms=close.retry_after_ms,
  313. )
  314. else:
  315. raise ChattoRealtimeError(
  316. "unexpected_frame",
  317. f"unknown server frame: {case!r}",
  318. )
  319. def _raise_for_control(frame: Any) -> None:
  320. """If ``frame`` is an error or close frame, translate it and raise."""
  321. case = frame.WhichOneof("frame")
  322. if case == "error":
  323. err = frame.error
  324. raise ChattoRealtimeError(err.code, err.message, fatal=err.fatal)
  325. if case == "close":
  326. close = frame.close
  327. raise ChattoRealtimeCloseError(
  328. close.code,
  329. close.message,
  330. reconnect=close.reconnect,
  331. retry_after_ms=close.retry_after_ms,
  332. )
  333. def _wrap_event(envelope: Any) -> RealtimeEvent:
  334. kind = envelope.WhichOneof("event") or ""
  335. pb_payload = getattr(envelope, kind, None) if kind else None
  336. payload = parse_payload(kind, pb_payload) if kind else None
  337. created_at = None
  338. if envelope.HasField("created_at"):
  339. created_at = parse_datetime(envelope.created_at.ToJsonString())
  340. actor_id: str | None = None
  341. if envelope.HasField("actor_id"):
  342. actor_id = envelope.actor_id
  343. return RealtimeEvent(
  344. id=envelope.id,
  345. created_at=created_at,
  346. actor_id=actor_id,
  347. kind=kind,
  348. payload=payload,
  349. raw=envelope,
  350. )
  351. async def stream_events(
  352. client: ChattoClient,
  353. *,
  354. protocol_version: int = REALTIME_PROTOCOL_VERSION,
  355. ) -> AsyncIterator[RealtimeEvent]:
  356. """Open a realtime connection and yield events until the server closes.
  357. Raises :class:`ChattoRealtimeCloseError` when the server sends a close frame,
  358. :class:`ChattoRealtimeError` on fatal protocol errors, or
  359. :class:`ChattoConnectError` if the initial HTTP handshake fails.
  360. """
  361. conn = RealtimeConnection(client, protocol_version=protocol_version)
  362. try:
  363. await conn.connect()
  364. async for event in conn.events():
  365. yield event
  366. finally:
  367. await conn.close()
  368. __all__ = [
  369. "REALTIME_PATH",
  370. "REALTIME_PROTOCOL_VERSION",
  371. "ChattoRealtimeCloseError",
  372. "ChattoRealtimeError",
  373. "RealtimeConnection",
  374. "RealtimeEvent",
  375. "RealtimePayload",
  376. "ServerHello",
  377. "realtime_url",
  378. "stream_events",
  379. ]