Compare commits
2 Commits
4a12bdbd40
...
cee1275e3c
| Author | SHA1 | Date | |
|---|---|---|---|
| cee1275e3c | |||
| 9b702df01f |
670
bot.py
670
bot.py
@@ -1,202 +1,512 @@
|
|||||||
"""
|
|
||||||
matrix-nio basics wrapper
|
|
||||||
"""
|
|
||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
import time
|
import aiofiles
|
||||||
|
import aioconsole
|
||||||
import traceback
|
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 AsyncClient, AsyncClientConfig, SyncResponse
|
||||||
from nio import WhoamiResponse
|
from nio import LoginResponse, LoginError, WhoamiResponse, WhoamiError
|
||||||
from nio import AsyncClient, AsyncClientConfig
|
|
||||||
|
|
||||||
from datatypes import AppConfig
|
from nio import RoomSendResponse, RoomSendError
|
||||||
import util
|
|
||||||
|
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
|
||||||
# DATA
|
by yourself.
|
||||||
#
|
"""
|
||||||
_app_config: AppConfig
|
NEXT_BATCH_DUMP_PERIOD = 120.0
|
||||||
_client: AsyncClient
|
MATRIX_SYNC_PERIOD = 5000
|
||||||
_task: asyncio.Task | None = None
|
|
||||||
_stop: asyncio.Event | None = None
|
|
||||||
_since: str | None = None
|
|
||||||
_last_since_save_time: float = 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
|
||||||
# CALLBACKS
|
def _validate_storage_directory(path: Path) -> None:
|
||||||
#
|
"""Checks if `path` is a valid storage directory and creates it.
|
||||||
async def _sync_callback(response: SyncResponse):
|
Raises an Exception if it is not.
|
||||||
global _since, _last_since_save_time
|
"""
|
||||||
t = time.time()
|
path.mkdir(parents=True, exist_ok=True)
|
||||||
_since = response.next_batch
|
if not path.is_dir():
|
||||||
if t - _last_since_save_time >= 120.0:
|
raise RuntimeError(f"Could not create directory {path}")
|
||||||
await util.set_next_batch(_app_config, _since)
|
|
||||||
_last_since_save_time = t
|
|
||||||
|
|
||||||
|
@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:
|
||||||
# PRIVATE
|
"""Builds `nio.AsyncClient` from `MatrixBotConfig`"""
|
||||||
#
|
# create the config for the client
|
||||||
async def _bot_login_using_token(token: str, device_id: str) -> bool:
|
client_config = AsyncClientConfig(
|
||||||
"""Tries to login using access_token. Returns True on success."""
|
store_name="nio_store_file",
|
||||||
util.log_info("Authorizing using access_token...")
|
encryption_enabled=True,
|
||||||
_client.restore_login(
|
store_sync_tokens=False
|
||||||
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
|
|
||||||
)
|
)
|
||||||
# stopped
|
# create the client
|
||||||
if stop_task in done:
|
client = AsyncClient(
|
||||||
util.log_info("Stopping sync_forever...")
|
homeserver=config.matrix_homeserver_url,
|
||||||
_client.stop_sync_forever()
|
user=config.matrix_username_localpart,
|
||||||
util.log_info("Waiting for sync_forever to quit...")
|
store_path=str(config.storage_directory),
|
||||||
await sync_task
|
config=client_config
|
||||||
sync_task.cancel()
|
)
|
||||||
break
|
return client
|
||||||
# something happened
|
|
||||||
|
@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:
|
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:
|
except:
|
||||||
traceback.print_exc()
|
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()
|
||||||
|
|
||||||
#
|
# prepare some private data
|
||||||
# PUBLIC
|
self._background_task: asyncio.Task | None = None
|
||||||
#
|
self._last_next_batch_dump: float = 0.0
|
||||||
async def start(config: AppConfig) -> bool:
|
self._last_next_batch: str | None = None
|
||||||
"""Starts the bot"""
|
self._cb_password = self._default_password_callback
|
||||||
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
|
|
||||||
|
|
||||||
def get_client() -> AsyncClient:
|
def start(self) -> None:
|
||||||
return _client
|
"""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:
|
async def stop(self) -> None:
|
||||||
"""Stop the bot"""
|
"""Stop the bot and wait for the bot stop."""
|
||||||
global _task, _stop
|
if self._background_task is None:
|
||||||
if _task is None or _stop is None:
|
return
|
||||||
return
|
self._stop_event.set()
|
||||||
_stop.set()
|
try:
|
||||||
try:
|
await self._background_task
|
||||||
await _task
|
except:
|
||||||
except asyncio.CancelledError:
|
traceback.print_exc()
|
||||||
pass
|
self._stop_event = None
|
||||||
except:
|
self._background_task = None
|
||||||
traceback.print_exc()
|
|
||||||
_task = None
|
|
||||||
_stop = None
|
|
||||||
try:
|
|
||||||
await _client.close()
|
|
||||||
except:
|
|
||||||
pass
|
|
||||||
await util.set_next_batch(_app_config, _since)
|
|
||||||
|
|
||||||
|
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)
|
||||||
148
logic.py
148
logic.py
@@ -1,148 +0,0 @@
|
|||||||
"""Main logic implementation"""
|
|
||||||
|
|
||||||
import traceback
|
|
||||||
from typing import Any
|
|
||||||
import html
|
|
||||||
|
|
||||||
from nio import AsyncClient
|
|
||||||
from nio import JoinResponse, RoomSendResponse, RoomSendError
|
|
||||||
from nio import MatrixInvitedRoom, InviteMemberEvent
|
|
||||||
from nio import MatrixRoom, RoomMessageText
|
|
||||||
|
|
||||||
from nio import OlmUnverifiedDeviceError
|
|
||||||
|
|
||||||
from datatypes import *
|
|
||||||
import util
|
|
||||||
|
|
||||||
|
|
||||||
#
|
|
||||||
# DATA
|
|
||||||
#
|
|
||||||
_client: AsyncClient
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
#
|
|
||||||
# PRIVATE
|
|
||||||
#
|
|
||||||
async def _verify_all_devices() -> None:
|
|
||||||
"""Verifies all known devices"""
|
|
||||||
for user_id in _client.device_store.users:
|
|
||||||
for device_id, olm_device in _client.device_store[user_id].items():
|
|
||||||
# can't trust ourselves
|
|
||||||
if device_id == _client.device_id and user_id == _client.user_id:
|
|
||||||
continue
|
|
||||||
# they are already verified
|
|
||||||
if olm_device.verified:
|
|
||||||
continue
|
|
||||||
# verify them
|
|
||||||
_client.verify_device(olm_device)
|
|
||||||
|
|
||||||
def _handle_html_in_kwargs(kwargs: dict[str, Any]) -> None:
|
|
||||||
"""Modified `kwargs` in-place so that `formatted_body` appears if needed"""
|
|
||||||
if "formatted_body" in kwargs or "body" not in kwargs:
|
|
||||||
return
|
|
||||||
is_html, text_without_html = util.check_and_remove_html(kwargs["body"])
|
|
||||||
if not is_html:
|
|
||||||
return
|
|
||||||
kwargs["format"] = "org.matrix.custom.html"
|
|
||||||
kwargs["formatted_body"] = kwargs["body"]
|
|
||||||
kwargs["body"] = text_without_html
|
|
||||||
|
|
||||||
async def _send_message_to(room_id: str, message_type: MessageType, **kwargs) -> str:
|
|
||||||
"""Sends a message to the room and returns event_id.
|
|
||||||
|
|
||||||
This function automatically detects `body` key in `kwargs` and checks
|
|
||||||
if it is a valid HTML. If it is a valid HTML, it will send it as such.
|
|
||||||
Moreover, `body` attribute will be cleaned from any HTML tags, so that
|
|
||||||
the text will be looking well. `formatted_body` attribute is added
|
|
||||||
automatically and you should not add it manually.
|
|
||||||
"""
|
|
||||||
try:
|
|
||||||
# handle HTML
|
|
||||||
_handle_html_in_kwargs(kwargs)
|
|
||||||
# try to send the message
|
|
||||||
result = await _client.room_send(
|
|
||||||
room_id=room_id,
|
|
||||||
message_type="m.room.message",
|
|
||||||
content={
|
|
||||||
"msgtype": message_type.value,
|
|
||||||
**kwargs
|
|
||||||
}
|
|
||||||
)
|
|
||||||
# success
|
|
||||||
if type(result) is RoomSendResponse:
|
|
||||||
return result.event_id
|
|
||||||
# error
|
|
||||||
elif type(result) is RoomSendError:
|
|
||||||
raise Exception(result)
|
|
||||||
# unknown error
|
|
||||||
else:
|
|
||||||
raise RuntimeError()
|
|
||||||
except OlmUnverifiedDeviceError as e:
|
|
||||||
# verify everyone and retry
|
|
||||||
await _verify_all_devices()
|
|
||||||
return await _send_message_to(room_id, message_type, **kwargs)
|
|
||||||
except:
|
|
||||||
raise
|
|
||||||
|
|
||||||
async def _send_text_to(room_id: str, text: str) -> str:
|
|
||||||
"""Sends a text message to the room. `text` may be HTML"""
|
|
||||||
return await _send_message_to(
|
|
||||||
room_id=room_id,
|
|
||||||
message_type=MessageType.TEXT,
|
|
||||||
body=text
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
#
|
|
||||||
# CALLBACKS
|
|
||||||
#
|
|
||||||
async def _message_callback(room: MatrixRoom, event: RoomMessageText) -> None:
|
|
||||||
"""Handle commands received from Matrix"""
|
|
||||||
try:
|
|
||||||
# do not process messages sent by ourselves
|
|
||||||
if event.sender == _client.user_id:
|
|
||||||
return
|
|
||||||
# prepare response
|
|
||||||
response = "<b>Получено сообщение</b>"
|
|
||||||
response += f"<br><br><b>Room ID:</b> {html.escape(room.room_id)}"
|
|
||||||
response += f"<br><b>Sender:</b> {html.escape(event.sender)}"
|
|
||||||
print(await _send_text_to(room.room_id, response))
|
|
||||||
except:
|
|
||||||
traceback.print_exc()
|
|
||||||
|
|
||||||
async def _invite_callback(room: MatrixInvitedRoom, event: InviteMemberEvent) -> None:
|
|
||||||
"""Happens when the bot is invited to somewhere"""
|
|
||||||
try:
|
|
||||||
result = await _client.join(room.room_id)
|
|
||||||
if type(result) is JoinResponse:
|
|
||||||
util.log_info(f"Joined the room {room.room_id}")
|
|
||||||
else:
|
|
||||||
util.log_error(f"Can't join room {room.room_id}")
|
|
||||||
except:
|
|
||||||
traceback.print_exc()
|
|
||||||
|
|
||||||
async def _generic_test_callback(*args, **kwargs) -> None:
|
|
||||||
"""Use this callback to check argument types"""
|
|
||||||
print("GENERIC TEST CALLBACK")
|
|
||||||
for a in args:
|
|
||||||
print(f" - {type(a)}")
|
|
||||||
for k in kwargs:
|
|
||||||
print(f" * {k} = {kwargs[k]}")
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
#
|
|
||||||
# PUBLIC
|
|
||||||
#
|
|
||||||
async def setup(client: AsyncClient) -> None:
|
|
||||||
global _client
|
|
||||||
_client = client
|
|
||||||
client.add_event_callback(_message_callback, RoomMessageText) # type: ignore
|
|
||||||
client.add_event_callback(_invite_callback, InviteMemberEvent) # type: ignore
|
|
||||||
|
|
||||||
async def stop() -> None:
|
|
||||||
"""Stop all ongoing processes"""
|
|
||||||
pass
|
|
||||||
2
main.py
2
main.py
@@ -11,7 +11,7 @@ from pathlib import Path
|
|||||||
import config
|
import config
|
||||||
import util
|
import util
|
||||||
|
|
||||||
from new_bot import MatrixBot
|
from bot import MatrixBot
|
||||||
from bot_types import MatrixBotConfig
|
from bot_types import MatrixBotConfig
|
||||||
|
|
||||||
import nio.events
|
import nio.events
|
||||||
|
|||||||
512
new_bot.py
512
new_bot.py
@@ -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)
|
|
||||||
Reference in New Issue
Block a user