Compare commits
22 Commits
159a43ebe6
...
v0.5.2
| Author | SHA1 | Date | |
|---|---|---|---|
| 16b849ebd8 | |||
| 35bd68516a | |||
| 86257a70e3 | |||
| 43b5d8e3a6 | |||
| 0e9b621895 | |||
| f3f1e24c0b | |||
| 2e4228b22a | |||
| 0b79d78715 | |||
| a90cd66d3e | |||
| cb520814d8 | |||
| 4287a0d20c | |||
| e10c920a56 | |||
| 156afd6b61 | |||
| 2d47483d55 | |||
| fc4f664a5c | |||
| fa97e4b098 | |||
| 1fe4434e16 | |||
| 4f0792b9aa | |||
| 31fbcb4697 | |||
| 3e15ae426c | |||
| b8e598715c | |||
| f58c8601d1 |
2
.gitignore
vendored
2
.gitignore
vendored
@@ -1,7 +1,9 @@
|
|||||||
__pycache__/
|
__pycache__/
|
||||||
|
session_storage/
|
||||||
*.vscode
|
*.vscode
|
||||||
.venv/
|
.venv/
|
||||||
dist/
|
dist/
|
||||||
*.egg-info/
|
*.egg-info/
|
||||||
*.swp
|
*.swp
|
||||||
*.swo
|
*.swo
|
||||||
|
*.tmp
|
||||||
153
README.md
153
README.md
@@ -1,69 +1,122 @@
|
|||||||
# mab
|
# 🤖 mab
|
||||||
|
|
||||||
**mab** *(MAtrix Bot)* is a **very** simple Python package that can be used to
|
**mab** *(MAtrix Bot)* is a **very** simple Python package that can be used to
|
||||||
develop **very** simple Matrix bots. I have decided to make something like this
|
develop **very** simple Matrix bots. It does not aim to be the best library out
|
||||||
because I wasn't satisfied by simplicity and usage of other libraries. So
|
there, but it aims to be convenient and usable for relatively serious projects.
|
||||||
this library does not aim to be "the best matrix bot library", it only aims to
|
|
||||||
be good enough for me.
|
|
||||||
|
|
||||||
## Features
|
## ✨ Features
|
||||||
|
|
||||||
The package supports the following features:
|
The library supports the following features:
|
||||||
|
- **Completely `asyncio` based**
|
||||||
- **Filter-based callback system**
|
- **Filter-based callback system**
|
||||||
- **Images sending**
|
- **Downloading and transparently decrypting files**
|
||||||
- **Videos sending with automatic thumbnail generation (requires `ffmpeg`)**
|
- **Sending images**
|
||||||
|
- **Sending videos with automatic thumbnail generation (requires `ffmpeg`)**
|
||||||
|
|
||||||
## Installation
|
## 📦 Installation
|
||||||
|
|
||||||
Use `pip` to install this package:
|
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
|
||||||
|
use `venv` or another Python virtual environment. Here are the commands to
|
||||||
|
install the latest version of the library:
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
python -m pip install git+https://git.tyukalov.su/nikita/mab@v0.3.0
|
apt install libmagic1-dev libolm-dev
|
||||||
|
python -m pip install git+https://git.tyukalov.su/nikita/mab@v0.5.1
|
||||||
```
|
```
|
||||||
|
|
||||||
You should specify package version you want to use, because `main` without tags
|
`libmagic1-dev` is needed for automatic file MIME type detection, `libolm-dev`
|
||||||
contain unstable code.
|
is needed for E2EE to work.
|
||||||
|
|
||||||
## Basic usage
|
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.
|
||||||
|
|
||||||
This is the most simple bot you can create. It would respond to any message
|
## 🚀 Usage
|
||||||
that starts with `!test`, `!hello` or `!hi`.
|
|
||||||
|
|
||||||
```python
|
If you use `mab`, your application will *most likely* be using **callbacks** to
|
||||||
import asyncio
|
react to user actions. `mab` uses filter-based callback system to avoid exposing
|
||||||
from mab import MatrixBot, MatrixBotConfig
|
raw `nio-matrix` event objects.
|
||||||
from mab import TextCommandFilter
|
|
||||||
from mab.types import RoomEventData
|
|
||||||
|
|
||||||
async def on_valid_command(data: RoomEventData) -> None:
|
This is the workflow you will most likely follow:
|
||||||
# do not respond to ourselves
|
1. **Define the callback as `async` function that take 1 argument of type
|
||||||
if event.sender == data.bot.get_client().user_id:
|
`EventContext`.** For example, this callback would print the caption of the
|
||||||
return
|
message:
|
||||||
text = f"Your message contains {len(data.event.body)} symbols"
|
```python
|
||||||
await data.bot.send_text_to_room(data.room, text)
|
from mab import *
|
||||||
|
|
||||||
async def main() -> None:
|
async def on_media_with_body(ctx: EventContext):
|
||||||
# create and start the bot
|
"""To be called when a message with image/video and caption is received."""
|
||||||
cfg = MatrixBotConfig(
|
print(ctx[CTX_BODY])
|
||||||
matrix_homeserver_url="matrix.domain.su",
|
```
|
||||||
matrix_username_localpart="nagibator666",
|
2. **Define the conditions your callback must be called on.** For example, you
|
||||||
storage_directory=Path("storage_nagibator666")
|
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()
|
||||||
)
|
)
|
||||||
bot = MatrixBot(matrix_bot_config)
|
```
|
||||||
bot.add_callback(
|
4. **Add the callback to your `MatrixBot` instance.** For example, if you would
|
||||||
TextCommandFilter(["test", "hello", "hi"]),
|
have used everything from above, then your code would look something like
|
||||||
on_valid_command
|
this:
|
||||||
)
|
```python
|
||||||
await bot.start()
|
from mab import *
|
||||||
# wait for Ctrl+C
|
|
||||||
try:
|
|
||||||
while True:
|
|
||||||
await asyncio.sleep(1)
|
|
||||||
except:
|
|
||||||
pass
|
|
||||||
# stop the bot
|
|
||||||
await bot.stop()
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
# let's assume you create your MatrixBot as `bot` variable here
|
||||||
asyncio.run(main())
|
|
||||||
|
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
|
||||||
|
is not tagged, it must be treated as versionless and should not be used for your
|
||||||
|
application.
|
||||||
|
- `X` **(Major)**: Breaking architectiral changes or complete rewrites. Existing
|
||||||
|
code will break. Note that `0.Y.Z` versions are considered **very unstable**,
|
||||||
|
the API may change at any time and some features do not work as expected.
|
||||||
|
- `Y` **(Minor)**: Breaking API changes, feature removals, or behavioral
|
||||||
|
modifications. Existing code will likely break.
|
||||||
|
- `Z` **(Patch)**: Backward-compatible feature additions, bug fixes, or internal
|
||||||
|
changes. Existing code will not break.
|
||||||
|
|
||||||
|
## 🛠️ Development
|
||||||
|
|
||||||
|
Here's the list of commands you should execute to get started with development
|
||||||
|
(including cloning the repository and installing required packages). Please note
|
||||||
|
that your workflow may use something other than `venv`.
|
||||||
|
```bash
|
||||||
|
apt install libmagic1-dev libolm-dev
|
||||||
|
git clone https://git.tyukalov.su/nikita/mab
|
||||||
|
cd mab
|
||||||
|
python3 -m venv .venv
|
||||||
|
. .venv/bin/activate
|
||||||
|
pip install -e .
|
||||||
```
|
```
|
||||||
18
examples/_environment.py
Normal file
18
examples/_environment.py
Normal file
@@ -0,0 +1,18 @@
|
|||||||
|
import os
|
||||||
|
import sys
|
||||||
|
|
||||||
|
def check_environment() -> None:
|
||||||
|
"""
|
||||||
|
This function checks if required environment variables are set. It prints
|
||||||
|
problem resolution guide and exits using `sys.exit(1)` on problem.
|
||||||
|
"""
|
||||||
|
if "MATRIX_HOMESERVER" not in os.environ:
|
||||||
|
print("Please set `MATRIX_HOMESERVER` environment variable!")
|
||||||
|
print("P.S. use something like this in your shell:")
|
||||||
|
print(" export MATRIX_HOMESERVER=\"https://matrix.server.net\"")
|
||||||
|
sys.exit(1)
|
||||||
|
if "MATRIX_USERNAME" not in os.environ:
|
||||||
|
print("Please set `MATRIX_USERNAME` environment variable!")
|
||||||
|
print("P.S. use something like this in your shell:")
|
||||||
|
print(" export MATRIX_USERNAME=\"megakiller228\"")
|
||||||
|
sys.exit(1)
|
||||||
130
examples/command_bot.py
Normal file
130
examples/command_bot.py
Normal 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())
|
||||||
43
examples/echo_bot.py
Normal file
43
examples/echo_bot.py
Normal file
@@ -0,0 +1,43 @@
|
|||||||
|
"""
|
||||||
|
This example implements Matrix bot that echoes all text messages it receives.
|
||||||
|
|
||||||
|
It uses environment variables to specify authorization data. Use Ctrl+C to stop
|
||||||
|
the bot.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import os
|
||||||
|
import logging
|
||||||
|
|
||||||
|
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, ctx[CTX_BODY])
|
||||||
|
|
||||||
|
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)
|
||||||
|
|
||||||
|
# run until Ctrl+C
|
||||||
|
try:
|
||||||
|
await bot.run()
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
pass
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
asyncio.run(main())
|
||||||
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())
|
||||||
129
examples/image_bot.py
Normal file
129
examples/image_bot.py
Normal file
@@ -0,0 +1,129 @@
|
|||||||
|
"""
|
||||||
|
This example implements Matrix bot that generates a pixelized noise image with
|
||||||
|
specified maximum R, G and B values.
|
||||||
|
|
||||||
|
It uses environment variables to specify authorization data. Use Ctrl+C to stop
|
||||||
|
the bot.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
from io import BytesIO
|
||||||
|
import os
|
||||||
|
import time
|
||||||
|
import logging
|
||||||
|
import random
|
||||||
|
from PIL import Image, ImageFilter
|
||||||
|
|
||||||
|
from mab import *
|
||||||
|
|
||||||
|
from _environment import check_environment
|
||||||
|
|
||||||
|
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[CTX_CMD_ARGS]]
|
||||||
|
except:
|
||||||
|
await data.bot.send_text(data.room, "Invalid arguments")
|
||||||
|
return
|
||||||
|
await data.bot.send_text(data.room, "Generating the noise...")
|
||||||
|
# create the basic noise
|
||||||
|
img = Image.new("RGB", (16, 16))
|
||||||
|
for x in range(img.width):
|
||||||
|
for y in range(img.height):
|
||||||
|
col = (random.random() * r, random.random() * g, random.random() * b)
|
||||||
|
img.putpixel(
|
||||||
|
(x, y),
|
||||||
|
tuple(int(c * 255) for c in col)
|
||||||
|
)
|
||||||
|
# pixelized upscale
|
||||||
|
img = img.resize((2048, 2048), resample=Image.Resampling.NEAREST)
|
||||||
|
# save to buffer
|
||||||
|
buf = BytesIO()
|
||||||
|
img.save(buf, format="PNG")
|
||||||
|
buf.seek(0)
|
||||||
|
buf = buf.read()
|
||||||
|
# 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> or send an image to blur"
|
||||||
|
)
|
||||||
|
|
||||||
|
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)
|
||||||
|
command_filter = MessageTypeFilter(MessageType.TEXT) & BodyCommandFilter(
|
||||||
|
verbs=["gen"],
|
||||||
|
min_args=3,
|
||||||
|
max_args=3
|
||||||
|
)
|
||||||
|
|
||||||
|
# callback for message that
|
||||||
|
# 1. are sent not by this bot
|
||||||
|
# 2. do match the command filter
|
||||||
|
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. are new messages (not edits)
|
||||||
|
bot.add_callback(
|
||||||
|
~SenderIsBotFilter() & BodyExistsFilter() & NewMessageFilter(),
|
||||||
|
on_wrong_message
|
||||||
|
)
|
||||||
|
|
||||||
|
# run until Ctrl+C
|
||||||
|
try:
|
||||||
|
await bot.run()
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
pass
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
asyncio.run(main())
|
||||||
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
|
|||||||
|
|
||||||
[project]
|
[project]
|
||||||
name = "mab"
|
name = "mab"
|
||||||
version = "0.3.0"
|
version = "0.5.2"
|
||||||
authors = [
|
authors = [
|
||||||
{ name = "Tyukalov Nikita", email = "nikita@tyukalov.su" }
|
{ name = "Tyukalov Nikita", email = "nikita@tyukalov.su" }
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -1,12 +1,14 @@
|
|||||||
from . import bot
|
from . import bot
|
||||||
from . import types
|
from . import types
|
||||||
|
|
||||||
from .types import MatrixBotConfig
|
from .types import MatrixBotConfig, MessageType
|
||||||
|
from .context import *
|
||||||
|
|
||||||
from .bot import MatrixBot
|
from .bot import MatrixBot
|
||||||
|
|
||||||
from .filters.base import *
|
from .filters.base import *
|
||||||
from .filters.text import *
|
from .filters.message import *
|
||||||
|
from .filters.body import *
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
# module names
|
# module names
|
||||||
@@ -15,18 +17,41 @@ __all__ = [
|
|||||||
|
|
||||||
# .types
|
# .types
|
||||||
"MatrixBotConfig",
|
"MatrixBotConfig",
|
||||||
|
"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
|
# .bot
|
||||||
"MatrixBot",
|
"MatrixBot",
|
||||||
|
|
||||||
# .filters.base
|
# .filters.base
|
||||||
"BaseEventFilter",
|
"BaseEventFilter",
|
||||||
|
"EventTypeFilter",
|
||||||
|
|
||||||
# .filters.text
|
# .filters.body
|
||||||
"TextFilter",
|
"BodyExistsFilter",
|
||||||
"FormattedTextFilter",
|
"BodyContainsFilter",
|
||||||
"TextContainsFilter",
|
"BodyStartsWithFilter",
|
||||||
"TextStartsWithFilter",
|
"BodyEndsWithFilter",
|
||||||
"TextEndsWithFilter",
|
"BodyCommandFilter",
|
||||||
"TextCommandFilter",
|
"BodyRegexFilter",
|
||||||
|
|
||||||
|
# .filters.message
|
||||||
|
"MessageTypeFilter",
|
||||||
|
"NewMessageFilter",
|
||||||
|
"EditedMessageFilter",
|
||||||
|
"RedactedMessageFilter",
|
||||||
|
"SenderIsFilter",
|
||||||
|
"SenderIsBotFilter",
|
||||||
|
"MessageHasFile",
|
||||||
]
|
]
|
||||||
@@ -10,7 +10,8 @@ from nio import MatrixInvitedRoom, InviteMemberEvent, JoinResponse
|
|||||||
from nio.events.room_events import Event as RoomEvent
|
from nio.events.room_events import Event as RoomEvent
|
||||||
|
|
||||||
from ._storage import Storage
|
from ._storage import Storage
|
||||||
from ..types import MatrixBotConfig, RoomEventData
|
from ..types import MatrixBotConfig
|
||||||
|
from ..context import EventContext
|
||||||
from ..filters.base import BaseEventFilter
|
from ..filters.base import BaseEventFilter
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
@@ -31,7 +32,7 @@ class Callbacks:
|
|||||||
filter: BaseEventFilter
|
filter: BaseEventFilter
|
||||||
"""Filter to use for matching"""
|
"""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"""
|
"""Callback that will be called if the filter matches"""
|
||||||
|
|
||||||
stop_matching: bool
|
stop_matching: bool
|
||||||
@@ -51,20 +52,19 @@ class Callbacks:
|
|||||||
for callback_info in self._filters:
|
for callback_info in self._filters:
|
||||||
if not isinstance(callback_info, self._FilterBasedCallback):
|
if not isinstance(callback_info, self._FilterBasedCallback):
|
||||||
continue
|
continue
|
||||||
|
event_data = EventContext(
|
||||||
|
room=room,
|
||||||
|
event=event,
|
||||||
|
bot=self._matrix_bot
|
||||||
|
)
|
||||||
try:
|
try:
|
||||||
if not await callback_info.filter(room, event, self._client):
|
if not await callback_info.filter(event_data):
|
||||||
continue
|
continue
|
||||||
except asyncio.CancelledError:
|
except asyncio.CancelledError:
|
||||||
raise
|
raise
|
||||||
except:
|
except:
|
||||||
self._logger.error(traceback.format_exc())
|
self._logger.error(traceback.format_exc())
|
||||||
continue
|
continue
|
||||||
event_data = RoomEventData(
|
|
||||||
room=room,
|
|
||||||
event=event,
|
|
||||||
filter=callback_info.filter,
|
|
||||||
bot=self._matrix_bot
|
|
||||||
)
|
|
||||||
try:
|
try:
|
||||||
# dump argument types
|
# dump argument types
|
||||||
if callback_info.callback is None:
|
if callback_info.callback is None:
|
||||||
@@ -144,7 +144,7 @@ class Callbacks:
|
|||||||
def add_room_event_callback(
|
def add_room_event_callback(
|
||||||
self,
|
self,
|
||||||
filter: BaseEventFilter,
|
filter: BaseEventFilter,
|
||||||
callback: Callable[[RoomEventData], Coroutine[Any, Any, None]] | None,
|
callback: Callable[[EventContext], Coroutine[Any, Any, None]] | None,
|
||||||
*,
|
*,
|
||||||
stop_matching: bool = True) -> None:
|
stop_matching: bool = True) -> None:
|
||||||
"""
|
"""
|
||||||
|
|||||||
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
|
||||||
@@ -191,11 +191,13 @@ class ClientManager:
|
|||||||
raise RuntimeError("The bot was never started")
|
raise RuntimeError("The bot was never started")
|
||||||
self._background_task.cancel()
|
self._background_task.cancel()
|
||||||
try:
|
try:
|
||||||
asyncio.shield(self._background_task)
|
await asyncio.shield(self._background_task)
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
pass
|
||||||
except:
|
except:
|
||||||
self._logger.error(traceback.format_exc())
|
self._logger.error(traceback.format_exc())
|
||||||
try:
|
try:
|
||||||
asyncio.shield(self._close_client())
|
await asyncio.shield(self._close_client())
|
||||||
except:
|
except:
|
||||||
self._logger.error(traceback.format_exc())
|
self._logger.error(traceback.format_exc())
|
||||||
self._background_task = None
|
self._background_task = None
|
||||||
|
|||||||
@@ -1,4 +1,5 @@
|
|||||||
import asyncio
|
import asyncio
|
||||||
|
from io import BytesIO
|
||||||
import logging
|
import logging
|
||||||
from html.parser import HTMLParser
|
from html.parser import HTMLParser
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
@@ -214,6 +215,69 @@ class ClientSender:
|
|||||||
}
|
}
|
||||||
return (await self.send_content(room, content)).event_id
|
return (await self.send_content(room, content)).event_id
|
||||||
|
|
||||||
|
async def send_image_bytes(self,
|
||||||
|
room: MatrixRoom | str,
|
||||||
|
data: bytes,
|
||||||
|
filename: str,
|
||||||
|
*,
|
||||||
|
text: str | None = None,
|
||||||
|
is_html: bool | None = None,
|
||||||
|
timeout: float | None = 60 * 60) -> str:
|
||||||
|
"""
|
||||||
|
Send the image 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 text to
|
||||||
|
- bytes - the image to send
|
||||||
|
- filename - filename to use for the file
|
||||||
|
- text - image 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
|
||||||
|
mime_type: str = magic.from_buffer(data, mime=True)
|
||||||
|
if not mime_type.startswith("image/"):
|
||||||
|
raise RuntimeError(f"Data has non-image mime-type")
|
||||||
|
# get image size
|
||||||
|
buffer = BytesIO(data)
|
||||||
|
with Image.open(buffer) as image:
|
||||||
|
width, height = image.size
|
||||||
|
buffer.seek(0)
|
||||||
|
# upload
|
||||||
|
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.image",
|
||||||
|
"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,
|
||||||
|
"w": width,
|
||||||
|
"h": height
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return (await self.send_content(room, content)).event_id
|
||||||
|
|
||||||
async def send_video(self,
|
async def send_video(self,
|
||||||
room: MatrixRoom | str,
|
room: MatrixRoom | str,
|
||||||
path: Path | str,
|
path: Path | str,
|
||||||
@@ -298,3 +362,122 @@ class ClientSender:
|
|||||||
}
|
}
|
||||||
# send
|
# send
|
||||||
return (await self.send_content(room, content)).event_id
|
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,15 +1,18 @@
|
|||||||
|
import asyncio
|
||||||
import logging
|
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 ..filters.base import BaseEventFilter
|
||||||
from ..types import *
|
from ..types import *
|
||||||
|
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
|
||||||
@@ -31,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
|
||||||
@@ -43,7 +47,7 @@ class MatrixBot:
|
|||||||
|
|
||||||
def add_callback(self,
|
def add_callback(self,
|
||||||
filter: BaseEventFilter,
|
filter: BaseEventFilter,
|
||||||
callback: Callable[[RoomEventData], Coroutine[Any, Any, None]] | None,
|
callback: Callable[[EventContext], Coroutine[Any, Any, None]] | None,
|
||||||
*,
|
*,
|
||||||
stop_matching: bool = True) -> None:
|
stop_matching: bool = True) -> None:
|
||||||
"""
|
"""
|
||||||
@@ -78,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(),
|
||||||
@@ -97,6 +105,20 @@ class MatrixBot:
|
|||||||
"""
|
"""
|
||||||
await self._client_manager.stop()
|
await self._client_manager.stop()
|
||||||
|
|
||||||
|
async def run(self) -> None:
|
||||||
|
"""
|
||||||
|
Start bot operation in foreground. You may cancel task running this
|
||||||
|
method to stop the bot.
|
||||||
|
|
||||||
|
Warning: calling `stop()` is not a supported way to stop the bot. You
|
||||||
|
should cancel this task instead.
|
||||||
|
"""
|
||||||
|
await self.start()
|
||||||
|
try:
|
||||||
|
await asyncio.Event().wait()
|
||||||
|
finally:
|
||||||
|
await self.stop()
|
||||||
|
|
||||||
def get_client(self) -> AsyncClient:
|
def get_client(self) -> AsyncClient:
|
||||||
"""
|
"""
|
||||||
Get AsyncClient.
|
Get AsyncClient.
|
||||||
@@ -162,6 +184,39 @@ class MatrixBot:
|
|||||||
timeout=timeout
|
timeout=timeout
|
||||||
)
|
)
|
||||||
|
|
||||||
|
async def send_image_bytes(self,
|
||||||
|
room: MatrixRoom | str,
|
||||||
|
data: bytes,
|
||||||
|
filename: str,
|
||||||
|
*,
|
||||||
|
text: str | None = None,
|
||||||
|
is_html: bool | None = None,
|
||||||
|
timeout: float | None = 60 * 60) -> str:
|
||||||
|
"""
|
||||||
|
Send the image 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 text to
|
||||||
|
- bytes - the image to send
|
||||||
|
- filename - filename to use for the file
|
||||||
|
- text - image 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_image_bytes(
|
||||||
|
room=room,
|
||||||
|
data=data,
|
||||||
|
filename=filename,
|
||||||
|
text=text,
|
||||||
|
is_html=is_html,
|
||||||
|
timeout=timeout
|
||||||
|
)
|
||||||
|
|
||||||
async def send_video(self,
|
async def send_video(self,
|
||||||
room: MatrixRoom | str,
|
room: MatrixRoom | str,
|
||||||
path: Path | str,
|
path: Path | str,
|
||||||
@@ -196,3 +251,109 @@ class MatrixBot:
|
|||||||
is_html=is_html,
|
is_html=is_html,
|
||||||
timeout=timeout
|
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
88
src/mab/context.py
Normal 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
|
||||||
@@ -1,13 +1,20 @@
|
|||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
import logging
|
import logging
|
||||||
|
from typing import Any, Type
|
||||||
|
|
||||||
from nio import AsyncClient
|
from nio import AsyncClient
|
||||||
from nio import MatrixRoom, Event
|
from nio import MatrixRoom, Event
|
||||||
|
|
||||||
|
from ..context import EventContext
|
||||||
|
|
||||||
class BaseEventFilter(ABC):
|
class BaseEventFilter(ABC):
|
||||||
"""Base class for all message filters"""
|
"""Base class for all message filters"""
|
||||||
_logger = logging.Logger("EventFilter")
|
_logger = logging.Logger("EventFilter")
|
||||||
|
|
||||||
|
def __init__(self, **kwargs):
|
||||||
|
for a in kwargs:
|
||||||
|
self._logger.critical(f"Unknown keyword argument for some filter is used: '{a}'={repr(kwargs[a])}")
|
||||||
|
|
||||||
# AND
|
# AND
|
||||||
def __and__(self, other):
|
def __and__(self, other):
|
||||||
if not isinstance(other, BaseEventFilter):
|
if not isinstance(other, BaseEventFilter):
|
||||||
@@ -30,7 +37,7 @@ class BaseEventFilter(ABC):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def __ror__(self, other):
|
def __ror__(self, other):
|
||||||
return self.__ror__(other)
|
return self.__or__(other)
|
||||||
|
|
||||||
# XOR
|
# XOR
|
||||||
def __xor__(self, other):
|
def __xor__(self, other):
|
||||||
@@ -52,30 +59,41 @@ class BaseEventFilter(ABC):
|
|||||||
)
|
)
|
||||||
|
|
||||||
# PAYLOAD
|
# PAYLOAD
|
||||||
@abstractmethod
|
|
||||||
def __repr__(self) -> str:
|
def __repr__(self) -> str:
|
||||||
"""This method must be redefined in derived classes to improve
|
|
||||||
debugging experience.
|
|
||||||
"""
|
"""
|
||||||
pass
|
This method may be redefined in derived classes to improve debugging
|
||||||
|
experience.
|
||||||
|
"""
|
||||||
|
return str(self.__class__.__name__)
|
||||||
|
|
||||||
@abstractmethod
|
async def __call__(self, context: EventContext) -> bool:
|
||||||
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
|
"""
|
||||||
"""This abstract method must be redefined in derived classes so that
|
This abstract method must be redefined in derived classes so that the
|
||||||
the filter operates according to its description. This method must
|
filter operates according to its description. This method must not raise
|
||||||
not raise exceptions. In case of exception it should log it using
|
exceptions. In case of exception it should log it using `self._logger`
|
||||||
`self._logger` and return False
|
and return False
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
room - room the event has happened in
|
- context - event context; your derived classes may add variables
|
||||||
event - the event to check againts this filter
|
to it (see `message.MessageTypeFilter` implementation
|
||||||
client - the client
|
for reference)
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
True if the event satisfies this filter
|
- True if the event satisfies this filter
|
||||||
False if the event does not satisfy this filter
|
- False if the event does not satisfy this filter
|
||||||
"""
|
"""
|
||||||
pass
|
return True
|
||||||
|
|
||||||
|
class EventTypeFilter(BaseEventFilter):
|
||||||
|
"""Event filter that checks if the event is an instance of some class"""
|
||||||
|
def __init__(self, type: Type, **kwargs):
|
||||||
|
super().__init__(**kwargs)
|
||||||
|
self._type = 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):
|
class CompoundEventFilter(BaseEventFilter):
|
||||||
"""Event filter that consists of multiple filters"""
|
"""Event filter that consists of multiple filters"""
|
||||||
@@ -104,8 +122,8 @@ class CompoundEventFilter(BaseEventFilter):
|
|||||||
CompoundEventFilter.OPERATOR_INVERT: [1],
|
CompoundEventFilter.OPERATOR_INVERT: [1],
|
||||||
}[op]
|
}[op]
|
||||||
|
|
||||||
def __init__(self, operator: str, arguments: list[BaseEventFilter]):
|
def __init__(self, operator: str, arguments: list[BaseEventFilter], **kwargs):
|
||||||
super().__init__()
|
super().__init__(**kwargs)
|
||||||
if not self._is_operator_valid(operator):
|
if not self._is_operator_valid(operator):
|
||||||
raise RuntimeError(f"Invalid operator `{operator}`")
|
raise RuntimeError(f"Invalid operator `{operator}`")
|
||||||
if not self._is_elements_count_valid(operator, len(arguments)):
|
if not self._is_elements_count_valid(operator, len(arguments)):
|
||||||
@@ -126,8 +144,10 @@ class CompoundEventFilter(BaseEventFilter):
|
|||||||
expression = f"~{reprs[0]}"
|
expression = f"~{reprs[0]}"
|
||||||
return f"({expression})"
|
return f"({expression})"
|
||||||
|
|
||||||
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
|
async def __call__(self, context: EventContext) -> bool:
|
||||||
evaluated = [await arg(room, event, client) for arg in self._arguments]
|
if not await super().__call__(context):
|
||||||
|
return False
|
||||||
|
evaluated = [await arg(context) for arg in self._arguments]
|
||||||
if self._operator == self.OPERATOR_AND:
|
if self._operator == self.OPERATOR_AND:
|
||||||
return all(evaluated)
|
return all(evaluated)
|
||||||
elif self._operator == self.OPERATOR_OR:
|
elif self._operator == self.OPERATOR_OR:
|
||||||
|
|||||||
@@ -1,11 +1,18 @@
|
|||||||
import re
|
import re
|
||||||
import traceback
|
import traceback
|
||||||
from .base import BaseEventFilter
|
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 AsyncClient
|
||||||
from nio import MatrixRoom, Event
|
from nio import MatrixRoom, Event
|
||||||
|
|
||||||
class TextFilter(BaseEventFilter):
|
class BodyExistsFilter(NewMessageFilter):
|
||||||
"""
|
"""
|
||||||
This filter returns True if all conditions are met:
|
This filter returns True if all conditions are met:
|
||||||
1. `event` has attribute `body`
|
1. `event` has attribute `body`
|
||||||
@@ -18,53 +25,33 @@ class TextFilter(BaseEventFilter):
|
|||||||
If `event.body` value equals to `event.source["content"]["filename"]` (if it
|
If `event.body` value equals to `event.source["content"]["filename"]` (if it
|
||||||
is present, of course) then this filter will not match it by default. You
|
is present, of course) then this filter will not match it by default. You
|
||||||
may disable `ignore_filename_in_body` to disable this feature.
|
may disable `ignore_filename_in_body` to disable this feature.
|
||||||
|
|
||||||
|
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):
|
def __init__(self, *, ignore_filename_in_body: bool = True, **kwargs):
|
||||||
super().__init__(**kwargs)
|
super().__init__(**kwargs)
|
||||||
self._ignore_filename_in_body = ignore_filename_in_body
|
self._ignore_filename_in_body = ignore_filename_in_body
|
||||||
|
|
||||||
def __repr__(self) -> str:
|
async def __call__(self, context: EventContext) -> bool:
|
||||||
return "TextFilter()"
|
if not await super().__call__(context):
|
||||||
|
|
||||||
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
|
|
||||||
if not hasattr(event, "body"):
|
|
||||||
return False
|
return False
|
||||||
if not isinstance(event.body, str): # type: ignore
|
if not hasattr(context.event, "body"):
|
||||||
return False
|
return False
|
||||||
if not event.body.strip(): # type: ignore
|
if not isinstance(context.event.body, str): # type: ignore
|
||||||
|
return False
|
||||||
|
if not context.event.body.strip(): # type: ignore
|
||||||
return False
|
return False
|
||||||
if self._ignore_filename_in_body:
|
if self._ignore_filename_in_body:
|
||||||
content = event.source["content"]
|
content = context.event.source["content"]
|
||||||
if "filename" in content and content["filename"] == event.body: # type: ignore
|
if "filename" in content and content["filename"] == context.event.body: # type: ignore
|
||||||
return False
|
return False
|
||||||
|
context[CTX_BODY] = context.event.body # type: ignore
|
||||||
return True
|
return True
|
||||||
|
|
||||||
class FormattedTextFilter(BaseEventFilter):
|
class BodyContainsFilter(BodyExistsFilter):
|
||||||
"""
|
|
||||||
This filter returns True if all conditions are met:
|
|
||||||
1. `event` has attribute `formatted_body`
|
|
||||||
2. `event.formatted_body` is instance of `str`
|
|
||||||
3. `event.formatted_body.strip()` evaluates to True
|
|
||||||
|
|
||||||
If this filter matches, you can access `event.formatted_body` and it stores
|
|
||||||
formatted text of the message.
|
|
||||||
"""
|
|
||||||
def __init__(self, **kwargs):
|
|
||||||
super().__init__(**kwargs)
|
|
||||||
|
|
||||||
def __repr__(self):
|
|
||||||
return "FormattedTextFilter()"
|
|
||||||
|
|
||||||
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
|
|
||||||
if not hasattr(event, "formatted_body"):
|
|
||||||
return False
|
|
||||||
if not isinstance(event.formatted_body, str): # type: ignore
|
|
||||||
return False
|
|
||||||
if not event.formatted_body.strip(): # type: ignore
|
|
||||||
return False
|
|
||||||
return True
|
|
||||||
|
|
||||||
class TextContainsFilter(TextFilter):
|
|
||||||
"""
|
"""
|
||||||
This filter returns True if `event.body` contains `needle` substring (or any
|
This filter returns True if `event.body` contains `needle` substring (or any
|
||||||
of neddle from the list). `event.body` will be converted to lower case if
|
of neddle from the list). `event.body` will be converted to lower case if
|
||||||
@@ -84,19 +71,16 @@ class TextContainsFilter(TextFilter):
|
|||||||
self._any_case = any_case
|
self._any_case = any_case
|
||||||
self._needle = needle
|
self._needle = needle
|
||||||
|
|
||||||
def __repr__(self) -> str:
|
async def __call__(self, context: EventContext) -> bool:
|
||||||
return f"TextContainsFilter({repr(self._needle)}, any_case={repr(self._any_case)})"
|
if not await super().__call__(context):
|
||||||
|
|
||||||
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
|
|
||||||
if not await super().__call__(room, event, client):
|
|
||||||
return False
|
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:
|
for n in self._needle:
|
||||||
if n in body:
|
if n in body:
|
||||||
return True
|
return True
|
||||||
return False
|
return False
|
||||||
|
|
||||||
class TextStartsWithFilter(TextFilter):
|
class BodyStartsWithFilter(BodyExistsFilter):
|
||||||
"""
|
"""
|
||||||
This filter returns True if `event.body` starts with `substring` (or any of
|
This filter returns True if `event.body` starts with `substring` (or any of
|
||||||
substrings from the list). The check will be case insensetive if `any_case`
|
substrings from the list). The check will be case insensetive if `any_case`
|
||||||
@@ -116,19 +100,16 @@ class TextStartsWithFilter(TextFilter):
|
|||||||
self._any_case = any_case
|
self._any_case = any_case
|
||||||
self._substring = substring
|
self._substring = substring
|
||||||
|
|
||||||
def __repr__(self):
|
async def __call__(self, context: EventContext) -> bool:
|
||||||
return f"TextStartsWithFilter({repr(self._substring)}, any_case={repr(self._any_case)})"
|
if not await super().__call__(context):
|
||||||
|
|
||||||
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
|
|
||||||
if not await super().__call__(room, event, client):
|
|
||||||
return False
|
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:
|
for s in self._substring:
|
||||||
if body.startswith(s):
|
if body.startswith(s):
|
||||||
return True
|
return True
|
||||||
return False
|
return False
|
||||||
|
|
||||||
class TextEndsWithFilter(TextFilter):
|
class BodyEndsWithFilter(BodyExistsFilter):
|
||||||
"""
|
"""
|
||||||
This filter returns True if `event.body` ends with `substring` (or any of
|
This filter returns True if `event.body` ends with `substring` (or any of
|
||||||
substrings from the list). The check will be case insensetive if `any_case`
|
substrings from the list). The check will be case insensetive if `any_case`
|
||||||
@@ -148,19 +129,16 @@ class TextEndsWithFilter(TextFilter):
|
|||||||
self._any_case = any_case
|
self._any_case = any_case
|
||||||
self._substring = substring
|
self._substring = substring
|
||||||
|
|
||||||
def __repr__(self):
|
async def __call__(self, context: EventContext) -> bool:
|
||||||
return f"TextEndsWithFilter({repr(self._substring)}, any_case={repr(self._any_case)})"
|
if not await super().__call__(context):
|
||||||
|
|
||||||
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
|
|
||||||
if not await super().__call__(room, event, client):
|
|
||||||
return False
|
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:
|
for s in self._substring:
|
||||||
if body.endswith(s):
|
if body.endswith(s):
|
||||||
return True
|
return True
|
||||||
return False
|
return False
|
||||||
|
|
||||||
class TextCommandFilter(TextFilter):
|
class BodyCommandFilter(BodyExistsFilter):
|
||||||
"""
|
"""
|
||||||
This filter returns True if all conditions are met:
|
This filter returns True if all conditions are met:
|
||||||
1. `event.body` contains at least `min_args + 1` words after split()
|
1. `event.body` contains at least `min_args + 1` words after split()
|
||||||
@@ -177,9 +155,10 @@ class TextCommandFilter(TextFilter):
|
|||||||
store all verbs in lower case. This filter will not match any verbs that
|
store all verbs in lower case. This filter will not match any verbs that
|
||||||
use mixed case of upper case.
|
use mixed case of upper case.
|
||||||
|
|
||||||
If this filter is matched, then it will set a new attribute for the event:
|
This filter sets the following context variables:
|
||||||
`event.command_args: list[str]`. You may use this attribute in your callback
|
- `CTX_CMD_PREFIX` - prefix that was used
|
||||||
for this event.
|
- `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):
|
def __init__(self, verbs: str | list[str], min_args: int = 0, max_args: int | None = None, prefix: str = "!", **kwargs):
|
||||||
super().__init__(**kwargs)
|
super().__init__(**kwargs)
|
||||||
@@ -190,13 +169,10 @@ class TextCommandFilter(TextFilter):
|
|||||||
self._max_args = max_args
|
self._max_args = max_args
|
||||||
self._prefix = prefix
|
self._prefix = prefix
|
||||||
|
|
||||||
def __repr__(self):
|
async def __call__(self, context: EventContext) -> bool:
|
||||||
return f"TextCommandFilter({repr(self._verbs)}, {repr(self._min_args)}, {repr(self._max_args)}, {repr(self._prefix)})"
|
if not await super().__call__(context):
|
||||||
|
|
||||||
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
|
|
||||||
if not await super().__call__(room, event, client):
|
|
||||||
return False
|
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
|
args_count = len(parts) - 1
|
||||||
if args_count < self._min_args:
|
if args_count < self._min_args:
|
||||||
return False
|
return False
|
||||||
@@ -207,11 +183,13 @@ class TextCommandFilter(TextFilter):
|
|||||||
cmd = parts[0][len(self._prefix):].lower()
|
cmd = parts[0][len(self._prefix):].lower()
|
||||||
for verb in self._verbs:
|
for verb in self._verbs:
|
||||||
if cmd == verb:
|
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 True
|
||||||
return False
|
return False
|
||||||
|
|
||||||
class TextRegexFilter(TextFilter):
|
class BodyRegexFilter(BodyExistsFilter):
|
||||||
"""
|
"""
|
||||||
This filter returns True if the `event.body` passes the regex.
|
This filter returns True if the `event.body` passes the regex.
|
||||||
"""
|
"""
|
||||||
@@ -221,14 +199,11 @@ class TextRegexFilter(TextFilter):
|
|||||||
regex = re.compile(regex)
|
regex = re.compile(regex)
|
||||||
self._regex = regex
|
self._regex = regex
|
||||||
|
|
||||||
def __repr__(self):
|
async def __call__(self, context: EventContext) -> bool:
|
||||||
return f"TextRegexFilter({repr(self._regex)})"
|
if not await super().__call__(context):
|
||||||
|
|
||||||
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
|
|
||||||
if not await super().__call__(room, event, client):
|
|
||||||
return False
|
return False
|
||||||
try:
|
try:
|
||||||
return self._regex.match(event.body) # type: ignore
|
return self._regex.match(context.event.body) is not None # type: ignore
|
||||||
except:
|
except:
|
||||||
self._logger.error(traceback.format_exc())
|
self._logger.error(traceback.format_exc())
|
||||||
return False
|
return False
|
||||||
151
src/mab/filters/message.py
Normal file
151
src/mab/filters/message.py
Normal file
@@ -0,0 +1,151 @@
|
|||||||
|
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
|
||||||
|
|
||||||
|
from nio import RedactionEvent
|
||||||
|
|
||||||
|
class MessageTypeFilter(BaseEventFilter):
|
||||||
|
"""
|
||||||
|
This filter should be used to match specific message types (text-only,
|
||||||
|
images, videos, files, etc) based on `event.source["content"]["msgtype"]`
|
||||||
|
value.
|
||||||
|
|
||||||
|
`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)
|
||||||
|
if isinstance(types, MessageType):
|
||||||
|
types = [types]
|
||||||
|
self._types = types
|
||||||
|
|
||||||
|
async def __call__(self, context: EventContext) -> bool:
|
||||||
|
if not await super().__call__(context):
|
||||||
|
return False
|
||||||
|
if "msgtype" not in context.event.source["content"]:
|
||||||
|
return False
|
||||||
|
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):
|
||||||
|
"""
|
||||||
|
This filter returns True if the event is a new message. Most filters are
|
||||||
|
derived from this base class because it ignores events about edited
|
||||||
|
messages.
|
||||||
|
"""
|
||||||
|
def __init__(self, **kwargs):
|
||||||
|
super().__init__(**kwargs)
|
||||||
|
|
||||||
|
async def __call__(self, context: EventContext) -> bool:
|
||||||
|
if not await super().__call__(context):
|
||||||
|
return False
|
||||||
|
return "m.new_content" not in context.event.source["content"]
|
||||||
|
|
||||||
|
class EditedMessageFilter(BaseEventFilter):
|
||||||
|
"""
|
||||||
|
This filter returns True if the event is an edited message. You may use this
|
||||||
|
filter to create callbacks that are called if the message gets edited.
|
||||||
|
"""
|
||||||
|
def __init__(self, **kwargs):
|
||||||
|
super().__init__(**kwargs)
|
||||||
|
|
||||||
|
async def __call__(self, context: EventContext) -> bool:
|
||||||
|
if not await super().__call__(context):
|
||||||
|
return False
|
||||||
|
return "m.new_content" in context.event.source["content"]
|
||||||
|
|
||||||
|
class RedactedMessageFilter(EventTypeFilter):
|
||||||
|
"""
|
||||||
|
This filter returns True if the event is a RedactionEvent.
|
||||||
|
"""
|
||||||
|
def __init__(self, **kwargs):
|
||||||
|
super().__init__(RedactionEvent, **kwargs)
|
||||||
|
raise NotImplementedError()
|
||||||
|
|
||||||
|
async def __call__(self, context: EventContext) -> bool:
|
||||||
|
raise NotImplementedError()
|
||||||
|
|
||||||
|
class SenderIsFilter(BaseEventFilter):
|
||||||
|
"""
|
||||||
|
This filter returns True if `event.sender` is any of specified senders.
|
||||||
|
`event.sender` is converted to lower case if `any_case` is True (default).
|
||||||
|
Supplied sender list is NEVER converted to lower case, so it is your duty to
|
||||||
|
use lower case if `any_case` is True.
|
||||||
|
|
||||||
|
`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)
|
||||||
|
if isinstance(sender, str):
|
||||||
|
sender = [sender]
|
||||||
|
self._sender = sender
|
||||||
|
self._any_case = any_case
|
||||||
|
|
||||||
|
async def __call__(self, context: EventContext) -> bool:
|
||||||
|
if not await super().__call__(context):
|
||||||
|
return False
|
||||||
|
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
|
||||||
|
|
||||||
|
class SenderIsBotFilter(BaseEventFilter):
|
||||||
|
"""
|
||||||
|
This filter returns True if `event.sender` is the client that has received
|
||||||
|
the event. You may use this filter to set callbacks for messages sent by
|
||||||
|
other users by using the following syntax:
|
||||||
|
```py
|
||||||
|
~SenderIsBotFilter()
|
||||||
|
```
|
||||||
|
"""
|
||||||
|
def __init__(self, **kwargs):
|
||||||
|
super().__init__(**kwargs)
|
||||||
|
|
||||||
|
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
|
||||||
|
|
||||||
|
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
|
||||||
@@ -1,94 +1,6 @@
|
|||||||
from .base import BaseEventFilter
|
from .base import BaseEventFilter
|
||||||
|
|
||||||
from nio import AsyncClient
|
from ..context import EventContext, CTX_ROOM_ENCRYPTED
|
||||||
from nio import MatrixRoom, Event
|
|
||||||
|
|
||||||
class RoomIdContainsFilter(BaseEventFilter):
|
|
||||||
"""
|
|
||||||
This filter returns True if the `room.room_id` contains `needle` (or
|
|
||||||
any of needles from the list). The check will be case insensetive if
|
|
||||||
`any_case` is True.
|
|
||||||
"""
|
|
||||||
def __init__(self, needle: str | list[str], *, any_case: bool = True):
|
|
||||||
super().__init__()
|
|
||||||
if type(needle) is str:
|
|
||||||
needle = [needle]
|
|
||||||
self._any_case = any_case
|
|
||||||
if self._any_case:
|
|
||||||
self._needle = [s.lower() for s in needle]
|
|
||||||
else:
|
|
||||||
self._needle = list(needle)
|
|
||||||
|
|
||||||
def __repr__(self):
|
|
||||||
return f"RoomIdContainsFilter({repr(self._needle)}, any_case={repr(self._any_case)})"
|
|
||||||
|
|
||||||
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
|
|
||||||
try:
|
|
||||||
room_id = room.room_id.lower() if self._any_case else room.room_id
|
|
||||||
for s in self._needle:
|
|
||||||
if s in room_id:
|
|
||||||
return True
|
|
||||||
return False
|
|
||||||
except:
|
|
||||||
return False
|
|
||||||
|
|
||||||
class RoomIdStartsWithFilter(BaseEventFilter):
|
|
||||||
"""
|
|
||||||
This filter returns True if the `room.room_id` starts with `substring`
|
|
||||||
(or any of substrings from the list). The check will be case insensetive
|
|
||||||
if `any_case` is True.
|
|
||||||
"""
|
|
||||||
def __init__(self, substring: str | list[str], *, any_case: bool = True):
|
|
||||||
super().__init__()
|
|
||||||
if type(substring) is str:
|
|
||||||
substring = [substring]
|
|
||||||
self._any_case = any_case
|
|
||||||
if self._any_case:
|
|
||||||
self._substring = [s.lower() for s in substring]
|
|
||||||
else:
|
|
||||||
self._substring = list(substring)
|
|
||||||
|
|
||||||
def __repr__(self):
|
|
||||||
return f"RoomIdStartsWithFilter({repr(self._substring)}, any_case={repr(self._any_case)})"
|
|
||||||
|
|
||||||
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
|
|
||||||
try:
|
|
||||||
room_id = room.room_id.lower() if self._any_case else room.room_id
|
|
||||||
for s in self._substring:
|
|
||||||
if room_id.startswith(s):
|
|
||||||
return True
|
|
||||||
return False
|
|
||||||
except:
|
|
||||||
return False
|
|
||||||
|
|
||||||
class RoomIdEndsWithFilter(BaseEventFilter):
|
|
||||||
"""
|
|
||||||
This filter returns True if the `room.room_id` ends with `substring` (or
|
|
||||||
any of substrings from the list). The check will be case insensetive
|
|
||||||
if `any_case` is True.
|
|
||||||
"""
|
|
||||||
def __init__(self, substring: str | list[str], *, any_case: bool = True):
|
|
||||||
super().__init__()
|
|
||||||
if type(substring) is str:
|
|
||||||
substring = [substring]
|
|
||||||
self._any_case = any_case
|
|
||||||
if self._any_case:
|
|
||||||
self._substring = [s.lower() for s in substring]
|
|
||||||
else:
|
|
||||||
self._substring = list(substring)
|
|
||||||
|
|
||||||
def __repr__(self):
|
|
||||||
return f"RoomIdEndsWithFilter({repr(self._substring)}, any_case={repr(self._any_case)})"
|
|
||||||
|
|
||||||
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
|
|
||||||
try:
|
|
||||||
room_id = room.room_id.lower() if self._any_case else room.room_id
|
|
||||||
for s in self._substring:
|
|
||||||
if room_id.endswith(s):
|
|
||||||
return True
|
|
||||||
return False
|
|
||||||
except:
|
|
||||||
return False
|
|
||||||
|
|
||||||
class RoomEncryptedFilter(BaseEventFilter):
|
class RoomEncryptedFilter(BaseEventFilter):
|
||||||
"""
|
"""
|
||||||
@@ -97,11 +9,11 @@ class RoomEncryptedFilter(BaseEventFilter):
|
|||||||
def __init__(self):
|
def __init__(self):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
|
|
||||||
def __repr__(self):
|
async def __call__(self, context: EventContext) -> bool:
|
||||||
return f"RoomEncryptedFilter()"
|
if not await super().__call__(context):
|
||||||
|
return False
|
||||||
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
|
|
||||||
try:
|
try:
|
||||||
return room.encrypted
|
context[CTX_ROOM_ENCRYPTED] = context.room.encrypted
|
||||||
|
return context.room.encrypted
|
||||||
except:
|
except:
|
||||||
return False
|
return False
|
||||||
@@ -2,16 +2,10 @@
|
|||||||
|
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
|
from enum import Enum
|
||||||
|
|
||||||
from nio import MatrixRoom, Event
|
|
||||||
from nio import UploadResponse
|
from nio import UploadResponse
|
||||||
|
|
||||||
from .filters.base import BaseEventFilter
|
|
||||||
|
|
||||||
from typing import TYPE_CHECKING
|
|
||||||
if TYPE_CHECKING:
|
|
||||||
from .bot import MatrixBot
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class MatrixBotConfig:
|
class MatrixBotConfig:
|
||||||
"""Configuration for MatrixBot"""
|
"""Configuration for MatrixBot"""
|
||||||
@@ -67,22 +61,6 @@ class VideoFileProperties:
|
|||||||
thumbnail: Path | str | bytes | None = None
|
thumbnail: Path | str | bytes | None = None
|
||||||
"""Path to the thumbnail or the raw JPEG thumbnail data"""
|
"""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
|
@dataclass
|
||||||
class UploadResult:
|
class UploadResult:
|
||||||
"""Result of data upload"""
|
"""Result of data upload"""
|
||||||
@@ -98,3 +76,30 @@ class UploadResult:
|
|||||||
|
|
||||||
filesize: int
|
filesize: int
|
||||||
"""Size of uploaded file"""
|
"""Size of uploaded file"""
|
||||||
|
|
||||||
|
class MessageType(Enum):
|
||||||
|
TEXT = "m.text"
|
||||||
|
EMOTE = "m.emote"
|
||||||
|
NOTICE = "m.notice"
|
||||||
|
IMAGE = "m.image"
|
||||||
|
FILE = "m.file"
|
||||||
|
AUDIO = "m.audio"
|
||||||
|
LOCATION = "m.location"
|
||||||
|
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)})"
|
||||||
Reference in New Issue
Block a user