Filter update, fixed bot stop, v0.4.0

- Fixed `asyncio.shield` not being awaited in `ClientManager.stop()` (that led to `next_batch` value not being saved and session not being closed properly)
- `BaseEventFilter` does not have any abstract methods anymore and can be used to use callback for any event
- `BaseEventFilter.__repr__` prints class name now
- `BaseEventFilter.__call__` returns True now
- Added `EventTypeFilter` which can be used to check if the event is an instance of some class
- Removed most room events
- Added `NewMessageFilter`, `EditedMessageFilter`, `RedactedMessageFilter`, `SenderIsFilter` and `SenderIsBotFilter`
- Removed `FormattedTextFilter`
- Text filters are derived from `NewMessageFilter` so they won't match edited messages anymore
This commit is contained in:
2026-09-09 19:43:29 +03:00
parent 159a43ebe6
commit f58c8601d1
8 changed files with 157 additions and 166 deletions

View File

@@ -18,7 +18,7 @@ The package supports the following features:
Use `pip` to install this package:
```bash
python -m pip install git+https://git.tyukalov.su/nikita/mab@v0.3.0
python -m pip install git+https://git.tyukalov.su/nikita/mab@v0.4.0
```
You should specify package version you want to use, because `main` without tags
@@ -27,18 +27,16 @@ contain unstable code.
## Basic usage
This is the most simple bot you can create. It would respond to any message
that starts with `!test`, `!hello` or `!hi`.
that contains `hello` and `hi` words.
```python
import asyncio
from pathlib import Path
from mab import MatrixBot, MatrixBotConfig
from mab import TextCommandFilter
from mab import TextContainsFilter, SenderIsBotFilter
from mab.types import RoomEventData
async def on_valid_command(data: RoomEventData) -> None:
# do not respond to ourselves
if event.sender == data.bot.get_client().user_id:
return
async def on_message(data: RoomEventData) -> None:
text = f"Your message contains {len(data.event.body)} symbols"
await data.bot.send_text_to_room(data.room, text)
@@ -51,8 +49,8 @@ async def main() -> None:
)
bot = MatrixBot(matrix_bot_config)
bot.add_callback(
TextCommandFilter(["test", "hello", "hi"]),
on_valid_command
~SenderIsBotFilter() & TextContainsFilter(["test", "hello", "hi"]),
on_message
)
await bot.start()
# wait for Ctrl+C

View File

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

View File

@@ -191,11 +191,13 @@ class ClientManager:
raise RuntimeError("The bot was never started")
self._background_task.cancel()
try:
asyncio.shield(self._background_task)
await asyncio.shield(self._background_task)
except asyncio.CancelledError:
pass
except:
self._logger.error(traceback.format_exc())
try:
asyncio.shield(self._close_client())
await asyncio.shield(self._close_client())
except:
self._logger.error(traceback.format_exc())
self._background_task = None

View File

@@ -0,0 +1,28 @@
from .base import *
from .message import *
from .room import *
from .text import *
__all__ = [
# base.py
"BaseEventFilter",
"EventTypeFilter",
# message.py
"NewMessageFilter",
"EditedMessageFilter",
"RedactedMessageFilter",
"SenderIsFilter",
"SenderIsBotFilter",
# room.py
"RoomEncryptedFilter",
# text.py
"TextFilter",
"TextContainsFilter",
"TextStartsWithFilter",
"TextEndsWithFilter",
"TextCommandFilter",
"TextRegexFilter"
]

View File

@@ -1,5 +1,6 @@
from abc import ABC, abstractmethod
import logging
from typing import Any, Type
from nio import AsyncClient
from nio import MatrixRoom, Event
@@ -52,30 +53,38 @@ class BaseEventFilter(ABC):
)
# PAYLOAD
@abstractmethod
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, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
"""This abstract method must be redefined in derived classes so that
the filter operates according to its description. This method must
not raise exceptions. In case of exception it should log it using
`self._logger` and return False
Args:
room - room the event has happened in
event - the event to check againts this filter
client - the client
Returns:
True if the event satisfies this filter
False if the event does not satisfy this filter
"""
pass
This abstract method must be redefined in derived classes so that the
filter operates according to its description. This method must not raise
exceptions. In case of exception it should log it using `self._logger`
and return False
Args:
- room - room the event has happened in
- event - the event to check againts this filter
- client - the client
Returns:
- True if the event satisfies this filter
- False if the event does not satisfy this filter
"""
return True
class EventTypeFilter(BaseEventFilter):
"""Event filter that checks if the event is an instance of some class"""
def __init__(self, type: Type):
self._type = type
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
return isinstance(event, self._type)
class CompoundEventFilter(BaseEventFilter):
"""Event filter that consists of multiple filters"""

View File

@@ -0,0 +1,87 @@
import traceback
from .base import BaseEventFilter, EventTypeFilter
from nio import AsyncClient
from nio import MatrixRoom, Event
from nio import RedactionEvent
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, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
if not await super().__call__(room, event, client):
return False
return "m.new_content" not in 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, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
if not await super().__call__(room, event, client):
return False
return "m.new_content" in event.source["content"]
class RedactedMessageFilter(EventTypeFilter):
"""
This filter returns True if the event is a RedactionEvent.
"""
def __init__(self, **kwargs):
super().__init__(RedactionEvent, **kwargs)
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
return await super().__call__(room, event, client)
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.
"""
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, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
if not await super().__call__(room, event, client):
return False
sender = event.sender.lower() if self._any_case else event.sender
for s in self._sender:
if sender == s:
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, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
if not await super().__call__(room, event, client):
return False
return client.user_id == event.sender

View File

@@ -3,93 +3,6 @@ from .base import BaseEventFilter
from nio import AsyncClient
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):
"""
This filter returns True if the room is encrypted.
@@ -97,9 +10,6 @@ class RoomEncryptedFilter(BaseEventFilter):
def __init__(self):
super().__init__()
def __repr__(self):
return f"RoomEncryptedFilter()"
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
try:
return room.encrypted

View File

@@ -1,11 +1,11 @@
import re
import traceback
from .base import BaseEventFilter
from .message import NewMessageFilter
from nio import AsyncClient
from nio import MatrixRoom, Event
class TextFilter(BaseEventFilter):
class TextFilter(NewMessageFilter):
"""
This filter returns True if all conditions are met:
1. `event` has attribute `body`
@@ -23,9 +23,6 @@ class TextFilter(BaseEventFilter):
super().__init__(**kwargs)
self._ignore_filename_in_body = ignore_filename_in_body
def __repr__(self) -> str:
return "TextFilter()"
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
if not hasattr(event, "body"):
return False
@@ -39,31 +36,6 @@ class TextFilter(BaseEventFilter):
return False
return True
class FormattedTextFilter(BaseEventFilter):
"""
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
@@ -84,9 +56,6 @@ class TextContainsFilter(TextFilter):
self._any_case = any_case
self._needle = needle
def __repr__(self) -> str:
return f"TextContainsFilter({repr(self._needle)}, any_case={repr(self._any_case)})"
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
if not await super().__call__(room, event, client):
return False
@@ -116,9 +85,6 @@ class TextStartsWithFilter(TextFilter):
self._any_case = any_case
self._substring = substring
def __repr__(self):
return f"TextStartsWithFilter({repr(self._substring)}, any_case={repr(self._any_case)})"
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
if not await super().__call__(room, event, client):
return False
@@ -148,9 +114,6 @@ class TextEndsWithFilter(TextFilter):
self._any_case = any_case
self._substring = substring
def __repr__(self):
return f"TextEndsWithFilter({repr(self._substring)}, any_case={repr(self._any_case)})"
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
if not await super().__call__(room, event, client):
return False
@@ -190,9 +153,6 @@ class TextCommandFilter(TextFilter):
self._max_args = max_args
self._prefix = prefix
def __repr__(self):
return f"TextCommandFilter({repr(self._verbs)}, {repr(self._min_args)}, {repr(self._max_args)}, {repr(self._prefix)})"
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
if not await super().__call__(room, event, client):
return False
@@ -221,9 +181,6 @@ class TextRegexFilter(TextFilter):
regex = re.compile(regex)
self._regex = regex
def __repr__(self):
return f"TextRegexFilter({repr(self._regex)})"
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
if not await super().__call__(room, event, client):
return False