diff --git a/README.md b/README.md index 2e97de0..3d37202 100644 --- a/README.md +++ b/README.md @@ -9,6 +9,7 @@ there, but it aims to be convenient and usable for relatively serious projects. The library supports the following features: - **Completely `asyncio` based** - **Filter-based callback system** +- **Downloading and transparently decrypting files** - **Sending images** - **Sending videos with automatic thumbnail generation (requires `ffmpeg`)** @@ -21,7 +22,7 @@ install the latest version of the library: ```bash apt install libmagic1-dev libolm-dev -python -m pip install git+https://git.tyukalov.su/nikita/mab@v0.5.0 +python -m pip install git+https://git.tyukalov.su/nikita/mab@v0.5.1 ``` `libmagic1-dev` is needed for automatic file MIME type detection, `libolm-dev` diff --git a/examples/image_gen_bot.py b/examples/image_bot.py similarity index 67% rename from examples/image_gen_bot.py rename to examples/image_bot.py index dede3c9..3e45c39 100644 --- a/examples/image_gen_bot.py +++ b/examples/image_bot.py @@ -9,9 +9,10 @@ the bot. import asyncio from io import BytesIO import os +import time import logging import random -from PIL import Image +from PIL import Image, ImageFilter from mab import * @@ -45,11 +46,41 @@ async def on_gen_command(data: EventContext) -> None: # send await data.bot.send_image_bytes(data.room, buf, "noise.png") +async def on_image(ctx: EventContext) -> None: + """This callback is called when an image is received.""" + # download + s = time.time() + data = await ctx.bot.download_file(ctx) + took_time = time.time() - s + if not isinstance(data, bytes): + await ctx.bot.send_text(ctx.room, "😧") + return + # convert to Image + with BytesIO(data) as buf: + img = Image.open(buf) + # apply effects + blur = ImageFilter.GaussianBlur( + radius=min(img.size[0] // 10, 5) + ) + img = img.filter(blur) + # save to buffer + buf = BytesIO() + img.save(buf, format="PNG") + buf.seek(0) + buf = buf.read() + # send + await ctx.bot.send_image_bytes( + ctx.room, + buf, + "blurred.png", + text="Download and decryption took %.4f seconds" % took_time + ) + async def on_wrong_message(data: EventContext) -> None: """This callback is called when a wrong message is received.""" await data.bot.send_text( data.room, - "Text me something like !gen 0.1 0.7 1.0" + "Text me something like !gen 0.1 0.7 1.0 or send an image to blur" ) async def main() -> None: @@ -76,12 +107,18 @@ async def main() -> None: bot.add_callback( ~SenderIsBotFilter() & command_filter, on_gen_command) + + # callback for new image message + bot.add_callback( + ~SenderIsBotFilter() & MessageTypeFilter(MessageType.IMAGE) & NewMessageFilter(), + on_image + ) # callback for message that # 1. are sent not by this bot - # 2. do NOT match the command filter + # 2. are new messages (not edits) bot.add_callback( - ~SenderIsBotFilter() & ~command_filter, + ~SenderIsBotFilter() & BodyExistsFilter() & NewMessageFilter(), on_wrong_message ) diff --git a/pyproject.toml b/pyproject.toml index 0cb7e83..ae641a2 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta" [project] name = "mab" -version = "0.5.0" +version = "0.5.1" authors = [ { name = "Tyukalov Nikita", email = "nikita@tyukalov.su" } ] diff --git a/src/mab/bot/_client_downloader.py b/src/mab/bot/_client_downloader.py new file mode 100644 index 0000000..601b8ca --- /dev/null +++ b/src/mab/bot/_client_downloader.py @@ -0,0 +1,241 @@ +import asyncio +import aiofiles +import logging +import os +import time +import threading +import traceback +import random +from pathlib import Path + +import hashlib +import hmac +import unpaddedbase64 +from Crypto.Cipher import AES +from Crypto.Util import Counter + +from nio import ( + AsyncClient, + Event, + DiskDownloadResponse, + MemoryDownloadResponse +) + +from ..types import MatrixBotConfig +from ..context import EventContext + +class ClientDownloader: + """This class downloads files""" + + def __init__(self): + self._logger = logging.getLogger("ClientSender") + self._config: MatrixBotConfig | None = None + self._client: AsyncClient | None = None + + async def setup(self, + config: MatrixBotConfig, + client: AsyncClient) -> None: + """ + Setup the sender. + + Args: + - config - config to use + - client - client to use + """ + self._config = config + self._client = client + + async def _decrypt_file(self, + src: Path, + dst: Path, + key: str, + iv: str, + sha256: str) -> None: + """ + Decrypts `src` and saves it as `dst` asynchronously. + + Args: + - src - path to the source (encrypted) file + - dst - path where the resulting (decrypted) file will be saved + - key - key, from event["content"]["file"]["key"]["k"] + - iv - initialization vector, from event["content"]["file"]["iv"] + - sha256 - SHA-256 digest, from + event["content"]["file"]["hashes"]["sha256"] + + Returns: + - Returns nothing on success + - Raises an exception on error + """ + async with aiofiles.open(src, "rb") as reader: + # check file hash + expected = unpaddedbase64.decode_base64(sha256) + digest = hashlib.sha256() + while chunk := await reader.read(16 * 1024): + digest.update(chunk) + if not hmac.compare_digest(digest.digest(), expected): + raise RuntimeError("SHA256 mismatch") + await reader.seek(0) + # decrypt + decoded_key = unpaddedbase64.decode_base64(key) + decoded_iv = unpaddedbase64.decode_base64(iv) + cipher = AES.new( + decoded_key, + AES.MODE_CTR, + counter=Counter.new( + nbits=64, + prefix=decoded_iv[:8], + initial_value=int.from_bytes(decoded_iv[8:], "big") + ) + ) + async with aiofiles.open(dst, "wb") as writer: + while chunk := await reader.read(16 * 1024): + await writer.write(cipher.decrypt(chunk)) + + async def _decrypt_bytes(self, + data: bytes, + key: str, + iv: str, + sha256: str) -> bytes: + """ + Decrypts `data`. It starts subthread that performs decryption. + + Args: + - data - data that needs to be decrypted + - key - key, from event["content"]["file"]["key"]["k"] + - iv - initialization vector, from event["content"]["file"]["iv"] + - sha265 - SHA-256 digest, from + event["content"]["file"]["hashes"]["sha256"] + + Returns: + - Returns decrypted data on success + - Raises an exception on error + """ + cancel = threading.Event() + # subfunction that will run in another thread + def decrypt() -> bytes: + chunk_size = 64 * 1024 + view = memoryview(data) + expected = unpaddedbase64.decode_base64(sha256) + digest = hashlib.sha256() + for offset in range(0, len(view), chunk_size): + if cancel.is_set(): + return b"" + digest.update(view[offset:offset + chunk_size]) + if not hmac.compare_digest(digest.digest(), expected): + raise RuntimeError("SHA-256 mismatch") + # prepare AES-CTR + decoded_key = unpaddedbase64.decode_base64(key) + decoded_iv = unpaddedbase64.decode_base64(iv) + cipher = AES.new( + decoded_key, + AES.MODE_CTR, + counter=Counter.new( + nbits=64, + prefix=decoded_iv[:8], + initial_value=int.from_bytes(decoded_iv[8:], "big"), + ), + ) + # decrypt + result = bytearray() + for offset in range(0, len(view), chunk_size): + if cancel.is_set(): + return b"" + result.extend( + cipher.decrypt(view[offset:offset + chunk_size]) + ) + return bytes(result) + # decrypt in thread + try: + return await asyncio.to_thread(decrypt) + except asyncio.CancelledError: + cancel.set() + raise + + async def download_file(self, + source: EventContext | Event, + *, + path: str | Path | None = None) -> Path | bytes: + """ + Download the file from the event. Automatically deciphers encrypted + media. + + Args: + - source - event that has a file in its `content` + - path - where to save the file to. Use `None` to store the file in + memory. Use path to specify file download path. + + Returns: + - Returns `Path` to the file on success (if `path` isn't `None`) + - Returns `bytes` of the file on success (if `path` is `None`) + - Raises an exception on error + """ + if self._client is None: + raise RuntimeError("ClientDownloader is not set up") + # prepare the args + if isinstance(source, EventContext): + source = source.event + content: dict = source.source["content"] + f = content.get("file") + # encypted + if f: + if not isinstance(f["url"], str): + raise RuntimeError("`source` does not contain file URL") + mxc = f["url"] + if f["key"]["alg"].upper() != "A256CTR": + raise NotImplementedError( + f"Unsupported encryption algorithm `{f['key']['alg']}`" + ) + key: str | None = f["key"]["k"] + iv: str | None = f["iv"] + sha256: str | None = f["hashes"]["sha256"] + # not encrypted + else: + mxc = content["url"] + key = None + iv = None + sha256 = None + # prepare path + if isinstance(path, str): + path = Path(path) + filename: str | None = \ + os.path.basename(path) if isinstance(path, Path) else None + # download + result = await self._client.download( + mxc, + filename=filename, + save_to=path + ) + # check the response + if isinstance(result, DiskDownloadResponse): + result = Path(result.body) + elif isinstance(result, MemoryDownloadResponse): + result = result.body + else: + raise RuntimeError(result) + # decrypt if needed + temp_path: Path | None = None + try: + if key is not None and iv is not None and sha256 is not None: + # on disk + if isinstance(result, Path): + r = unpaddedbase64.encode_base64(random.randbytes(8)) + temp_path = result.with_name( + f".{result.name}.{int(time.time())}{r}.tmp" + ) + await self._decrypt_file( + result, temp_path, key, iv, sha256 + ) + result.unlink() + temp_path.replace(result) + temp_path = None + # in memory + else: + result = await self._decrypt_bytes( + result, key, iv, sha256 + ) + finally: + if isinstance(temp_path, Path) and isinstance(result, Path): + result.unlink(missing_ok=True) + temp_path.unlink(missing_ok=True) + # return the result + return result \ No newline at end of file diff --git a/src/mab/bot/bot.py b/src/mab/bot/bot.py index fda11b2..3a79e29 100644 --- a/src/mab/bot/bot.py +++ b/src/mab/bot/bot.py @@ -3,7 +3,7 @@ import logging from typing import Callable, Coroutine, Any -from nio import AsyncClient, MatrixRoom +from nio import AsyncClient, MatrixRoom, Event from ..filters.base import BaseEventFilter from ..types import * @@ -12,6 +12,7 @@ from ..context import EventContext from ._validation import Validator from ._storage import Storage from ._client_auth import ClientAuth +from ._client_downloader import ClientDownloader from ._client_manager import ClientManager from ._client_uploader import ClientUploader from ._client_sender import ClientSender @@ -33,6 +34,7 @@ class MatrixBot: self._client_auth = ClientAuth(self._storage) self._client_manager = ClientManager(self._client_auth, self._storage) self._client_uploader = ClientUploader(self._storage) + self._client_downloader = ClientDownloader() self._client_sender = ClientSender() self._callbacks = Callbacks(self._storage, self) # validate the config and save it @@ -80,6 +82,10 @@ class MatrixBot: self._config, self._client_manager.get_client() ) + await self._client_downloader.setup( + self._config, + self._client_manager.get_client() + ) await self._client_sender.setup( self._config, self._client_manager.get_client(), @@ -244,4 +250,27 @@ class MatrixBot: text=text, is_html=is_html, timeout=timeout + ) + + async def download_file(self, + source: EventContext | Event, + *, + path: str | Path | None = None) -> Path | bytes: + """ + Download the file from the event. Automatically dechiphers encrypted + media. + + Args: + - source - event that has a file in its `content` + - path - where to save the file to. Use `None` to store the file in + memory. Use path to specify file download path. + + Returns: + - Returns `Path` to the file on success (if `path` isn't `None`) + - Returns `bytes` of the file on success (if `path` is `None`) + - Raises an exception on error + """ + return await self._client_downloader.download_file( + source=source, + path=path ) \ No newline at end of file