| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433 |
- """Chatto realtime WebSocket client.
- Chatto exposes a binary-protobuf realtime channel at ``/api/realtime`` (see
- ``proto/chatto/realtime/v1/realtime.proto``). This module speaks that
- protocol: it opens the WebSocket, exchanges ``hello`` frames, subscribes to
- the caller's authorized live-event stream, and yields decoded events.
- Usage::
- async with await ChattoClient.login(...) as client:
- async for event in stream_events(client):
- print(event.kind, event.payload)
- Requires the ``chattolib[realtime]`` extra (which pulls in ``websockets`` and
- ``protobuf``).
- """
- from __future__ import annotations
- from collections.abc import AsyncIterator
- from dataclasses import dataclass
- from datetime import datetime
- from typing import TYPE_CHECKING, Any, Literal, overload
- from chattolib import _pb # noqa: F401 — installs pb import path
- from chattolib.exceptions import ChattoError
- from chattolib.realtime_types import RealtimePayload, parse_payload
- from chattolib.types import parse_datetime
- if TYPE_CHECKING:
- from chattolib.client import ChattoClient
- from chattolib.realtime_types import (
- AssetDeletedPayload,
- AssetProcessingPayload,
- CallEventPayload,
- MentionNotificationPayload,
- MessageEditedPayload,
- MessagePostedPayload,
- MessageRetractedPayload,
- NewDirectMessageNotificationPayload,
- NotificationCreatedPayload,
- NotificationDismissedPayload,
- NotificationLevelChangedPayload,
- PresenceChangedPayload,
- ReactionPayload,
- RoomEventPayload,
- RoomGroupsUpdatedPayload,
- RoomMarkedAsReadPayload,
- RoomUniversalChangedPayload,
- ServerMemberDeletedPayload,
- ServerUpdatedPayload,
- ServerUserPreferencesUpdatedPayload,
- SessionTerminatedPayload,
- ThreadCreatedPayload,
- ThreadFollowChangedPayload,
- TypingPayload,
- UserCustomStatusClearedPayload,
- UserCustomStatusSetPayload,
- UserProfileUpdatedPayload,
- )
- REALTIME_PATH = "/api/realtime"
- REALTIME_PROTOCOL_VERSION = 1
- class ChattoRealtimeError(ChattoError):
- """Server returned a protocol error over the realtime WebSocket."""
- def __init__(self, code: str, message: str, *, fatal: bool = False) -> None:
- self.code = code
- self.message = message
- self.fatal = fatal
- super().__init__(f"{code}: {message}")
- class ChattoRealtimeCloseError(ChattoError):
- """Server sent a close frame."""
- def __init__(
- self,
- code: str,
- message: str,
- *,
- reconnect: bool = False,
- retry_after_ms: int = 0,
- ) -> None:
- self.code = code
- self.message = message
- self.reconnect = reconnect
- self.retry_after_ms = retry_after_ms
- super().__init__(f"{code}: {message}")
- def realtime_url(base_url: str) -> str:
- """Convert an HTTP(S) base URL to the realtime WebSocket URL."""
- base = base_url.rstrip("/")
- if base.startswith("https://"):
- return "wss://" + base[len("https://") :] + REALTIME_PATH
- if base.startswith("http://"):
- return "ws://" + base[len("http://") :] + REALTIME_PATH
- return "wss://" + base + REALTIME_PATH
- @dataclass
- class ServerHello:
- """Server's response to the initial hello frame."""
- protocol_version: int
- server_version: str
- heartbeat_interval_seconds: int
- capabilities: list[str]
- @dataclass
- class RealtimeEvent:
- """One live event delivered over the realtime WebSocket.
- ``kind`` names the ``oneof event`` case set on the envelope
- (``message_posted``, ``reaction_added``, ``presence_changed``, …).
- ``payload`` is the typed payload dataclass for that kind (see
- :mod:`chattolib.realtime_types`), or ``None`` for unknown/future kinds.
- ``raw`` exposes the original protobuf envelope as an escape hatch.
- """
- id: str
- created_at: datetime | None
- actor_id: str | None
- kind: str
- payload: RealtimePayload | None
- raw: Any # the full RealtimeEventEnvelope
- @overload
- def get(self, kind: Literal["message_posted"]) -> MessagePostedPayload | None: ...
- @overload
- def get(self, kind: Literal["message_edited"]) -> MessageEditedPayload | None: ...
- @overload
- def get(self, kind: Literal["message_retracted"]) -> MessageRetractedPayload | None: ...
- @overload
- def get(
- self, kind: Literal["reaction_added", "reaction_removed"]
- ) -> ReactionPayload | None: ...
- @overload
- def get(self, kind: Literal["user_typing"]) -> TypingPayload | None: ...
- @overload
- def get(self, kind: Literal["presence_changed"]) -> PresenceChangedPayload | None: ...
- @overload
- def get(
- self,
- kind: Literal[
- "room_created",
- "room_updated",
- "room_deleted",
- "room_archived",
- "room_unarchived",
- "user_joined_room",
- "user_left_room",
- ],
- ) -> RoomEventPayload | None: ...
- @overload
- def get(
- self, kind: Literal["room_universal_changed"]
- ) -> RoomUniversalChangedPayload | None: ...
- @overload
- def get(self, kind: Literal["notification_created"]) -> NotificationCreatedPayload | None: ...
- @overload
- def get(
- self, kind: Literal["notification_dismissed"]
- ) -> NotificationDismissedPayload | None: ...
- @overload
- def get(
- self, kind: Literal["notification_level_changed"]
- ) -> NotificationLevelChangedPayload | None: ...
- @overload
- def get(self, kind: Literal["thread_follow_changed"]) -> ThreadFollowChangedPayload | None: ...
- @overload
- def get(self, kind: Literal["thread_created"]) -> ThreadCreatedPayload | None: ...
- @overload
- def get(self, kind: Literal["room_marked_as_read"]) -> RoomMarkedAsReadPayload | None: ...
- @overload
- def get(self, kind: Literal["server_updated"]) -> ServerUpdatedPayload | None: ...
- @overload
- def get(self, kind: Literal["user_profile_updated"]) -> UserProfileUpdatedPayload | None: ...
- @overload
- def get(self, kind: Literal["user_custom_status_set"]) -> UserCustomStatusSetPayload | None: ...
- @overload
- def get(
- self, kind: Literal["user_custom_status_cleared"]
- ) -> UserCustomStatusClearedPayload | None: ...
- @overload
- def get(
- self, kind: Literal["server_user_preferences_updated"]
- ) -> ServerUserPreferencesUpdatedPayload | None: ...
- @overload
- def get(self, kind: Literal["room_groups_updated"]) -> RoomGroupsUpdatedPayload | None: ...
- @overload
- def get(self, kind: Literal["server_member_deleted"]) -> ServerMemberDeletedPayload | None: ...
- @overload
- def get(
- self,
- kind: Literal[
- "asset_processing_started", "asset_processing_succeeded", "asset_processing_failed"
- ],
- ) -> AssetProcessingPayload | None: ...
- @overload
- def get(self, kind: Literal["asset_deleted"]) -> AssetDeletedPayload | None: ...
- @overload
- def get(
- self,
- kind: Literal[
- "call_started", "call_participant_joined", "call_participant_left", "call_ended"
- ],
- ) -> CallEventPayload | None: ...
- @overload
- def get(self, kind: Literal["mention_notification"]) -> MentionNotificationPayload | None: ...
- @overload
- def get(
- self, kind: Literal["new_direct_message_notification"]
- ) -> NewDirectMessageNotificationPayload | None: ...
- @overload
- def get(self, kind: Literal["session_terminated"]) -> SessionTerminatedPayload | None: ...
- @overload
- def get(self, kind: str) -> RealtimePayload | None: ...
- def get(self, kind: str) -> RealtimePayload | None:
- """Return ``payload`` if this event is of ``kind``, else ``None``."""
- return self.payload if self.kind == kind else None
- class RealtimeConnection:
- """Live realtime WebSocket session.
- Prefer :func:`stream_events` for the common case; use this class directly
- when you also need to send client pings, close cleanly, or inspect the
- negotiated :class:`ServerHello`.
- """
- def __init__(
- self,
- client: ChattoClient,
- *,
- protocol_version: int = REALTIME_PROTOCOL_VERSION,
- ) -> None:
- self._client = client
- self._protocol_version = protocol_version
- self._ws: Any = None
- self._server_hello: ServerHello | None = None
- @property
- def server_hello(self) -> ServerHello | None:
- return self._server_hello
- async def __aenter__(self) -> RealtimeConnection:
- await self.connect()
- return self
- async def __aexit__(self, *exc: Any) -> None:
- await self.close()
- async def connect(self) -> None:
- try:
- import websockets
- except ImportError as exc: # pragma: no cover
- raise ChattoError(
- "The realtime channel requires the 'websockets' package. "
- "Install with `pip install chattolib[realtime]`."
- ) from exc
- from chattolib._pb.chatto.realtime.v1 import realtime_pb2
- url = realtime_url(self._client.base_url)
- headers: dict[str, str] = {}
- if self._client.session_cookie:
- headers["Cookie"] = f"chatto_session={self._client.session_cookie}"
- self._ws = await websockets.connect(url, additional_headers=headers)
- client_hello = realtime_pb2.RealtimeClientFrame()
- client_hello.hello.protocol_version = self._protocol_version
- if self._client.token:
- client_hello.hello.bearer_token = self._client.token
- await self._ws.send(client_hello.SerializeToString())
- first = await self._ws.recv()
- first_frame = realtime_pb2.RealtimeServerFrame()
- first_frame.ParseFromString(first)
- if first_frame.WhichOneof("frame") != "hello":
- _raise_for_control(first_frame)
- raise ChattoRealtimeError(
- "unexpected_frame",
- f"expected server hello, got {first_frame.WhichOneof('frame')!r}",
- )
- hello = first_frame.hello
- self._server_hello = ServerHello(
- protocol_version=hello.protocol_version,
- server_version=hello.server_version,
- heartbeat_interval_seconds=hello.heartbeat_interval_seconds,
- capabilities=list(hello.capabilities),
- )
- subscribe = realtime_pb2.RealtimeClientFrame()
- subscribe.subscribe_events.SetInParent()
- await self._ws.send(subscribe.SerializeToString())
- async def close(self) -> None:
- if self._ws is None:
- return
- try:
- await self._ws.close()
- finally:
- self._ws = None
- async def ping(self, nonce: str = "") -> None:
- """Send a client ping. The server replies with a matching pong."""
- if self._ws is None:
- raise ChattoError("realtime connection is closed")
- from chattolib._pb.chatto.realtime.v1 import realtime_pb2
- frame = realtime_pb2.RealtimeClientFrame()
- frame.ping.nonce = nonce
- await self._ws.send(frame.SerializeToString())
- async def events(self) -> AsyncIterator[RealtimeEvent]:
- """Yield decoded live events until the connection closes."""
- if self._ws is None:
- raise ChattoError("realtime connection is not open")
- from chattolib._pb.chatto.realtime.v1 import realtime_pb2
- async for raw in self._ws:
- if isinstance(raw, str):
- # The protocol is binary; a text frame indicates a protocol violation.
- raise ChattoRealtimeError(
- "unexpected_text_frame",
- "server sent a text frame; realtime protocol expects binary",
- )
- frame = realtime_pb2.RealtimeServerFrame()
- frame.ParseFromString(raw)
- case = frame.WhichOneof("frame")
- if case == "event":
- yield _wrap_event(frame.event)
- elif case in ("heartbeat", "pong", "subscribed"):
- continue
- elif case == "error":
- err = frame.error
- exc = ChattoRealtimeError(err.code, err.message, fatal=err.fatal)
- if err.fatal:
- raise exc
- # Non-fatal errors are surfaced but the stream continues.
- # Callers can still receive them if desired; for now we log
- # by re-raising on fatal only.
- continue
- elif case == "close":
- close = frame.close
- raise ChattoRealtimeCloseError(
- close.code,
- close.message,
- reconnect=close.reconnect,
- retry_after_ms=close.retry_after_ms,
- )
- else:
- raise ChattoRealtimeError(
- "unexpected_frame",
- f"unknown server frame: {case!r}",
- )
- def _raise_for_control(frame: Any) -> None:
- """If ``frame`` is an error or close frame, translate it and raise."""
- case = frame.WhichOneof("frame")
- if case == "error":
- err = frame.error
- raise ChattoRealtimeError(err.code, err.message, fatal=err.fatal)
- if case == "close":
- close = frame.close
- raise ChattoRealtimeCloseError(
- close.code,
- close.message,
- reconnect=close.reconnect,
- retry_after_ms=close.retry_after_ms,
- )
- def _wrap_event(envelope: Any) -> RealtimeEvent:
- kind = envelope.WhichOneof("event") or ""
- pb_payload = getattr(envelope, kind, None) if kind else None
- payload = parse_payload(kind, pb_payload) if kind else None
- created_at = None
- if envelope.HasField("created_at"):
- created_at = parse_datetime(envelope.created_at.ToJsonString())
- actor_id: str | None = None
- if envelope.HasField("actor_id"):
- actor_id = envelope.actor_id
- return RealtimeEvent(
- id=envelope.id,
- created_at=created_at,
- actor_id=actor_id,
- kind=kind,
- payload=payload,
- raw=envelope,
- )
- async def stream_events(
- client: ChattoClient,
- *,
- protocol_version: int = REALTIME_PROTOCOL_VERSION,
- ) -> AsyncIterator[RealtimeEvent]:
- """Open a realtime connection and yield events until the server closes.
- Raises :class:`ChattoRealtimeCloseError` when the server sends a close frame,
- :class:`ChattoRealtimeError` on fatal protocol errors, or
- :class:`ChattoConnectError` if the initial HTTP handshake fails.
- """
- conn = RealtimeConnection(client, protocol_version=protocol_version)
- try:
- await conn.connect()
- async for event in conn.events():
- yield event
- finally:
- await conn.close()
- __all__ = [
- "REALTIME_PATH",
- "REALTIME_PROTOCOL_VERSION",
- "ChattoRealtimeCloseError",
- "ChattoRealtimeError",
- "RealtimeConnection",
- "RealtimeEvent",
- "RealtimePayload",
- "ServerHello",
- "realtime_url",
- "stream_events",
- ]
|