Started to implement good bot class
This commit is contained in:
270
new_bot.py
Normal file
270
new_bot.py
Normal file
@@ -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])
|
||||
Reference in New Issue
Block a user