15 Commits

Author SHA1 Message Date
16b849ebd8 Added file uploading support, v0.5.2
- Added example `file_bot.py`
- Added *.tmp to .gitignore
- Added CTX_FILE_SIZE, CTX_FILE_MIME and CTX_FILE_NAME
- Added MessageHasFile filter
- Added `send_file` and `send_file_bytes` for ClientSender and Bot
2026-09-23 19:43:30 +03:00
35bd68516a Fixed README.md link to image_bot.py 2026-09-13 04:18:43 +03:00
86257a70e3 Added overloads for MatrixBot.download_file 2026-09-13 03:39:23 +03:00
43b5d8e3a6 Fixed downloader bugs related to temp file name 2026-09-13 03:28:27 +03:00
0e9b621895 Added file downloading, updated to v0.5.1 2026-09-13 03:17:12 +03:00
f3f1e24c0b Updated to v0.5.0 2026-09-12 23:34:57 +03:00
2e4228b22a Added example for commands 2026-09-12 23:34:06 +03:00
0b79d78715 Fixed RoomEncryptedFilter 2026-09-12 23:01:48 +03:00
a90cd66d3e Fixed CTX_MESSAGE_TYPE typing 2026-09-12 22:56:45 +03:00
cb520814d8 Added description of filter system to README.md 2026-09-12 22:49:55 +03:00
4287a0d20c Improved filters system
- `RoomEventData` is renamed to `EventContext` and moved to context.py
- `EventContext.filter` is removed
- Added context variables system which improves type hints and simplifies callbacks code
- `BodyCommandFilter` sets context variables from now on
- Added `super().__call__` invocation to filters implements in base.py
- Examples are updated to include required changes
2026-09-12 22:19:28 +03:00
e10c920a56 Fixed README.md 2026-09-12 20:24:18 +03:00
156afd6b61 Fixed README references example that don't exist 2026-09-12 20:23:44 +03:00
2d47483d55 Fixed BodyRegexFilter returning re.Match 2026-09-12 20:22:02 +03:00
fc4f664a5c Fixed BaseEventFilter.__ror__ recursion 2026-09-12 20:20:41 +03:00
18 changed files with 1061 additions and 130 deletions

3
.gitignore vendored
View File

@@ -5,4 +5,5 @@ session_storage/
dist/
*.egg-info/
*.swp
*.swo
*.swo
*.tmp

View File

@@ -9,10 +9,11 @@ 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`)**
## 🚀 Usage
## 📦 Installation
Use `apt` to install required system packages and `pip` to install the package.
You may need to use `root` privileges to use `apt`. It's highly recommended you
@@ -21,18 +22,78 @@ 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.4.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/shell_bot.py`](examples/shell_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
`session_storage` directory in working directory.
## 🚀 Usage
If you use `mab`, your application will *most likely* be using **callbacks** to
react to user actions. `mab` uses filter-based callback system to avoid exposing
raw `nio-matrix` event objects.
This is the workflow you will most likely follow:
1. **Define the callback as `async` function that take 1 argument of type
`EventContext`.** For example, this callback would print the caption of the
message:
```python
from mab import *
async def on_media_with_body(ctx: EventContext):
"""To be called when a message with image/video and caption is received."""
print(ctx[CTX_BODY])
```
2. **Define the conditions your callback must be called on.** For example, you
may want your callback to be called when `the sender is not the bot` and
`the message contains textual body` and (`the message is an image` or
`the message is a video`).
3. **Define the conditions as `filters`.** Most of them are pretty
straightforward. For example, if you want to use the conditions from above:
```python
from mab import *
filters = (
~SenderIsBotFilter()
& MessageTypeFilter([MessageType.IMAGE, MessageType.VIDEO])
& BodyExistsFilter()
)
```
4. **Add the callback to your `MatrixBot` instance.** For example, if you would
have used everything from above, then your code would look something like
this:
```python
from mab import *
# let's assume you create your MatrixBot as `bot` variable here
async def on_media_with_body(ctx: EventContext):
"""To be called when a message with image/video and caption is received."""
print(ctx[CTX_BODY])
filters = (
~SenderIsBotFilter()
& MessageTypeFilter([MessageType.IMAGE, MessageType.VIDEO])
& BodyExistsFilter()
)
bot.add_callback(filters, on_media_with_body)
...
```
Filters support bitwise operators to implement complex matching logic. Some
filters set context variables which can be accessed by
`context[CTX_KEY_NAME]`-like syntax. Possible variables are defined in
[this file](src/mab/context.py). Filters are implemented in
[files of this directory](src/mab/filters/).
## 🏷️ Versioning
Releases are tagged in this repository using the `vX.Y.Z` format. If the commit

130
examples/command_bot.py Normal file
View File

@@ -0,0 +1,130 @@
"""
This example implements Matrix bot that can execute some commands.
It uses environment variables to specify authorization data. Use Ctrl+C to stop
the bot.
"""
import asyncio
import time
import html
import os
import traceback
import logging
from mab import *
from _environment import check_environment
async def on_help_command(ctx: EventContext) -> None:
"""!help"""
HELP_MESSAGE = (
"<strong>Here is the list of the commands:</strong><br>"
"<ul>"
"<li><code>!help</code> - this help message</li>"
"<li><code>!time</code> - get UNIX timestamp</li>"
"<li><code>!raise</code> - raise <code>RuntimeError()</code></li>"
"<li><code>!assert</code> - perform <code>assert</code> that will fail</li>"
"<li><code>!mul A B [C] [D]...</code> - multiply A, B... and so on</li>"
"<li><code>!args arg1 [arg2] ... [arg5]</code> - command that takes 1..5 arguments</li>"
"</ul>"
)
await ctx.bot.send_text(ctx.room, HELP_MESSAGE)
async def on_time_command(ctx: EventContext) -> None:
"""!time"""
await ctx.bot.send_text(
ctx.room,
f"Current UNIX timestamp is <strong>{int(time.time())}</strong>"
)
async def on_raise_command(ctx: EventContext) -> None:
"""!raise"""
await ctx.bot.send_text(
ctx.room,
"<strong>Executing <code>raise RuntimeError()</code>...</strong>"
)
raise RuntimeError()
async def on_assert_command(ctx: EventContext) -> None:
"""!assert"""
await ctx.bot.send_text(
ctx.room,
"<strong>Executing <code>assert False</code>...</strong>"
)
assert False
async def on_mul_command(ctx: EventContext) -> None:
"""!mul"""
try:
numbers = [float(v) for v in ctx[CTX_CMD_ARGS]]
v = numbers[0]
for n in numbers[1:]:
v *= n
response = " * ".join(html.escape("%.2f" % n) for n in numbers)
response += f" = <strong>{html.escape(str(v))}<strong>"
await ctx.bot.send_text(ctx.room, response)
except Exception as e:
await ctx.bot.send_text(ctx.room, f"Could not process the command: {e}")
async def on_args_command(ctx: EventContext) -> None:
"""!args"""
try:
response = (
f"Prefix: <code>{ctx[CTX_CMD_PREFIX]}</code><br>"
f"Verb: <code>{ctx[CTX_CMD_VERB]}</code><br>"
f"Arguments: <code>{len(ctx[CTX_CMD_ARGS])}</code><br>"
f"Arguments are:<br><ol>"
)
for arg in ctx[CTX_CMD_ARGS]:
response += f"<li><code>{html.escape(arg)}</code></li>"
response += "</ol>"
await ctx.bot.send_text(
ctx.room,
response
)
except:
await ctx.bot.send_text(
ctx.room,
f"Could not process the command: {traceback.format_exc()}"
)
async def invalid_usage(ctx: EventContext) -> None:
"""This callback is called when the bot used incorrectly."""
await ctx.bot.send_text(ctx.room, "Use <code>!help</code>")
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)
COMMANDS = {
on_help_command: BodyCommandFilter(["help", "?"]),
on_time_command: BodyCommandFilter("time"),
on_raise_command: BodyCommandFilter("raise"),
on_assert_command: BodyCommandFilter("assert"),
on_mul_command: BodyCommandFilter("mul", min_args=2),
on_args_command: BodyCommandFilter("args", min_args=1, max_args=5),
}
for callback, filter in COMMANDS.items():
f = ~SenderIsBotFilter() & filter
bot.add_callback(f, callback)
bot.add_callback(~SenderIsBotFilter() & NewMessageFilter(), invalid_usage)
# run until Ctrl+C
try:
await bot.run()
except asyncio.CancelledError:
pass
if __name__ == "__main__":
asyncio.run(main())

View File

@@ -9,21 +9,13 @@ import asyncio
import os
import logging
from mab import (
MatrixBot,
MatrixBotConfig,
RoomEventData,
BodyExistsFilter,
MessageTypeFilter,
SenderIsBotFilter,
MessageType
)
from mab import *
from _environment import check_environment
async def on_text_message(data: RoomEventData) -> None:
async def on_text_message(ctx: EventContext) -> None:
"""This callback is called when a text message arrives."""
await data.bot.send_text(data.room, data.event.body) # type: ignore
await ctx.bot.send_text(ctx.room, ctx[CTX_BODY])
async def main() -> None:
"""Application entry point"""

97
examples/file_bot.py Normal file
View 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())

View File

@@ -9,27 +9,20 @@ 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 (
MatrixBot,
MatrixBotConfig,
RoomEventData,
MessageTypeFilter,
BodyCommandFilter,
SenderIsBotFilter,
MessageType
)
from mab import *
from _environment import check_environment
async def on_gen_command(data: RoomEventData) -> None:
async def on_gen_command(data: EventContext) -> None:
"""This callback is called when `!gen R G B` command is received."""
# convert R, G and B to floats
try:
r, g, b = [float(v) for v in data.event.command_args] # type: ignore
r, g, b = [float(v) for v in data[CTX_CMD_ARGS]]
except:
await data.bot.send_text(data.room, "Invalid arguments")
return
@@ -53,11 +46,38 @@ async def on_gen_command(data: RoomEventData) -> None:
# send
await data.bot.send_image_bytes(data.room, buf, "noise.png")
async def on_wrong_message(data: RoomEventData) -> None:
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:
@@ -84,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
)

View File

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

View File

@@ -1,7 +1,8 @@
from . import bot
from . import types
from .types import MatrixBotConfig, RoomEventData, MessageType
from .types import MatrixBotConfig, MessageType
from .context import *
from .bot import MatrixBot
@@ -16,9 +17,20 @@ __all__ = [
# .types
"MatrixBotConfig",
"RoomEventData",
"MessageType",
# .context
"EventContext",
"CTX_BODY",
"CTX_MESSAGE_TYPE",
"CTX_SENDER",
"CTX_CMD_PREFIX",
"CTX_CMD_VERB",
"CTX_CMD_ARGS",
"CTX_FILE_SIZE",
"CTX_FILE_MIME",
"CTX_FILE_NAME",
# .bot
"MatrixBot",
@@ -41,4 +53,5 @@ __all__ = [
"RedactedMessageFilter",
"SenderIsFilter",
"SenderIsBotFilter",
"MessageHasFile",
]

View File

@@ -10,7 +10,8 @@ from nio import MatrixInvitedRoom, InviteMemberEvent, JoinResponse
from nio.events.room_events import Event as RoomEvent
from ._storage import Storage
from ..types import MatrixBotConfig, RoomEventData
from ..types import MatrixBotConfig
from ..context import EventContext
from ..filters.base import BaseEventFilter
if TYPE_CHECKING:
@@ -31,7 +32,7 @@ class Callbacks:
filter: BaseEventFilter
"""Filter to use for matching"""
callback: Callable[[RoomEventData], Coroutine[Any, Any, None]] | None
callback: Callable[[EventContext], Coroutine[Any, Any, None]] | None
"""Callback that will be called if the filter matches"""
stop_matching: bool
@@ -51,20 +52,19 @@ class Callbacks:
for callback_info in self._filters:
if not isinstance(callback_info, self._FilterBasedCallback):
continue
event_data = EventContext(
room=room,
event=event,
bot=self._matrix_bot
)
try:
if not await callback_info.filter(room, event, self._client):
if not await callback_info.filter(event_data):
continue
except asyncio.CancelledError:
raise
except:
self._logger.error(traceback.format_exc())
continue
event_data = RoomEventData(
room=room,
event=event,
filter=callback_info.filter,
bot=self._matrix_bot
)
try:
# dump argument types
if callback_info.callback is None:
@@ -144,7 +144,7 @@ class Callbacks:
def add_room_event_callback(
self,
filter: BaseEventFilter,
callback: Callable[[RoomEventData], Coroutine[Any, Any, None]] | None,
callback: Callable[[EventContext], Coroutine[Any, Any, None]] | None,
*,
stop_matching: bool = True) -> None:
"""

View 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

View File

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

View File

@@ -1,16 +1,18 @@
import asyncio
import logging
from typing import Callable, Coroutine, Any
from typing import Callable, Coroutine, Any, overload
from nio import AsyncClient
from nio import AsyncClient, MatrixRoom, Event
from ..filters.base import BaseEventFilter
from ..types import *
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
@@ -32,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
@@ -44,7 +47,7 @@ class MatrixBot:
def add_callback(self,
filter: BaseEventFilter,
callback: Callable[[RoomEventData], Coroutine[Any, Any, None]] | None,
callback: Callable[[EventContext], Coroutine[Any, Any, None]] | None,
*,
stop_matching: bool = True) -> None:
"""
@@ -79,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(),
@@ -243,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
)

88
src/mab/context.py Normal file
View File

@@ -0,0 +1,88 @@
"""This module implements logic for event context"""
from dataclasses import dataclass
from typing import Any, TYPE_CHECKING
from .types import ContextDataKey, MessageType
from nio import MatrixRoom, Event
if TYPE_CHECKING:
from .bot import MatrixBot
#
# Possible context variables
#
CTX_BODY = ContextDataKey[str]("CTX_BODY")
"""Value of `event.body`"""
CTX_MESSAGE_TYPE = ContextDataKey[MessageType]("CTX_MESSAGE_TYPE")
"""Value of `msgtype` for the event"""
CTX_SENDER = ContextDataKey[str]("CTX_SENDER")
"""Value of `event.sender`"""
CTX_CMD_PREFIX = ContextDataKey[str]("CTX_CMD_PREFIX")
"""Command prefix that was used when matching"""
CTX_CMD_VERB = ContextDataKey[str]("CTX_CMD_VERB")
"""The verb that was used to execute the command"""
CTX_CMD_ARGS = ContextDataKey[list[str]]("CTX_CMD_ARGS")
"""Arguments that were passed with the command"""
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
#
@dataclass
class EventContext:
"""The class holding information about an event that happened in the room"""
room: MatrixRoom
"""The room the event has happened in"""
event: Event
"""The event that has happened in the room"""
bot: "MatrixBot"
"""The bot that is the source of the event"""
def __setitem__[T](self, key: ContextDataKey[T], value: T | None) -> None:
"""Set a value inside the context data storage. `None` removes it"""
if not hasattr(self, "_datastore"):
self._datastore: dict[ContextDataKey, Any] = {}
if value is None:
del self._datastore[key]
else:
self._datastore[key] = value
def __getitem__[T](self, key: ContextDataKey[T]) -> T:
"""
Get a value inside the context data storage.
Raises RuntimeError if the value is not present.
"""
if not hasattr(self, "_datastore") or key not in self._datastore:
raise RuntimeError(f"Context does not contain {repr(key)}")
return self._datastore[key]
def __contains__[T](self, key: ContextDataKey[T]) -> bool:
"""Check if context data storage contains the value"""
if not hasattr(self, "_datastore"):
return False
if key not in self._datastore:
return False
return True

View File

@@ -5,6 +5,8 @@ from typing import Any, Type
from nio import AsyncClient
from nio import MatrixRoom, Event
from ..context import EventContext
class BaseEventFilter(ABC):
"""Base class for all message filters"""
_logger = logging.Logger("EventFilter")
@@ -35,7 +37,7 @@ class BaseEventFilter(ABC):
)
def __ror__(self, other):
return self.__ror__(other)
return self.__or__(other)
# XOR
def __xor__(self, other):
@@ -64,7 +66,7 @@ class BaseEventFilter(ABC):
"""
return str(self.__class__.__name__)
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
async def __call__(self, context: EventContext) -> bool:
"""
This abstract method must be redefined in derived classes so that the
filter operates according to its description. This method must not raise
@@ -72,9 +74,9 @@ class BaseEventFilter(ABC):
and return False
Args:
- room - room the event has happened in
- event - the event to check againts this filter
- client - the client
- context - event context; your derived classes may add variables
to it (see `message.MessageTypeFilter` implementation
for reference)
Returns:
- True if the event satisfies this filter
@@ -88,8 +90,10 @@ class EventTypeFilter(BaseEventFilter):
super().__init__(**kwargs)
self._type = type
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
return isinstance(event, self._type)
async def __call__(self, context: EventContext) -> bool:
if not await super().__call__(context):
return False
return isinstance(context.event, self._type)
class CompoundEventFilter(BaseEventFilter):
"""Event filter that consists of multiple filters"""
@@ -140,8 +144,10 @@ class CompoundEventFilter(BaseEventFilter):
expression = f"~{reprs[0]}"
return f"({expression})"
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
evaluated = [await arg(room, event, client) for arg in self._arguments]
async def __call__(self, context: EventContext) -> bool:
if not await super().__call__(context):
return False
evaluated = [await arg(context) for arg in self._arguments]
if self._operator == self.OPERATOR_AND:
return all(evaluated)
elif self._operator == self.OPERATOR_OR:

View File

@@ -1,6 +1,13 @@
import re
import traceback
from .message import NewMessageFilter
from ..context import EventContext
from ..context import (
CTX_BODY,
CTX_CMD_PREFIX,
CTX_CMD_VERB,
CTX_CMD_ARGS
)
from nio import AsyncClient
from nio import MatrixRoom, Event
@@ -21,24 +28,27 @@ class BodyExistsFilter(NewMessageFilter):
This filter will match any message that has `body` in it, including images,
videos, files, etc.
This filter sets `CTX_BODY` context variable.
"""
def __init__(self, *, ignore_filename_in_body: bool = True, **kwargs):
super().__init__(**kwargs)
self._ignore_filename_in_body = ignore_filename_in_body
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
if not await super().__call__(room, event, client):
async def __call__(self, context: EventContext) -> bool:
if not await super().__call__(context):
return False
if not hasattr(event, "body"):
if not hasattr(context.event, "body"):
return False
if not isinstance(event.body, str): # type: ignore
if not isinstance(context.event.body, str): # type: ignore
return False
if not event.body.strip(): # type: ignore
if not context.event.body.strip(): # type: ignore
return False
if self._ignore_filename_in_body:
content = event.source["content"]
if "filename" in content and content["filename"] == event.body: # type: ignore
content = context.event.source["content"]
if "filename" in content and content["filename"] == context.event.body: # type: ignore
return False
context[CTX_BODY] = context.event.body # type: ignore
return True
class BodyContainsFilter(BodyExistsFilter):
@@ -61,10 +71,10 @@ class BodyContainsFilter(BodyExistsFilter):
self._any_case = any_case
self._needle = needle
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
if not await super().__call__(room, event, client):
async def __call__(self, context: EventContext) -> bool:
if not await super().__call__(context):
return False
body = event.body.lower() if self._any_case else event.body # type: ignore
body = context.event.body.lower() if self._any_case else context.event.body # type: ignore
for n in self._needle:
if n in body:
return True
@@ -90,10 +100,10 @@ class BodyStartsWithFilter(BodyExistsFilter):
self._any_case = any_case
self._substring = substring
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
if not await super().__call__(room, event, client):
async def __call__(self, context: EventContext) -> bool:
if not await super().__call__(context):
return False
body = event.body.lower() if self._any_case else event.body # type: ignore
body = context.event.body.lower() if self._any_case else context.event.body # type: ignore
for s in self._substring:
if body.startswith(s):
return True
@@ -119,10 +129,10 @@ class BodyEndsWithFilter(BodyExistsFilter):
self._any_case = any_case
self._substring = substring
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
if not await super().__call__(room, event, client):
async def __call__(self, context: EventContext) -> bool:
if not await super().__call__(context):
return False
body = event.body.lower() if self._any_case else event.body # type: ignore
body = context.event.body.lower() if self._any_case else context.event.body # type: ignore
for s in self._substring:
if body.endswith(s):
return True
@@ -145,9 +155,10 @@ class BodyCommandFilter(BodyExistsFilter):
store all verbs in lower case. This filter will not match any verbs that
use mixed case of upper case.
If this filter is matched, then it will set a new attribute for the event:
`event.command_args: list[str]`. You may use this attribute in your callback
for this event.
This filter sets the following context variables:
- `CTX_CMD_PREFIX` - prefix that was used
- `CTX_CMD_VERB` - verb that was used
- `CTX_CMD_ARGS` - arguments that were passed
"""
def __init__(self, verbs: str | list[str], min_args: int = 0, max_args: int | None = None, prefix: str = "!", **kwargs):
super().__init__(**kwargs)
@@ -158,10 +169,10 @@ class BodyCommandFilter(BodyExistsFilter):
self._max_args = max_args
self._prefix = prefix
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
if not await super().__call__(room, event, client):
async def __call__(self, context: EventContext) -> bool:
if not await super().__call__(context):
return False
parts = [p.strip() for p in event.body.split() if p.strip()] # type: ignore
parts = [p.strip() for p in context.event.body.split() if p.strip()] # type: ignore
args_count = len(parts) - 1
if args_count < self._min_args:
return False
@@ -172,7 +183,9 @@ class BodyCommandFilter(BodyExistsFilter):
cmd = parts[0][len(self._prefix):].lower()
for verb in self._verbs:
if cmd == verb:
setattr(event, "command_args", parts[1:])
context[CTX_CMD_PREFIX] = self._prefix
context[CTX_CMD_VERB] = verb
context[CTX_CMD_ARGS] = parts[1:]
return True
return False
@@ -186,11 +199,11 @@ class BodyRegexFilter(BodyExistsFilter):
regex = re.compile(regex)
self._regex = regex
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
if not await super().__call__(room, event, client):
async def __call__(self, context: EventContext) -> bool:
if not await super().__call__(context):
return False
try:
return self._regex.match(event.body) # type: ignore
return self._regex.match(context.event.body) is not None # type: ignore
except:
self._logger.error(traceback.format_exc())
return False

View File

@@ -1,6 +1,12 @@
import traceback
from .base import BaseEventFilter, EventTypeFilter
from ..types import MessageType
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
@@ -15,6 +21,8 @@ class MessageTypeFilter(BaseEventFilter):
`types` list is stored by reference so you may modify the behavior of this
filter dynamically.
This filter sets `CTX_MESSAGE_TYPE` variable in the context.
"""
def __init__(self, types: list[MessageType] | MessageType, **kwargs):
super().__init__(**kwargs)
@@ -22,14 +30,16 @@ class MessageTypeFilter(BaseEventFilter):
types = [types]
self._types = types
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
if not await super().__call__(room, event, client):
async def __call__(self, context: EventContext) -> bool:
if not await super().__call__(context):
return False
if "msgtype" not in event.source["content"]:
if "msgtype" not in context.event.source["content"]:
return False
return (
event.source["content"]["msgtype"] in [t.value for t in self._types]
)
msgtype = context.event.source["content"]["msgtype"]
if not msgtype in [t.value for t in self._types]:
return False
context[CTX_MESSAGE_TYPE] = MessageType(msgtype)
return True
class NewMessageFilter(BaseEventFilter):
"""
@@ -40,10 +50,10 @@ class NewMessageFilter(BaseEventFilter):
def __init__(self, **kwargs):
super().__init__(**kwargs)
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
if not await super().__call__(room, event, client):
async def __call__(self, context: EventContext) -> bool:
if not await super().__call__(context):
return False
return "m.new_content" not in event.source["content"]
return "m.new_content" not in context.event.source["content"]
class EditedMessageFilter(BaseEventFilter):
"""
@@ -53,10 +63,10 @@ class EditedMessageFilter(BaseEventFilter):
def __init__(self, **kwargs):
super().__init__(**kwargs)
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
if not await super().__call__(room, event, client):
async def __call__(self, context: EventContext) -> bool:
if not await super().__call__(context):
return False
return "m.new_content" in event.source["content"]
return "m.new_content" in context.event.source["content"]
class RedactedMessageFilter(EventTypeFilter):
"""
@@ -64,9 +74,10 @@ class RedactedMessageFilter(EventTypeFilter):
"""
def __init__(self, **kwargs):
super().__init__(RedactionEvent, **kwargs)
raise NotImplementedError()
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
return await super().__call__(room, event, client)
async def __call__(self, context: EventContext) -> bool:
raise NotImplementedError()
class SenderIsFilter(BaseEventFilter):
"""
@@ -77,6 +88,8 @@ class SenderIsFilter(BaseEventFilter):
`senders` list is stored by reference so you can modify behavior of this
filter dynamically.
This filter sets `CTX_SENDER` variable in the context.
"""
def __init__(self, sender: list[str] | str, *, any_case: bool = True, **kwargs):
super().__init__(**kwargs)
@@ -85,12 +98,13 @@ class SenderIsFilter(BaseEventFilter):
self._sender = sender
self._any_case = any_case
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
if not await super().__call__(room, event, client):
async def __call__(self, context: EventContext) -> bool:
if not await super().__call__(context):
return False
sender = event.sender.lower() if self._any_case else event.sender
sender = context.event.sender.lower() if self._any_case else context.event.sender
for s in self._sender:
if sender == s:
context[CTX_SENDER] = sender
return True
return False
@@ -106,7 +120,32 @@ class SenderIsBotFilter(BaseEventFilter):
def __init__(self, **kwargs):
super().__init__(**kwargs)
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
if not await super().__call__(room, event, client):
async def __call__(self, context: EventContext) -> bool:
if not await super().__call__(context):
return False
return client.user_id == 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

View File

@@ -1,7 +1,6 @@
from .base import BaseEventFilter
from nio import AsyncClient
from nio import MatrixRoom, Event
from ..context import EventContext, CTX_ROOM_ENCRYPTED
class RoomEncryptedFilter(BaseEventFilter):
"""
@@ -10,8 +9,11 @@ class RoomEncryptedFilter(BaseEventFilter):
def __init__(self):
super().__init__()
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
async def __call__(self, context: EventContext) -> bool:
if not await super().__call__(context):
return False
try:
return room.encrypted
context[CTX_ROOM_ENCRYPTED] = context.room.encrypted
return context.room.encrypted
except:
return False

View File

@@ -4,15 +4,8 @@ from pathlib import Path
from dataclasses import dataclass
from enum import Enum
from nio import MatrixRoom, Event
from nio import UploadResponse
from .filters.base import BaseEventFilter
from typing import TYPE_CHECKING
if TYPE_CHECKING:
from .bot import MatrixBot
@dataclass
class MatrixBotConfig:
"""Configuration for MatrixBot"""
@@ -68,22 +61,6 @@ class VideoFileProperties:
thumbnail: Path | str | bytes | None = None
"""Path to the thumbnail or the raw JPEG thumbnail data"""
@dataclass
class RoomEventData:
"""Dataclass that hold information about event that happened in the room"""
room: MatrixRoom
"""The room the event has happened in"""
event: Event
"""The event that has happened in the room"""
filter: BaseEventFilter
"""The filter that invoked this event"""
bot: "MatrixBot"
"""The bot that is the source of the event"""
@dataclass
class UploadResult:
"""Result of data upload"""
@@ -108,4 +85,21 @@ class MessageType(Enum):
FILE = "m.file"
AUDIO = "m.audio"
LOCATION = "m.location"
VIDEO = "m.video"
VIDEO = "m.video"
class ContextDataKey[T]:
"""
Instances of this class represent a single possible data key that can be
stored inside EventContext.
"""
def __init__(self, name: str) -> None:
"""
Initialize a ContextDataKey
Args:
- name - name that will be used internally
"""
self._name = name
def __repr__(self) -> str:
return f"ContextDataKey[{type(T)}]({repr(self._name)})"