Compare commits

..

2 Commits

Author SHA1 Message Date
cee1275e3c Deleted logic.py 2026-08-21 07:08:53 +03:00
9b702df01f Deleted old bot.py 2026-08-21 06:57:31 +03:00
4 changed files with 491 additions and 841 deletions

670
bot.py
View File

@@ -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: except:
traceback.print_exc() traceback.print_exc()
await asyncio.sleep(1)
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:
# PUBLIC """Reads `next_batch` value from disk. Returns None if file does not exist."""
# path = self._config.storage_directory / "next_batch"
async def start(config: AppConfig) -> bool: if not path.is_file():
"""Starts the bot""" return None
global _client async with aiofiles.open(path, "r") as f:
global _task, _stop, _app_config return (await f.read()).strip()
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: async def _write_session_data(self, *, access_token: str, device_id: str) -> None:
return _client """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 stop() -> None: async def _read_session_data(self) -> dict[str, Any] | None:
"""Stop the bot""" """Read session data from disk."""
global _task, _stop path = self._config.storage_directory / "session_data.json"
if _task is None or _stop is None: if not path.is_file():
return return None
_stop.set() async with aiofiles.open(path, "r") as f:
try: j = json.loads(await f.read())
await _task return j
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)
#
# 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)

148
logic.py
View File

@@ -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

View File

@@ -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

View File

@@ -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)