diff --git a/bot.py b/bot.py index adae7f1..728f6a9 100644 --- a/bot.py +++ b/bot.py @@ -1,202 +1,512 @@ -""" -matrix-nio basics wrapper -""" - import asyncio -import time +import aiofiles +import aioconsole import traceback -from pathlib import Path +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 LoginError, LoginResponse, SyncResponse -from nio import WhoamiResponse -from nio import AsyncClient, AsyncClientConfig +from nio import AsyncClient, AsyncClientConfig, SyncResponse +from nio import LoginResponse, LoginError, WhoamiResponse, WhoamiError -from datatypes import AppConfig -import util +from nio import RoomSendResponse, RoomSendError + +from nio import OlmUnverifiedDeviceError + +from nio import MatrixInvitedRoom, InviteMemberEvent +from nio import JoinResponse + +import nio.events + +from bot_types import * +class MatrixBot: + """Asynchronous Matrix Bot Implementation. -# -# DATA -# -_app_config: AppConfig -_client: AsyncClient -_task: asyncio.Task | None = None -_stop: asyncio.Event | None = None -_since: str | None = None -_last_since_save_time: float = 0 + 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") -# -# CALLBACKS -# -async def _sync_callback(response: SyncResponse): - global _since, _last_since_save_time - t = time.time() - _since = response.next_batch - if t - _last_since_save_time >= 120.0: - await util.set_next_batch(_app_config, _since) - _last_since_save_time = t + @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) - -# -# PRIVATE -# -async def _bot_login_using_token(token: str, device_id: str) -> bool: - """Tries to login using access_token. Returns True on success.""" - util.log_info("Authorizing using access_token...") - _client.restore_login( - user_id=f"@{_app_config.matrix_user}:{util.get_hostname_from_url(_app_config.matrix_homeserver)}", - device_id=device_id, - access_token=token - ) - result = await _client.whoami() - if type(result) is not WhoamiResponse: - return False - util.log_info(f"Logged in as {result.user_id} using access_token") - return True - -async def _bot_login_using_password(password: str) -> tuple[str, str] | tuple[None, None]: - """Tries to login using password. Returns (access_token, device_id) on success.""" - util.log_info("Authorizing using password...") - result = await _client.login(password=password) - if type(result) is LoginResponse: - util.log_info(f"Authorized using password") - return result.access_token, result.device_id - elif type(result) is LoginError: - util.log_error(f"Failed to authorize using password: {result.message}") - return None, None - else: - raise RuntimeError(f"Invalid login result: {result}") - -async def _bot_login(config: AppConfig) -> bool: - """Tries to login""" - try: - # get the session token and try to use it - session_token, device_id = await util.get_session_data(config) - if session_token is not None and device_id is not None: - if await _bot_login_using_token(session_token, device_id): - return True - await util.set_session_data(config, None) - util.log_warning("Existing access_token is deleted") - # get the password and try to use it - password = await util.get_password() - if password is None: - util.log_error("No password provided (consider using MATRIX_PASSWORD environment variable)") - return False - access_token, device_id = await _bot_login_using_password(password) - if access_token is not None and device_id is not None: - await util.set_session_data(config, (access_token, device_id)) - util.log_warning("Saved new access_token and device_id") - return True - # can't login - return False - except asyncio.CancelledError: - return False - except: - traceback.print_exc() - return False - -async def _bot_loop(config: AppConfig) -> None: - """Bot loop""" - global _client, _since - # app stop task - if _stop is None: - raise RuntimeError("_stop can't be None") - stop_task = asyncio.create_task(_stop.wait()) - # setup the callback for syncing - _client.add_response_callback(_sync_callback, SyncResponse) # type: ignore - # login - login_task = asyncio.create_task(_bot_login(config)) - done, _ = await asyncio.wait( - [login_task, stop_task], - return_when=asyncio.FIRST_COMPLETED - ) - # stopped - if stop_task in done: - login_task.cancel() - return - # failed to login - if login_task.exception() or not login_task.result(): - util.request_app_stop("Can't authorize into matrix") - return - # load initial `next_batch` - _since = await util.get_next_batch(config) - # sync forever - while True: - # sync - sync_task = asyncio.create_task(_client.sync_forever(timeout=5000, since=_since)) - done, _ = await asyncio.wait( - [sync_task, stop_task], - return_when=asyncio.FIRST_COMPLETED + @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 ) - # stopped - if stop_task in done: - util.log_info("Stopping sync_forever...") - _client.stop_sync_forever() - util.log_info("Waiting for sync_forever to quit...") - await sync_task - sync_task.cancel() - break - # something happened + # 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: - sync_task.result() + 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() + + async def _callback_autojoin(self, room: MatrixInvitedRoom, event: InviteMemberEvent): + try: + # event.state_key must be our username + if event.state_key != self._client.user_id: + return + # membership status must be invite + if event.membership != "invite": + return + result = await self._client.join(room.room_id) + if type(result) is JoinResponse: + self._logger.info(f"Autojoined the room {room.room_id}") + else: + self._logger.error(f"Can't autojoin the room {room.room_id}") + except: + self._logger.error(traceback.format_exc()) + + # + # LIFECYCLE + # + def _setup_client_callbacks(self) -> None: + """Setup internal client callbacks""" + self._client.add_response_callback(self._callback_sync, SyncResponse) # type: ignore + + if self._config.auto_join_any_room_on_invite: + self._client.add_event_callback(self._callback_autojoin, InviteMemberEvent) # 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() - await asyncio.sleep(1) + # + # 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() -# -# PUBLIC -# -async def start(config: AppConfig) -> bool: - """Starts the bot""" - global _client - global _task, _stop, _app_config - if _task is not None: - return False - _app_config = config - # create the bot - store_dir = Path.cwd() / config.store_dir - store_dir.mkdir(parents=True, exist_ok=True) - client_config = AsyncClientConfig( - store_name="storefile", - encryption_enabled=True, - store_sync_tokens=False - ) - _client = AsyncClient( - homeserver=config.matrix_homeserver, - user=config.matrix_user, - store_path=str(store_dir), - config=client_config - ) - _stop = asyncio.Event() - _task = asyncio.create_task(_bot_loop(config)) - return True + # 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 get_client() -> AsyncClient: - return _client + 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() -> None: - """Stop the bot""" - global _task, _stop - if _task is None or _stop is None: - return - _stop.set() - try: - await _task - except asyncio.CancelledError: - pass - except: - traceback.print_exc() - _task = None - _stop = None - try: - await _client.close() - except: - pass - await util.set_next_batch(_app_config, _since) + 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) \ No newline at end of file diff --git a/main.py b/main.py index e1155e0..a5c3151 100644 --- a/main.py +++ b/main.py @@ -11,7 +11,7 @@ from pathlib import Path import config import util -from new_bot import MatrixBot +from bot import MatrixBot from bot_types import MatrixBotConfig import nio.events diff --git a/new_bot.py b/new_bot.py deleted file mode 100644 index 728f6a9..0000000 --- a/new_bot.py +++ /dev/null @@ -1,512 +0,0 @@ -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 - -from nio import MatrixInvitedRoom, InviteMemberEvent -from nio import JoinResponse - -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() - - async def _callback_autojoin(self, room: MatrixInvitedRoom, event: InviteMemberEvent): - try: - # event.state_key must be our username - if event.state_key != self._client.user_id: - return - # membership status must be invite - if event.membership != "invite": - return - result = await self._client.join(room.room_id) - if type(result) is JoinResponse: - self._logger.info(f"Autojoined the room {room.room_id}") - else: - self._logger.error(f"Can't autojoin the room {room.room_id}") - except: - self._logger.error(traceback.format_exc()) - - # - # LIFECYCLE - # - def _setup_client_callbacks(self) -> None: - """Setup internal client callbacks""" - self._client.add_response_callback(self._callback_sync, SyncResponse) # type: ignore - - if self._config.auto_join_any_room_on_invite: - self._client.add_event_callback(self._callback_autojoin, InviteMemberEvent) # 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) \ No newline at end of file