from __future__ import annotations import asyncio import json import logging import websockets from websockets.asyncio.client import ClientConnection from websockets.exceptions import ConnectionClosed from yuxi.channel.extensions.mattermost.client import MattermostClient from yuxi.channel.extensions.mattermost.errors import MattermostAuthError from yuxi.channel.extensions.mattermost.types import MattermostPost logger = logging.getLogger(__name__) PING_INTERVAL_SECONDS = 30 PONG_TIMEOUT_SECONDS = 10 HEALTH_CHECK_INTERVAL_SECONDS = 30 class MattermostWebSocketMonitor: def __init__( self, client: MattermostClient, account_id: str = "", on_posted=None, on_reaction_added=None, on_reaction_removed=None, on_post_edited=None, on_post_deleted=None, on_user_added=None, on_user_removed=None, on_user_updated=None, on_channel_created=None, on_channel_deleted=None, on_channel_updated=None, on_typing=None, on_status_change=None, on_thread_updated=None, ): self.client = client self.account_id = account_id self._on_posted = on_posted self._on_reaction_added = on_reaction_added self._on_reaction_removed = on_reaction_removed self._on_post_edited = on_post_edited self._on_post_deleted = on_post_deleted self._on_user_added = on_user_added self._on_user_removed = on_user_removed self._on_user_updated = on_user_updated self._on_channel_created = on_channel_created self._on_channel_deleted = on_channel_deleted self._on_channel_updated = on_channel_updated self._on_typing = on_typing self._on_status_change = on_status_change self._on_thread_updated = on_thread_updated self._ws: ClientConnection | None = None self._abort = asyncio.Event() self._connected = asyncio.Event() self.bot_user_id: str = "" self.bot_last_update_at: int = 0 @property def is_connected(self) -> bool: return self._ws is not None and not self._abort.is_set() async def connect(self) -> None: self._abort.clear() self._connected.clear() ws_url = self.client.base_url.replace("http", "ws") + "/api/v4/websocket" logger.info("Mattermost WS connecting to %s", ws_url) self._ws = await websockets.connect(ws_url, ping_interval=None) await self._authenticate() logger.info("Mattermost WS authenticated for account %s", self.account_id) self._connected.set() try: await asyncio.gather( self._event_loop(), self._ping_loop(), self._health_check_loop(), ) except asyncio.CancelledError: pass except ConnectionClosed: logger.warning("Mattermost WS connection closed for account %s", self.account_id) finally: await self._cleanup() async def disconnect(self) -> None: self._abort.set() await self._cleanup() async def _authenticate(self) -> None: auth_msg = { "seq": 1, "action": "authentication_challenge", "data": {"token": self.client.bot_token}, } await self._ws.send(json.dumps(auth_msg)) response = await asyncio.wait_for(self._ws.recv(), timeout=10) data = json.loads(response) if data.get("status") != "OK": raise MattermostAuthError(401, "WebSocket authentication failed") async def _event_loop(self) -> None: while not self._abort.is_set() and self._ws: try: message = await asyncio.wait_for(self._ws.recv(), timeout=60) await self._handle_message(message) except TimeoutError: continue except ConnectionClosed: break async def _ping_loop(self) -> None: while not self._abort.is_set() and self._ws: try: await asyncio.sleep(PING_INTERVAL_SECONDS) await self._ws.send(json.dumps({"seq": 1, "action": "ping"})) try: pong = await asyncio.wait_for(self._ws.recv(), timeout=PONG_TIMEOUT_SECONDS) data = json.loads(pong) if data.get("event") != "pong" and data.get("status") != "OK": logger.warning("Unexpected pong response: %s", data) except TimeoutError: logger.warning("Pong timeout for account %s", self.account_id) break except ConnectionClosed: break async def _health_check_loop(self) -> None: while not self._abort.is_set(): await asyncio.sleep(HEALTH_CHECK_INTERVAL_SECONDS) try: me = await self.client.fetch_me() update_at = me.get("update_at", 0) if self.bot_last_update_at and update_at != self.bot_last_update_at: logger.warning( "Bot update_at changed (%d -> %d), reconnecting", self.bot_last_update_at, update_at, ) break self.bot_last_update_at = update_at except Exception: pass async def _handle_message(self, raw: str) -> None: try: data = json.loads(raw) except json.JSONDecodeError: return event = data.get("event", "") if event == "hello": return event_data = data.get("data", {}) if event == "posted": post_data = event_data.get("post", "") if post_data and self._on_posted: post = _parse_post(post_data) if post: await self._on_posted(post, event_data) elif event == "reaction_added" and self._on_reaction_added: await self._on_reaction_added(event_data) elif event == "reaction_removed" and self._on_reaction_removed: await self._on_reaction_removed(event_data) elif event == "post_edited" and self._on_post_edited: post_data = event_data.get("post", "") post = _parse_post(post_data) if post: await self._on_post_edited(post, event_data) elif event == "post_deleted" and self._on_post_deleted: await self._on_post_deleted(event_data) elif event == "user_added" and self._on_user_added: await self._on_user_added(event_data) elif event == "user_removed" and self._on_user_removed: await self._on_user_removed(event_data) elif event == "user_updated" and self._on_user_updated: await self._on_user_updated(event_data) elif event == "channel_created" and self._on_channel_created: await self._on_channel_created(event_data) elif event == "channel_deleted" and self._on_channel_deleted: await self._on_channel_deleted(event_data) elif event == "channel_updated" and self._on_channel_updated: await self._on_channel_updated(event_data) elif event == "typing" and self._on_typing: await self._on_typing(event_data) elif event == "status_change" and self._on_status_change: await self._on_status_change(event_data) elif event == "thread_updated" and self._on_thread_updated: await self._on_thread_updated(event_data) async def _cleanup(self) -> None: if self._ws: try: await self._ws.close() except Exception: pass self._ws = None self._connected.clear() def _parse_post(post_data: str | dict) -> MattermostPost | None: if isinstance(post_data, str): try: post_data = json.loads(post_data) except json.JSONDecodeError: return None if not isinstance(post_data, dict): return None return MattermostPost.from_dict(post_data)