Compare commits
5 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 16b849ebd8 | |||
| 35bd68516a | |||
| 86257a70e3 | |||
| 43b5d8e3a6 | |||
| 0e9b621895 |
3
.gitignore
vendored
3
.gitignore
vendored
@@ -5,4 +5,5 @@ session_storage/
|
||||
dist/
|
||||
*.egg-info/
|
||||
*.swp
|
||||
*.swo
|
||||
*.swo
|
||||
*.tmp
|
||||
@@ -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,13 +22,13 @@ 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`
|
||||
is needed for E2EE to work.
|
||||
|
||||
Please inspect [`examples/image_gen_bot.py`](examples/image_gen_bot.py),
|
||||
Please inspect [`examples/image_bot.py`](examples/image_bot.py),
|
||||
[`examples/echo_bot.py`](examples/echo_bot.py) or open [`examples/`](examples/)
|
||||
directory to find usage examples. Examples require that you set
|
||||
`MATRIX_HOMESERVER` and `MATRIX_USERNAME` environment variables. Examples create
|
||||
|
||||
97
examples/file_bot.py
Normal file
97
examples/file_bot.py
Normal file
@@ -0,0 +1,97 @@
|
||||
"""
|
||||
This example implements Matrix bot that calculates SHA256 for a file sent by
|
||||
user.
|
||||
|
||||
It uses environment variables to specify authorization data. Use Ctrl+C to stop
|
||||
the bot.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import aiofiles
|
||||
import os
|
||||
import logging
|
||||
import hashlib
|
||||
|
||||
from mab import *
|
||||
|
||||
from _environment import check_environment
|
||||
|
||||
async def on_text_message(ctx: EventContext) -> None:
|
||||
"""This callback is called when a text message arrives."""
|
||||
await ctx.bot.send_text(ctx.room, "Please send a file/image/video")
|
||||
|
||||
async def on_file_message(ctx: EventContext) -> None:
|
||||
"""This callback is called when a file arrives."""
|
||||
sha = hashlib.sha256()
|
||||
does_temp_exist = False
|
||||
if ctx[CTX_FILE_SIZE] > 1_000_000:
|
||||
await ctx.bot.send_text(
|
||||
ctx.room,
|
||||
"The file is larger than 1 MB, downloading to filesystem"
|
||||
)
|
||||
await ctx.bot.download_file(ctx, path="temp.tmp")
|
||||
async with aiofiles.open("temp.tmp", "rb") as f:
|
||||
while True:
|
||||
data = await f.read(64 * 1024)
|
||||
if not data:
|
||||
break
|
||||
sha.update(data)
|
||||
does_temp_exist = True
|
||||
else:
|
||||
await ctx.bot.send_text(
|
||||
ctx.room,
|
||||
"The file is smaller than 1 MB, downloading to RAM"
|
||||
)
|
||||
content = await ctx.bot.download_file(ctx, path=None)
|
||||
sha.update(content)
|
||||
# result
|
||||
response = f"SHA256 for file `{ctx[CTX_FILE_NAME]}`"
|
||||
response += f" ({ctx[CTX_FILE_SIZE]} bytes, {ctx[CTX_FILE_MIME]})"
|
||||
await ctx.bot.send_file_bytes(
|
||||
room=ctx.room,
|
||||
data=sha.hexdigest().encode("utf-8"),
|
||||
filename="hash of the file.txt",
|
||||
mime_type="text/plain",
|
||||
text=response
|
||||
)
|
||||
# resend the file to test uploading
|
||||
if does_temp_exist:
|
||||
await ctx.bot.send_file(
|
||||
room=ctx.room,
|
||||
path="temp.tmp",
|
||||
filename=ctx[CTX_FILE_NAME],
|
||||
text="This is the file you have sent, but it was reuploaded"
|
||||
)
|
||||
try:
|
||||
os.unlink("temp.tmp")
|
||||
except:
|
||||
pass
|
||||
|
||||
|
||||
async def main() -> None:
|
||||
"""Application entry point"""
|
||||
logging.basicConfig(level=logging.INFO)
|
||||
logging.getLogger("nio").setLevel(logging.CRITICAL+1)
|
||||
check_environment()
|
||||
|
||||
config = MatrixBotConfig(
|
||||
matrix_homeserver_url=os.environ["MATRIX_HOMESERVER"],
|
||||
matrix_username_localpart=os.environ["MATRIX_USERNAME"],
|
||||
storage_directory="session_storage"
|
||||
)
|
||||
bot = MatrixBot(config)
|
||||
bot.add_callback(
|
||||
~SenderIsBotFilter() & BodyExistsFilter() & MessageTypeFilter(MessageType.TEXT),
|
||||
on_text_message)
|
||||
bot.add_callback(
|
||||
~SenderIsBotFilter() & MessageHasFile() & NewMessageFilter(),
|
||||
on_file_message)
|
||||
|
||||
# run until Ctrl+C
|
||||
try:
|
||||
await bot.run()
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
|
||||
if __name__ == "__main__":
|
||||
asyncio.run(main())
|
||||
@@ -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,38 @@ 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
|
||||
# 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 <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:
|
||||
@@ -76,12 +104,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
|
||||
)
|
||||
|
||||
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
|
||||
|
||||
[project]
|
||||
name = "mab"
|
||||
version = "0.5.0"
|
||||
version = "0.5.2"
|
||||
authors = [
|
||||
{ name = "Tyukalov Nikita", email = "nikita@tyukalov.su" }
|
||||
]
|
||||
|
||||
@@ -27,6 +27,9 @@ __all__ = [
|
||||
"CTX_CMD_PREFIX",
|
||||
"CTX_CMD_VERB",
|
||||
"CTX_CMD_ARGS",
|
||||
"CTX_FILE_SIZE",
|
||||
"CTX_FILE_MIME",
|
||||
"CTX_FILE_NAME",
|
||||
|
||||
# .bot
|
||||
"MatrixBot",
|
||||
@@ -50,4 +53,5 @@ __all__ = [
|
||||
"RedactedMessageFilter",
|
||||
"SenderIsFilter",
|
||||
"SenderIsBotFilter",
|
||||
"MessageHasFile",
|
||||
]
|
||||
237
src/mab/bot/_client_downloader.py
Normal file
237
src/mab/bot/_client_downloader.py
Normal file
@@ -0,0 +1,237 @@
|
||||
import asyncio
|
||||
import aiofiles
|
||||
import logging
|
||||
import os
|
||||
import threading
|
||||
from uuid import uuid4
|
||||
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 downloader.
|
||||
|
||||
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):
|
||||
temp_path = result.with_name(
|
||||
f".{result.name}.{uuid4().hex}.tmp"
|
||||
)
|
||||
await self._decrypt_file(
|
||||
result, temp_path, key, iv, sha256
|
||||
)
|
||||
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
|
||||
@@ -361,4 +361,123 @@ class ClientSender:
|
||||
}
|
||||
}
|
||||
# send
|
||||
return (await self.send_content(room, content)).event_id
|
||||
|
||||
async def send_file(self,
|
||||
room: MatrixRoom | str,
|
||||
path: Path | str,
|
||||
*,
|
||||
filename: str | None = None,
|
||||
mime_type: str | None = None,
|
||||
text: str | None = None,
|
||||
is_html: bool | None = None,
|
||||
timeout: float | None = 60 * 60):
|
||||
"""
|
||||
Send the file to `room`. Please note that formatted text is displayed
|
||||
incorrectly in some clients as of September 8th, 2026.
|
||||
|
||||
Args:
|
||||
- room - the room to send the file to
|
||||
- path - path to the file
|
||||
- filename - filename to use for upload (`None` to use basename from
|
||||
`path`)
|
||||
- mime_type - mime type to use (`None` for autodetect using content
|
||||
of `path`)
|
||||
- text - caption to use (`None` to disable)
|
||||
- is_html - whether the text is HTML-formatted (`None` for auto)
|
||||
- timeout - upload timeout in seconds (`None` to disable)
|
||||
|
||||
Returns:
|
||||
- `event_id` of sent message on success
|
||||
- Raises an exception on error
|
||||
"""
|
||||
if self._client is None or self._config is None:
|
||||
raise RuntimeError("ClientSender is not set up")
|
||||
# get filename if not set
|
||||
if not filename:
|
||||
filename = os.path.basename(path)
|
||||
# caption must not actually be empty
|
||||
if text is None or not text.strip():
|
||||
text = filename
|
||||
is_html = False
|
||||
# get the mime type if not specified
|
||||
if mime_type is None:
|
||||
mime_type = magic.from_file(path, mime=True)
|
||||
# upload
|
||||
async with asyncio.timeout(timeout):
|
||||
upload_result = await self._uploader.upload_file(
|
||||
path, mime_type=mime_type, filename=filename)
|
||||
# send
|
||||
content = {
|
||||
"msgtype": "m.file",
|
||||
"filename": filename,
|
||||
**self._process_html_text(text, is_html),
|
||||
"file": {
|
||||
"url": upload_result.response.content_uri,
|
||||
"mimetype": mime_type,
|
||||
**upload_result.keys
|
||||
},
|
||||
"info": {
|
||||
"mimetype": mime_type,
|
||||
"size": upload_result.filesize
|
||||
}
|
||||
}
|
||||
return (await self.send_content(room, content)).event_id
|
||||
|
||||
async def send_file_bytes(self,
|
||||
room: MatrixRoom | str,
|
||||
data: bytes,
|
||||
filename: str,
|
||||
*,
|
||||
mime_type: str | None = None,
|
||||
text: str | None = None,
|
||||
is_html: bool | None = None,
|
||||
timeout: float | None = 60 * 60) -> str:
|
||||
"""
|
||||
Send the file to `room`.
|
||||
|
||||
Args:
|
||||
- room - the room to send the file to
|
||||
- data - the content of the file to send
|
||||
- filename - filename to use for the file
|
||||
- mime_type - mime type to use (`None` for autodetect using `data`
|
||||
content)
|
||||
- text - caption to use (`None` to disable)
|
||||
- is_html - whether the text is HTML-formatted (`None` for auto)
|
||||
- timeout - upload timeout in seconds (`None` to disable)
|
||||
|
||||
Returns:
|
||||
- `event_id` of sent message on success
|
||||
- Raises an exception on error
|
||||
"""
|
||||
# caption must not actually be empty
|
||||
if text is None or not text.strip():
|
||||
text = filename
|
||||
is_html = False
|
||||
# check if the file is image
|
||||
if not mime_type:
|
||||
mime_type = magic.from_buffer(data, mime=True)
|
||||
# upload
|
||||
buffer = BytesIO(data)
|
||||
async with asyncio.timeout(timeout):
|
||||
upload_result = await self._uploader.upload_using_provider(
|
||||
provider=buffer,
|
||||
mime_type=mime_type,
|
||||
filename=filename,
|
||||
filesize=len(data))
|
||||
# prepare the content and send
|
||||
content = {
|
||||
"msgtype": "m.file",
|
||||
"filename": filename,
|
||||
**self._process_html_text(text, is_html),
|
||||
"file": {
|
||||
"url": upload_result.response.content_uri,
|
||||
"mimetype": mime_type,
|
||||
**upload_result.keys
|
||||
},
|
||||
"info": {
|
||||
"mimetype": mime_type,
|
||||
"size": upload_result.filesize
|
||||
}
|
||||
}
|
||||
return (await self.send_content(room, content)).event_id
|
||||
@@ -1,9 +1,9 @@
|
||||
import asyncio
|
||||
import logging
|
||||
|
||||
from typing import Callable, Coroutine, Any
|
||||
from typing import Callable, Coroutine, Any, overload
|
||||
|
||||
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,110 @@ class MatrixBot:
|
||||
text=text,
|
||||
is_html=is_html,
|
||||
timeout=timeout
|
||||
)
|
||||
|
||||
async def send_file(self,
|
||||
room: MatrixRoom | str,
|
||||
path: Path | str,
|
||||
*,
|
||||
filename: str | None = None,
|
||||
mime_type: str | None = None,
|
||||
text: str | None = None,
|
||||
is_html: bool | None = None,
|
||||
timeout: float | None = 60 * 60):
|
||||
"""
|
||||
Send the file to `room`. Please note that formatted text is displayed
|
||||
incorrectly in some clients as of September 8th, 2026.
|
||||
|
||||
Args:
|
||||
- room - the room to send the file to
|
||||
- path - path to the file
|
||||
- filename - filename to use for upload (`None` to use basename from
|
||||
`path`)
|
||||
- mime_type - mime type to use (`None` for auto)
|
||||
- text - caption to use (`None` to disable)
|
||||
- is_html - whether the text is HTML-formatted (`None` for auto)
|
||||
- timeout - upload timeout in seconds (`None` to disable)
|
||||
|
||||
Returns:
|
||||
- `event_id` of sent message on success
|
||||
- Raises an exception on error
|
||||
"""
|
||||
return await self._client_sender.send_file(
|
||||
room=room,
|
||||
path=path,
|
||||
filename=filename,
|
||||
mime_type=mime_type,
|
||||
text=text,
|
||||
is_html=is_html,
|
||||
timeout=timeout
|
||||
)
|
||||
|
||||
async def send_file_bytes(self,
|
||||
room: MatrixRoom | str,
|
||||
data: bytes,
|
||||
filename: str,
|
||||
*,
|
||||
mime_type: str | None = None,
|
||||
text: str | None = None,
|
||||
is_html: bool | None = None,
|
||||
timeout: float | None = 60 * 60) -> str:
|
||||
"""
|
||||
Send the file to `room`.
|
||||
|
||||
Args:
|
||||
- room - the room to send the file to
|
||||
- data - the content of the file to send
|
||||
- filename - filename to use for the file
|
||||
- text - caption to use (`None` to disable)
|
||||
- is_html - whether the text is HTML-formatted (`None` for auto)
|
||||
- timeout - upload timeout in seconds (`None` to disable)
|
||||
|
||||
Returns:
|
||||
- `event_id` of sent message on success
|
||||
- Raises an exception on error
|
||||
"""
|
||||
return await self._client_sender.send_file_bytes(
|
||||
room=room,
|
||||
data=data,
|
||||
filename=filename,
|
||||
mime_type=mime_type,
|
||||
text=text,
|
||||
is_html=is_html,
|
||||
timeout=timeout
|
||||
)
|
||||
|
||||
@overload
|
||||
async def download_file(self,
|
||||
source: EventContext | Event,
|
||||
*,
|
||||
path: str | Path) -> Path: ...
|
||||
|
||||
@overload
|
||||
async def download_file(self,
|
||||
source: EventContext | Event,
|
||||
*,
|
||||
path: None = None) -> bytes: ...
|
||||
|
||||
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
|
||||
)
|
||||
@@ -34,6 +34,15 @@ CTX_CMD_ARGS = ContextDataKey[list[str]]("CTX_CMD_ARGS")
|
||||
CTX_ROOM_ENCRYPTED = ContextDataKey[bool]("CTX_ROOM_ENCRYPTED")
|
||||
"""True if the room is encrypted"""
|
||||
|
||||
CTX_FILE_SIZE = ContextDataKey[int]("CTX_FILE_SIZE")
|
||||
"""Size of the file attached to the message (bytes)"""
|
||||
|
||||
CTX_FILE_MIME = ContextDataKey[str]("CTX_FILE_MIME")
|
||||
"""Mime type of the file attached to the message"""
|
||||
|
||||
CTX_FILE_NAME = ContextDataKey[str]("CTX_FILE_NAME")
|
||||
"""Name of the file attached to the message"""
|
||||
|
||||
|
||||
#
|
||||
# EventContext implementation
|
||||
|
||||
@@ -1,7 +1,12 @@
|
||||
import traceback
|
||||
from .base import BaseEventFilter, EventTypeFilter
|
||||
from ..types import MessageType
|
||||
from ..context import EventContext, CTX_MESSAGE_TYPE, CTX_SENDER
|
||||
from ..context import (EventContext,
|
||||
CTX_MESSAGE_TYPE,
|
||||
CTX_SENDER,
|
||||
CTX_FILE_SIZE,
|
||||
CTX_FILE_MIME,
|
||||
CTX_FILE_NAME)
|
||||
|
||||
from nio import AsyncClient
|
||||
from nio import MatrixRoom, Event
|
||||
@@ -118,4 +123,29 @@ class SenderIsBotFilter(BaseEventFilter):
|
||||
async def __call__(self, context: EventContext) -> bool:
|
||||
if not await super().__call__(context):
|
||||
return False
|
||||
return context.bot.get_client().user_id == context.event.sender
|
||||
return context.bot.get_client().user_id == context.event.sender
|
||||
|
||||
class MessageHasFile(BaseEventFilter):
|
||||
"""
|
||||
This filter returns True if the message contains file that can be
|
||||
downloaded.
|
||||
|
||||
This filter sets `CTX_FILE_NAME`, `CTX_FILE_SIZE` and `CTX_FILE_MIME`
|
||||
variables in the context (if they are present in the).
|
||||
"""
|
||||
def __init__(self, **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
|
||||
async def __call__(self, context: EventContext) -> bool:
|
||||
if not await super().__call__(context):
|
||||
return False
|
||||
content: dict | None = context.event.source.get("content")
|
||||
if content is None:
|
||||
return False
|
||||
# `file` if encrypted, `url` if not encrypted
|
||||
if "file" not in content and "url" not in content:
|
||||
return False
|
||||
context[CTX_FILE_NAME] = content.get("filename") or content.get("body")
|
||||
context[CTX_FILE_SIZE] = content["info"].get("size")
|
||||
context[CTX_FILE_MIME] = content["info"].get("mimetype")
|
||||
return True
|
||||
Reference in New Issue
Block a user