7 Commits

Author SHA1 Message Date
f58c8601d1 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
2026-09-09 19:43:29 +03:00
159a43ebe6 Fixed _callbacks.py did not reraise CancelledError 2026-09-09 18:19:52 +03:00
9f4cd4948a Text filters update and stability
- **kwargs are propagated to base classes in text filters from now on
- `_callbacks.py` prints filter exceptions from now on
2026-09-09 18:18:07 +03:00
d14e110525 Updated text filters
- Improved inheritance (most text filters are derived from `TextFilter` from now on)
- Added `TextRegexFilter`
- Most text filters store string lists by reference from now on
2026-09-09 17:25:23 +03:00
3d74cb737b Added some room filters 2026-09-09 01:50:54 +03:00
7d6230881b Preparing to update filter system
- Filters are asynchronous from now on
- Filters are provided with the AsyncClient from now on
2026-09-09 01:03:51 +03:00
5fd87879ca Fixed MatrixBotConfig.allow_ainput_password
`MatrixBotConfig.allow_ainput_password` was ignored before this commit
because it was forgotten about during refactoring
2026-09-09 00:53:33 +03:00
10 changed files with 337 additions and 155 deletions

View File

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

View File

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

View File

@@ -51,7 +51,13 @@ 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
if not callback_info.filter(room, event): try:
if not await callback_info.filter(room, event, self._client):
continue
except asyncio.CancelledError:
raise
except:
self._logger.error(traceback.format_exc())
continue continue
event_data = RoomEventData( event_data = RoomEventData(
room=room, room=room,

View File

@@ -17,11 +17,14 @@ class ClientAuth:
# #
# PRIVATE # PRIVATE
# #
@staticmethod async def _default_password_callback(self) -> str:
async def _default_password_callback() -> str: if self._config is None:
raise RuntimeError("No config")
if "MATRIX_PASSWORD" in os.environ: if "MATRIX_PASSWORD" in os.environ:
return os.environ["MATRIX_PASSWORD"] return os.environ["MATRIX_PASSWORD"]
return await aioconsole.ainput("Matrix password: ") if self._config.allow_ainput_password:
return await aioconsole.ainput("Matrix password: ")
raise RuntimeError("Can't get password")
async def _login_using_session_data(self, client: AsyncClient) -> None: async def _login_using_session_data(self, client: AsyncClient) -> None:
""" """
@@ -98,11 +101,13 @@ class ClientAuth:
self._logger = logging.getLogger("ClientAuth") self._logger = logging.getLogger("ClientAuth")
self._storage = storage self._storage = storage
self._full_matrix_username: str | None = None self._full_matrix_username: str | None = None
self._config: MatrixBotConfig | None = None
async def setup(self, config: MatrixBotConfig) -> None: async def setup(self, config: MatrixBotConfig) -> None:
""" """
Setup `ClientAuth` object using `config`. Setup `ClientAuth` object using `config`.
""" """
self._config = config
self._full_matrix_username = Utils.build_full_matrix_username(config) self._full_matrix_username = Utils.build_full_matrix_username(config)
async def login(self, client: AsyncClient) -> None: async def login(self, client: AsyncClient) -> None:

View File

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

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,6 +1,8 @@
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 MatrixRoom, Event from nio import MatrixRoom, Event
class BaseEventFilter(ABC): class BaseEventFilter(ABC):
@@ -51,28 +53,38 @@ 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, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
def __call__(self, room: MatrixRoom, event: Event) -> 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:
event - the event to check againts this filter
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): class CompoundEventFilter(BaseEventFilter):
"""Event filter that consists of multiple filters""" """Event filter that consists of multiple filters"""
@@ -123,8 +135,8 @@ class CompoundEventFilter(BaseEventFilter):
expression = f"~{reprs[0]}" expression = f"~{reprs[0]}"
return f"({expression})" return f"({expression})"
def __call__(self, room: MatrixRoom, event: Event) -> bool: async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
evaluated = [arg(room, event) for arg in self._arguments] evaluated = [await arg(room, event, client) 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:

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

17
src/mab/filters/room.py Normal file
View File

@@ -0,0 +1,17 @@
from .base import BaseEventFilter
from nio import AsyncClient
from nio import MatrixRoom, Event
class RoomEncryptedFilter(BaseEventFilter):
"""
This filter returns True if the room is encrypted.
"""
def __init__(self):
super().__init__()
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
try:
return room.encrypted
except:
return False

View File

@@ -1,164 +1,191 @@
from .base import BaseEventFilter import re
import traceback
from .message import NewMessageFilter
from nio import AsyncClient
from nio import MatrixRoom, Event from nio import MatrixRoom, Event
class TextFilter(BaseEventFilter): class TextFilter(NewMessageFilter):
""" """
This filter returns True if the event contains `body` attribute. This filter returns True if all conditions are met:
`body` attribute contains unformatted text, string. 1. `event` has attribute `body`
2. `event.body` is instance of `str`
3. `event.body.strip()` evaluates to True
If this filter matches, you can access `event.body` and it stores
unformatted text of the message.
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
may disable `ignore_filename_in_body` to disable this feature.
""" """
def __init__(self): def __init__(self, *, ignore_filename_in_body: bool = True, **kwargs):
super().__init__() super().__init__(**kwargs)
self._ignore_filename_in_body = ignore_filename_in_body
def __repr__(self): async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
return "TextFilter()" if not hasattr(event, "body"):
return False
if not isinstance(event.body, str): # type: ignore
return False
if not event.body.strip(): # type: ignore
return False
if self._ignore_filename_in_body:
content = event.source["content"]
if "filename" in content and content["filename"] == event.body: # type: ignore
return False
return True
def __call__(self, room: MatrixRoom, event: Event) -> bool: class TextContainsFilter(TextFilter):
return hasattr(event, "body") and type(event.body) is str # type: ignore
class FormattedTextFilter(BaseEventFilter):
""" """
This filter returns True if the event contains valid `formatted_body` This filter returns True if `event.body` contains `needle` substring (or any
attribute. `formatted_body` attribute contains formatted text, string. of neddle from the list). `event.body` will be converted to lower case if
""" `any_case` is True.
def __init__(self):
super().__init__()
def __repr__(self): `needle` list is stored by reference so you can dynamically edit behavior of
return "FormattedTextFilter()" this filter.
def __call__(self, room: MatrixRoom, event: Event) -> bool: Please note that only `event.body` is converted to lower case if `any_case`
return hasattr(event, "formatted_body") and type(event.formatted_body) is str # type: ignore is True. That means you must ensure that `needle` is lower case. The filter
will never match otherwise.
class TextContainsFilter(BaseEventFilter):
""" """
This filter returns True if the `event.body` contains `needle` def __init__(self, needle: str | list[str], *, any_case: bool = True, **kwargs):
substring (or any of neddle from the list). The check will be case super().__init__(**kwargs)
insensetive if `any_case` is True.
"""
def __init__(self, needle: str | list[str], *, any_case: bool = True):
super().__init__()
if type(needle) is str: if type(needle) is str:
needle = [needle] needle = [needle]
self._any_case = any_case self._any_case = any_case
if self._any_case: self._needle = needle
self._needle = [s.lower() for s in needle]
else:
self._needle = list(needle)
def __repr__(self): async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
return f"TextContainsFilter({repr(self._needle)}, any_case={repr(self._any_case)})" if not await super().__call__(room, event, client):
def __call__(self, room: MatrixRoom, event: Event) -> bool:
try:
body = event.body.lower() if self._any_case else event.body # type: ignore
for n in self._needle:
if n in body:
return True
return False
except:
return False return False
body = event.body.lower() if self._any_case else event.body # type: ignore
for n in self._needle:
if n in body:
return True
return False
class TextStartsWithFilter(BaseEventFilter): class TextStartsWithFilter(TextFilter):
""" """
This filter returns True if the `event.body` starts with `substring` (or This filter returns True if `event.body` starts with `substring` (or any of
any of substrings from the list). The check will be case insensetive if substrings from the list). The check will be case insensetive if `any_case`
`any_case` is True. is True.
`substring` list is stored by reference so you can dynamically edit behavior
of this filter.
Please note that only `event.body` is converted to lower case if `any_case`
is True. That means you must ensure that `needle` is lower case. The filter
will never match otherwise.
""" """
def __init__(self, substring: str | list[str], *, any_case: bool = True): def __init__(self, substring: str | list[str], *, any_case: bool = True, **kwargs):
super().__init__() super().__init__(**kwargs)
if type(substring) is str: if type(substring) is str:
substring = [substring] substring = [substring]
self._any_case = any_case self._any_case = any_case
if self._any_case: self._substring = substring
self._substring = [s.lower() for s in substring]
else:
self._substring = list(substring)
def __repr__(self): async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
return f"TextStartsWithFilter({repr(self._substring)}, any_case={repr(self._any_case)})" if not await super().__call__(room, event, client):
def __call__(self, room: MatrixRoom, event: Event) -> bool:
try:
body = event.body.lower() if self._any_case else event.body # type: ignore
for s in self._substring:
if body.startswith(s):
return True
return False
except:
return False return False
body = event.body.lower() if self._any_case else event.body # type: ignore
for s in self._substring:
if body.startswith(s):
return True
return False
class TextEndsWithFilter(BaseEventFilter): class TextEndsWithFilter(TextFilter):
""" """
This filter returns True if the `event.body` ends with `substring` (or This filter returns True if `event.body` ends with `substring` (or any of
any of substrings from the list). The check will be case insensetive if substrings from the list). The check will be case insensetive if `any_case`
`any_case` is True. is True.
`substring` list is stored by reference so you can dynamically edit behavior
of this filter.
Please note that only `event.body` is converted to lower case if `any_case`
is True. That means you must ensure that `needle` is lower case. The filter
will never match otherwise.
""" """
def __init__(self, substring: str | list[str], *, any_case: bool = True): def __init__(self, substring: str | list[str], *, any_case: bool = True, **kwargs):
super().__init__() super().__init__(**kwargs)
if type(substring) is str: if type(substring) is str:
substring = [substring] substring = [substring]
self._any_case = any_case self._any_case = any_case
if self._any_case: self._substring = substring
self._substring = [s.lower() for s in substring]
else:
self._substring = list(substring)
def __repr__(self): async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
return f"TextEndsWithFilter({repr(self._substring)}, any_case={repr(self._any_case)})" if not await super().__call__(room, event, client):
def __call__(self, room: MatrixRoom, event: Event) -> bool:
try:
body = event.body.lower() if self._any_case else event.body # type: ignore
for s in self._substring:
if body.endswith(s):
return True
return False
except:
return False return False
body = event.body.lower() if self._any_case else event.body # type: ignore
for s in self._substring:
if body.endswith(s):
return True
return False
class TextCommandFilter(BaseEventFilter): class TextCommandFilter(TextFilter):
""" """
This filter returns True if all conditions are True: 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()
2. `event.body` conrains at most `max_args + 1` words after split() 2. `event.body` contains at most `max_args + 1` words after split()
3. First element of splitted `event.body` starts with `prefix` 3. First element of splitted `event.body` starts with `prefix`
4. First element of splitted `event.body` (after stripping `prefix`) 4. First element of splitted `event.body` (after lstripping `prefix`) starts
starts with any of strings in `verbs` list (case-insensitive) with any of strings in `verbs` list
Remarks: `verbs` list is stored by reference so you can dynamically edit behavior of
- If this filter is satified, then it will set a new attribute for the this filter.
event: `event.command_args: list[str]`. You may use this attribute in
your callback for this event. Please note that prefix is checked case sensetively. However, event.body is
converted to lower case when `verbs` matching is performed. So you must
store all verbs in lower case. This filter will not match any verbs that
use mixed case of upper case.
If this filter is matched, then it will set a new attribute for the event:
`event.command_args: list[str]`. You may use this attribute in your callback
for this event.
""" """
def __init__(self, verbs: str | list[str], min_args: int = 0, max_args: int | None = None, prefix: str = "!"): def __init__(self, verbs: str | list[str], min_args: int = 0, max_args: int | None = None, prefix: str = "!", **kwargs):
super().__init__() super().__init__(**kwargs)
if type(verbs) is str: if type(verbs) is str:
verbs = [verbs] verbs = [verbs]
verbs = [v.lower() for v in verbs]
self._verbs = verbs self._verbs = verbs
self._min_args = min_args self._min_args = min_args
self._max_args = max_args self._max_args = max_args
self._prefix = prefix self._prefix = prefix
def __repr__(self): async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
return f"TextCommandFilter({repr(self._verbs)}, {repr(self._min_args)}, {repr(self._max_args)}, {repr(self._prefix)})" if not await super().__call__(room, event, client):
def __call__(self, room: MatrixRoom, event: Event) -> bool:
try:
parts = [p.strip() for p in event.body.split() if p.strip()] # type: ignore
args_count = len(parts) - 1
if args_count < self._min_args:
return False
if self._max_args is not None and args_count > self._max_args:
return False
if not parts[0].startswith(self._prefix):
return False
cmd = parts[0][len(self._prefix):].lower()
for verb in self._verbs:
if cmd == verb:
setattr(event, "command_args", parts[1:])
return True
return False return False
parts = [p.strip() for p in event.body.split() if p.strip()] # type: ignore
args_count = len(parts) - 1
if args_count < self._min_args:
return False
if self._max_args is not None and args_count > self._max_args:
return False
if not parts[0].startswith(self._prefix):
return False
cmd = parts[0][len(self._prefix):].lower()
for verb in self._verbs:
if cmd == verb:
setattr(event, "command_args", parts[1:])
return True
return False
class TextRegexFilter(TextFilter):
"""
This filter returns True if the `event.body` passes the regex.
"""
def __init__(self, regex: re.Pattern | str, **kwargs):
super().__init__(**kwargs)
if isinstance(regex, str):
regex = re.compile(regex)
self._regex = regex
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
if not await super().__call__(room, event, client):
return False
try:
return self._regex.match(event.body) # type: ignore
except: except:
self._logger.error(traceback.format_exc())
return False return False