import asyncio import aiofiles import aioconsole import traceback import logging import time import json import os import re from html.parser import HTMLParser from urllib.parse import urlparse from typing import Any, Callable, Awaitable from nio import AsyncClient, AsyncClientConfig, SyncResponse from nio import LoginResponse, LoginError, WhoamiResponse, WhoamiError from nio import RoomSendResponse, RoomSendError from nio import OlmUnverifiedDeviceError import nio.events from bot_types import * class MatrixBot: """Asynchronous Matrix Bot Implementation. Use objects of this class to build your bots. Manage the event loop by yourself. """ NEXT_BATCH_DUMP_PERIOD = 120.0 MATRIX_SYNC_PERIOD = 5000 # # PRIVATE # @staticmethod def _validate_matrix_homeserver_url(url: str) -> None: """Checks if `url` is a valid matrix homeserver URL. Raises an Exception if it is not. """ parsed = urlparse(url) if parsed.scheme not in ("http", "https"): raise RuntimeError( f"Scheme {parsed.scheme} is not a valid scheme for matrix homeserver URL" ) if not parsed.netloc: raise RuntimeError(f"{url} is not a valid matrix homeserver URL") if parsed.path != "": raise RuntimeError(f"{url} must have empty path (remove `{parsed.path}` after the hostname)") @staticmethod def _validate_matrix_username_localpart(localpart: str) -> None: """Checks if `username` is a valid localpart of matrix username. Raises an Exception if it is not. """ pattern = r"^[a-z0-9._=\-]+$" if not bool(re.fullmatch(pattern, localpart)) or len(localpart) > 255: raise RuntimeError(f"{localpart} is not a valid matrix username localpart") @staticmethod def _validate_storage_directory(path: Path) -> None: """Checks if `path` is a valid storage directory and creates it. Raises an Exception if it is not. """ path.mkdir(parents=True, exist_ok=True) if not path.is_dir(): raise RuntimeError(f"Could not create directory {path}") @staticmethod def _validate_bot_config(config: MatrixBotConfig) -> None: """Checks if `config` has errors. Raises an Exception if it does. """ MatrixBot._validate_matrix_homeserver_url(config.matrix_homeserver_url) MatrixBot._validate_matrix_username_localpart(config.matrix_username_localpart) MatrixBot._validate_storage_directory(config.storage_directory) @staticmethod def _build_client(config: MatrixBotConfig) -> AsyncClient: """Builds `nio.AsyncClient` from `MatrixBotConfig`""" # create the config for the client client_config = AsyncClientConfig( store_name="nio_store_file", encryption_enabled=True, store_sync_tokens=False ) # create the client client = AsyncClient( homeserver=config.matrix_homeserver_url, user=config.matrix_username_localpart, store_path=str(config.storage_directory), config=client_config ) return client @staticmethod def _build_matrix_username(config: MatrixBotConfig) -> str: """Builds full matrix username.""" homeserver_name = urlparse(config.matrix_homeserver_url).hostname localpart = config.matrix_username_localpart return f"@{localpart}:{homeserver_name}" @staticmethod def _check_and_remove_html(text_to_check: str) -> tuple[bool, str]: """Checks if `text_to_check` is HTML and sanitizes it. Returns: tuple[bool, str] where the first element is True if `text_to_check` contains valid HTML, and the second element is text without HTML (or just copy of `text_to_check` if it does not contain HTML) """ has_tags = False text_fragments = [] class Extractor(HTMLParser): def handle_starttag(self, tag, attrs): nonlocal has_tags has_tags = True def handle_data(self, data): text_fragments.append(data) parser = Extractor(convert_charrefs=True) parser.feed(text_to_check) try: if has_tags: return (True, " ".join("".join(text_fragments).split())) except: traceback.print_exc() return (False, text_to_check) @staticmethod async def _wait_for_task_and_stop_event(payload_task: asyncio.Task, stop_wait_task: asyncio.Task, print_exc: bool = False) -> Any: """Waits for `payload_task` to finish. Uses `stop_wait_task` to check if payload task should be cancelled. Returns value returned by `payload_task` task. Raises exception raised by `payload_task` task. If `stop_wait_task` finishes, then `asyncio.CancelledError` is raised. """ try: done, _ = await asyncio.wait( [payload_task, stop_wait_task], return_when=asyncio.FIRST_COMPLETED ) if stop_wait_task in done: payload_task.cancel() await payload_task raise asyncio.CancelledError() return payload_task.result() except asyncio.CancelledError: raise except: if print_exc: traceback.print_exc() raise @staticmethod async def _default_password_callback() -> str: """Gets password from `MATRIX_PASSWORD` envvar if it is set. Asks the user for the password otherwise.""" if "MATRIX_PASSWORD" in os.environ: return os.environ["MATRIX_PASSWORD"] print("--- A password is required (btw you might use MATRIX_PASSWORD envvar) ---") return await aioconsole.ainput("Password: ") @staticmethod async def _debug_event_callback(*args, **kwargs) -> None: """Just prints types of arguments""" try: print(f"_debug_event_callback ({len(args)} args, {len(kwargs)} kwargs)") for a in args: print(f" - {type(a)}") for k in kwargs: print(f" * {k} = {kwargs[k]}") except: traceback.print_exc() async def _write_next_batch(self, next_batch: str) -> None: """Writes `next_batch` value to disk.""" path = self._config.storage_directory / "next_batch" async with aiofiles.open(path, "w") as f: await f.write(next_batch) self._logger.debug("next_batch value is written to the disk") async def _read_next_batch(self) -> str | None: """Reads `next_batch` value from disk. Returns None if file does not exist.""" path = self._config.storage_directory / "next_batch" if not path.is_file(): return None async with aiofiles.open(path, "r") as f: return (await f.read()).strip() async def _write_session_data(self, *, access_token: str, device_id: str) -> None: """Write session data to disk.""" path = self._config.storage_directory / "session_data.json" data = { "access_token": access_token, "device_id": device_id } async with aiofiles.open(path, "w") as f: await f.write(json.dumps(data, indent=4)) self._logger.debug("Session data is writter to the disk") async def _read_session_data(self) -> dict[str, Any] | None: """Read session data from disk.""" path = self._config.storage_directory / "session_data.json" if not path.is_file(): return None async with aiofiles.open(path, "r") as f: j = json.loads(await f.read()) return j # # CALLBACKS # async def _callback_sync(self, response: SyncResponse) -> None: """This callback is called when AsyncClient syncs with the server""" current_time = time.time() delta_time = current_time - self._last_next_batch_dump self._last_next_batch = response.next_batch if delta_time >= self.NEXT_BATCH_DUMP_PERIOD: self._last_next_batch_dump = current_time try: await self._write_next_batch(self._last_next_batch) except: traceback.print_exc() # # LIFECYCLE # def _setup_client_callbacks(self) -> None: # setup the callbacks self._client.add_response_callback(self._callback_sync, SyncResponse) # type: ignore async def _client_login_session_data(self, session_data: dict[str, Any]) -> None: """Login using session data. Raises and exception on failure.""" self._logger.debug("Using stored session data to log in") # build user id username = self._build_matrix_username(self._config) self._client.restore_login( user_id=username, **session_data ) result = await self._client.whoami() if type(result) is WhoamiError: self._logger.error(f"Can't log in using stored session data: '{result.message}'") raise RuntimeError(result.message) elif type(result) is not WhoamiResponse: self._logger.error("Can't log in using stored session data, unknown error") raise RuntimeError("Unknown response for whoami request") self._logger.debug("Logged in using stored session data") async def _client_login_password(self) -> None: """Login using password and save result to disk on success. Raises an exception on failure. """ self._logger.debug("Using password to log in") # get the password password = await self._cb_password() result = await self._client.login(password=password) if type(result) is LoginResponse: self._logger.debug("Logged in using password") await self._write_session_data( access_token=result.access_token, device_id=result.device_id ) elif type(result) is LoginError: self._logger.error(f"Can't log in using password: '{result.message}'") raise RuntimeError(result.message) else: self._logger.error(f"Can't log in using password, unknown error") raise RuntimeError("Unknown login result") async def _client_login(self) -> None: """This function logs in.""" # check if we have session data stored on the disk session_data = await self._read_session_data() # session data is present, try to log in if session_data is not None: self._logger.debug("Some session data found on the disk") try: await self._client_login_session_data(session_data) return except asyncio.CancelledError: raise except: pass # no session data - login using password try: self._logger.debug("No session data found on the disk OR invalid data") await self._client_login_password() return except asyncio.CancelledError: raise except: pass # can't login self._logger.error("Can't log in using available methods") raise RuntimeError("All login methods have failed, can't continue") async def _client_destroy(self) -> None: """Gracefully destroys the client.""" try: self._logger.debug("Closing the client") await self._client.close() if self._last_next_batch is not None: self._logger.debug("Saving next_batch") await self._write_next_batch(self._last_next_batch) except: traceback.print_exc() async def _client_cancellable_sync_forever(self, *args, **kwargs) -> Any: """Behaves exactly like AsyncClient.sync_forever, but supports task cancellation""" sync_forever_task = asyncio.create_task( self._client.sync_forever(*args, **kwargs) ) try: return await sync_forever_task except asyncio.CancelledError: try: self._client.stop_sync_forever() except Exception: pass finally: sync_forever_task.cancel() await asyncio.gather(sync_forever_task, return_exceptions=True) raise async def _background_coroutine(self) -> None: """This function implements bot lifecycle.""" # we should stop when this task stops self._logger.debug("_background_coroutine is started") stop_wait_task = asyncio.create_task(self._stop_event.wait()) # perform login login_task = asyncio.create_task(self._client_login()) try: await self._wait_for_task_and_stop_event(login_task, stop_wait_task) except asyncio.CancelledError: self._logger.debug("Background task is cancelled during login") await self._client_destroy() return except: traceback.print_exc() self._logger.info("Succesfully logged in") # sync forever self._logger.info("Syncing forever") sync_task = asyncio.create_task( self._client_cancellable_sync_forever( timeout=self.MATRIX_SYNC_PERIOD, since=(await self._read_next_batch()) ) ) try: await self._wait_for_task_and_stop_event(sync_task, stop_wait_task) except asyncio.CancelledError: self._logger.debug("Sync task is cancelled") await self._client_destroy() return except: traceback.print_exc() # # PUBLIC # def __init__(self, config: MatrixBotConfig) -> None: # check if config is valid self._validate_bot_config(config) # may raise an Exception # save the config self._config: MatrixBotConfig = config # create the logger self._logger = logging.getLogger(self._build_matrix_username(config)) self._logger.setLevel(logging.DEBUG) # create the client self._client: AsyncClient = self._build_client(self._config) self._setup_client_callbacks() # prepare some private data self._background_task: asyncio.Task | None = None self._last_next_batch_dump: float = 0.0 self._last_next_batch: str | None = None self._cb_password = self._default_password_callback def start(self) -> None: """Start the bot. Starts the bot in background task. Raises an exception if there are problems (for example, the bot is already started). The bot will do everything to keep itself running, including restarts. Use `stop()` to stop the bot. """ if self._background_task is not None: raise RuntimeError("The bot is already started!") self._stop_event = asyncio.Event() self._background_task = asyncio.create_task( self._background_coroutine() ) async def stop(self) -> None: """Stop the bot and wait for the bot stop.""" if self._background_task is None: return self._stop_event.set() try: await self._background_task except: traceback.print_exc() self._stop_event = None self._background_task = None def verify_all_known_devices(self) -> bool: """Verifies all known devices. Returns: True if there were unverified devices that are verified now. """ result = False for user_id in self._client.device_store.users: for device_id, olm_device in self._client.device_store[user_id].items(): # can't trust ourselves if device_id == self._client.device_id and user_id == self._client.user_id: continue # they are already verified if olm_device.verified: continue # verify them self._client.verify_device(olm_device) result = True return result def add_event_callback(self, callback: Callable[[Any], Awaitable[None]] | None, event_class: nio.events.Event) -> None: """Added event callback for events of specified class. Use `None` instead of callback to print parameter types you need to use in your callback.""" if callback is None: callback = self._debug_event_callback self._client.add_event_callback(callback, event_class) # type: ignore def get_client(self) -> AsyncClient: """Get AsyncClient in use""" return self._client async def send_text_to_room(self, room_id: str, text: str, is_html: bool | None = None, **kwargs) -> str: """Sends a text message to the room and handle HTML as specified. Args: text - text to send is_html - True if text is HTML; False if text is not HTML; None if the value should be guessed kwargs - passed as `m.room.message` content keys Returns: event_id of the message on success. Raises an exception on error. """ text_with_html = text # guess if text is HTML if is_html is None: is_html, text = self._check_and_remove_html(text) # text is HTML elif is_html: _, text = self._check_and_remove_html(text) # create `content` for `room_send()` content = { "msgtype": "m.text", "body": text, **kwargs } if is_html: content["format"] = "org.matrix.custom.html" content["formatted_body"] = content["body"] # try to send the message try: result = await self._client.room_send( room_id=room_id, message_type="m.room.message", content=content ) except OlmUnverifiedDeviceError: if self._config.auto_verify_all_known_devices: if not self.verify_all_known_devices(): raise return await self.send_text_to_room(room_id, text, is_html, **kwargs) else: raise # success if type(result) is RoomSendResponse: return result.event_id # error elif type(result) is RoomSendError: raise RuntimeError(result) # unknown error else: raise RuntimeError("Unknown error has occured", result)