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