Added file downloading, updated to v0.5.1

This commit is contained in:
2026-09-13 03:17:12 +03:00
parent f3f1e24c0b
commit 0e9b621895
5 changed files with 315 additions and 7 deletions

View File

@@ -9,6 +9,7 @@ there, but it aims to be convenient and usable for relatively serious projects.
The library supports the following features: The library supports the following features:
- **Completely `asyncio` based** - **Completely `asyncio` based**
- **Filter-based callback system** - **Filter-based callback system**
- **Downloading and transparently decrypting files**
- **Sending images** - **Sending images**
- **Sending videos with automatic thumbnail generation (requires `ffmpeg`)** - **Sending videos with automatic thumbnail generation (requires `ffmpeg`)**
@@ -21,7 +22,7 @@ install the latest version of the library:
```bash ```bash
apt install libmagic1-dev libolm-dev 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` `libmagic1-dev` is needed for automatic file MIME type detection, `libolm-dev`

View File

@@ -9,9 +9,10 @@ the bot.
import asyncio import asyncio
from io import BytesIO from io import BytesIO
import os import os
import time
import logging import logging
import random import random
from PIL import Image from PIL import Image, ImageFilter
from mab import * from mab import *
@@ -45,11 +46,41 @@ async def on_gen_command(data: EventContext) -> None:
# send # send
await data.bot.send_image_bytes(data.room, buf, "noise.png") 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: async def on_wrong_message(data: EventContext) -> None:
"""This callback is called when a wrong message is received.""" """This callback is called when a wrong message is received."""
await data.bot.send_text( await data.bot.send_text(
data.room, data.room,
"Text me something like <code>!gen 0.1 0.7 1.0</code>" "Text me something like <code>!gen 0.1 0.7 1.0</code> or send an image to blur"
) )
async def main() -> None: async def main() -> None:
@@ -77,11 +108,17 @@ async def main() -> None:
~SenderIsBotFilter() & command_filter, ~SenderIsBotFilter() & command_filter,
on_gen_command) on_gen_command)
# callback for new image message
bot.add_callback(
~SenderIsBotFilter() & MessageTypeFilter(MessageType.IMAGE) & NewMessageFilter(),
on_image
)
# callback for message that # callback for message that
# 1. are sent not by this bot # 1. are sent not by this bot
# 2. do NOT match the command filter # 2. are new messages (not edits)
bot.add_callback( bot.add_callback(
~SenderIsBotFilter() & ~command_filter, ~SenderIsBotFilter() & BodyExistsFilter() & NewMessageFilter(),
on_wrong_message on_wrong_message
) )

View File

@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
[project] [project]
name = "mab" name = "mab"
version = "0.5.0" version = "0.5.1"
authors = [ authors = [
{ name = "Tyukalov Nikita", email = "nikita@tyukalov.su" } { name = "Tyukalov Nikita", email = "nikita@tyukalov.su" }
] ]

View File

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

View File

@@ -3,7 +3,7 @@ import logging
from typing import Callable, Coroutine, Any from typing import Callable, Coroutine, Any
from nio import AsyncClient, MatrixRoom from nio import AsyncClient, MatrixRoom, Event
from ..filters.base import BaseEventFilter from ..filters.base import BaseEventFilter
from ..types import * from ..types import *
@@ -12,6 +12,7 @@ from ..context import EventContext
from ._validation import Validator from ._validation import Validator
from ._storage import Storage from ._storage import Storage
from ._client_auth import ClientAuth from ._client_auth import ClientAuth
from ._client_downloader import ClientDownloader
from ._client_manager import ClientManager from ._client_manager import ClientManager
from ._client_uploader import ClientUploader from ._client_uploader import ClientUploader
from ._client_sender import ClientSender from ._client_sender import ClientSender
@@ -33,6 +34,7 @@ class MatrixBot:
self._client_auth = ClientAuth(self._storage) self._client_auth = ClientAuth(self._storage)
self._client_manager = ClientManager(self._client_auth, self._storage) self._client_manager = ClientManager(self._client_auth, self._storage)
self._client_uploader = ClientUploader(self._storage) self._client_uploader = ClientUploader(self._storage)
self._client_downloader = ClientDownloader()
self._client_sender = ClientSender() self._client_sender = ClientSender()
self._callbacks = Callbacks(self._storage, self) self._callbacks = Callbacks(self._storage, self)
# validate the config and save it # validate the config and save it
@@ -80,6 +82,10 @@ class MatrixBot:
self._config, self._config,
self._client_manager.get_client() self._client_manager.get_client()
) )
await self._client_downloader.setup(
self._config,
self._client_manager.get_client()
)
await self._client_sender.setup( await self._client_sender.setup(
self._config, self._config,
self._client_manager.get_client(), self._client_manager.get_client(),
@@ -245,3 +251,26 @@ class MatrixBot:
is_html=is_html, is_html=is_html,
timeout=timeout 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
)