diff --git a/bot_types.py b/bot_types.py new file mode 100644 index 0000000..5dba236 --- /dev/null +++ b/bot_types.py @@ -0,0 +1,17 @@ +"""Data types required for the bot""" + +from pathlib import Path +from dataclasses import dataclass + +@dataclass +class MatrixBotConfig: + """Configuration for MatrixBot""" + + matrix_homeserver_url: str + """Homeserver, for example: `https://matrix.domain.su`""" + + matrix_username_localpart: str + """Localpart of matrix username (without homeserver), for example: `valid-username`""" + + storage_directory: Path + """Path to the storage directory (will be created if needed)""" \ No newline at end of file diff --git a/new_bot.py b/new_bot.py new file mode 100644 index 0000000..880f562 --- /dev/null +++ b/new_bot.py @@ -0,0 +1,270 @@ +import asyncio +import aiofiles +import aioconsole +import traceback +import time +import json +import os +import re +from urllib.parse import urlparse +from typing import Any + +from nio import AsyncClient, AsyncClientConfig, SyncResponse +from nio import LoginResponse, LoginError, WhoamiResponse, WhoamiError + +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 + + # + # 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 + 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: ") + + + 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) + + 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)) + + 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 + if delta_time >= self.NEXT_BATCH_DUMP_PERIOD: + self._last_next_batch_dump = current_time + try: + await self._write_next_batch(response.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.""" + # 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: + raise RuntimeError(result.message) + elif type(result) is not WhoamiResponse: + raise RuntimeError("Unknown response for whoami request") + + async def _client_login_password(self) -> None: + """Login using password and save result to disk on success. + Raises an exception on failure. + """ + # get the password + password = await self._cb_password() + result = await self._client.login(password=password) + if type(result) is LoginResponse: + await self._write_session_data( + access_token=result.access_token, + device_id=result.device_id + ) + elif type(result) is LoginError: + raise RuntimeError(result.message) + else: + 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: + try: + await self._client_login_session_data(session_data) + return + except: + pass + # no session data - login using password + try: + await self._client_login_password() + return + except: + pass + # can't login + raise RuntimeError("All login methods have failed, can't continue") + + async def _background_coroutine(self) -> None: + """This function implements bot lifecycle.""" + # we should stop when this task stops + stop_wait_task = asyncio.create_task(self._stop_event.wait()) + self._client = self._build_client(self._config) + self._setup_client_callbacks() + # perform login + try: + await self._client_login() + 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 = config + + # prepare some private data + self._background_task: asyncio.Task | None = None + self._client: AsyncClient = None # type: ignore + self._last_next_batch_dump: float = 0.0 + 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() + ) + + def stop(self) -> bool: + """Stop the bot. + Signals the background task to stop. Returns True on success (the + signal is sent), False on failure (no background task). This + function won't wait for the task to stop. Use `wait_stop` to wait. + """ + if self._background_task is None: + return False + self._stop_event.set() + return True + + async def wait_stop(self) -> None: + """Wait for the background task to stop. + You must call `stop()` by yourself. This function will never return + otherwise. It just waits for the background task stop. + """ + if self._background_task is None: + return + await asyncio.wait([self._background_task]) \ No newline at end of file