Compare commits
1 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 0e9b621895 |
@@ -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`
|
||||||
|
|||||||
@@ -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
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -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" }
|
||||||
]
|
]
|
||||||
|
|||||||
241
src/mab/bot/_client_downloader.py
Normal file
241
src/mab/bot/_client_downloader.py
Normal 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
|
||||||
@@ -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
|
||||||
|
)
|
||||||
Reference in New Issue
Block a user