4 Commits

Author SHA1 Message Date
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
5 changed files with 268 additions and 93 deletions

View File

@@ -51,7 +51,7 @@ class Callbacks:
for callback_info in self._filters:
if not isinstance(callback_info, self._FilterBasedCallback):
continue
if not callback_info.filter(room, event):
if not await callback_info.filter(room, event, self._client):
continue
event_data = RoomEventData(
room=room,

View File

@@ -17,11 +17,14 @@ class ClientAuth:
#
# PRIVATE
#
@staticmethod
async def _default_password_callback() -> str:
async def _default_password_callback(self) -> str:
if self._config is None:
raise RuntimeError("No config")
if "MATRIX_PASSWORD" in os.environ:
return os.environ["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:
"""
@@ -98,11 +101,13 @@ class ClientAuth:
self._logger = logging.getLogger("ClientAuth")
self._storage = storage
self._full_matrix_username: str | None = None
self._config: MatrixBotConfig | None = None
async def setup(self, config: MatrixBotConfig) -> None:
"""
Setup `ClientAuth` object using `config`.
"""
self._config = config
self._full_matrix_username = Utils.build_full_matrix_username(config)
async def login(self, client: AsyncClient) -> None:

View File

@@ -1,6 +1,7 @@
from abc import ABC, abstractmethod
import logging
from nio import AsyncClient
from nio import MatrixRoom, Event
class BaseEventFilter(ABC):
@@ -59,14 +60,16 @@ class BaseEventFilter(ABC):
pass
@abstractmethod
def __call__(self, room: MatrixRoom, event: Event) -> bool:
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
@@ -123,8 +126,8 @@ class CompoundEventFilter(BaseEventFilter):
expression = f"~{reprs[0]}"
return f"({expression})"
def __call__(self, room: MatrixRoom, event: Event) -> bool:
evaluated = [arg(room, event) for arg in self._arguments]
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
evaluated = [await arg(room, event, client) for arg in self._arguments]
if self._operator == self.OPERATOR_AND:
return all(evaluated)
elif self._operator == self.OPERATOR_OR:

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

@@ -0,0 +1,107 @@
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.
"""
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
except:
return False

View File

@@ -1,25 +1,44 @@
import re
import traceback
from .base import BaseEventFilter
from nio import AsyncClient
from nio import MatrixRoom, Event
class TextFilter(BaseEventFilter):
"""
This filter returns True if the event contains `body` attribute.
`body` attribute contains unformatted text, string.
This filter returns True if all conditions are met:
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.
"""
def __init__(self):
super().__init__()
def __repr__(self):
def __repr__(self) -> str:
return "TextFilter()"
def __call__(self, room: MatrixRoom, event: Event) -> bool:
return hasattr(event, "body") and type(event.body) is str # type: ignore
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
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
return True
class FormattedTextFilter(BaseEventFilter):
"""
This filter returns True if the event contains valid `formatted_body`
attribute. `formatted_body` attribute contains formatted text, string.
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):
super().__init__()
@@ -27,115 +46,136 @@ class FormattedTextFilter(BaseEventFilter):
def __repr__(self):
return "FormattedTextFilter()"
def __call__(self, room: MatrixRoom, event: Event) -> bool:
return hasattr(event, "formatted_body") and type(event.formatted_body) is str # type: ignore
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(BaseEventFilter):
class TextContainsFilter(TextFilter):
"""
This filter returns True if the `event.body` contains `needle`
substring (or any of neddle from the list). The check will be case
insensetive if `any_case` is True.
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
`any_case` is True.
`needle` 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, 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)
self._needle = needle
def __repr__(self):
def __repr__(self) -> str:
return f"TextContainsFilter({repr(self._needle)}, any_case={repr(self._any_case)})"
def __call__(self, room: MatrixRoom, event: Event) -> bool:
try:
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
if not await super().__call__(room, event, client):
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
except:
return False
class TextStartsWithFilter(BaseEventFilter):
class TextStartsWithFilter(TextFilter):
"""
This filter returns True if the `event.body` starts with `substring` (or
any of substrings from the list). The check will be case insensetive if
`any_case` is True.
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`
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):
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)
self._substring = substring
def __repr__(self):
return f"TextStartsWithFilter({repr(self._substring)}, any_case={repr(self._any_case)})"
def __call__(self, room: MatrixRoom, event: Event) -> bool:
try:
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
if not await super().__call__(room, event, client):
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
except:
return False
class TextEndsWithFilter(BaseEventFilter):
class TextEndsWithFilter(TextFilter):
"""
This filter returns True if the `event.body` ends with `substring` (or
any of substrings from the list). The check will be case insensetive if
`any_case` is True.
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`
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):
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)
self._substring = substring
def __repr__(self):
return f"TextEndsWithFilter({repr(self._substring)}, any_case={repr(self._any_case)})"
def __call__(self, room: MatrixRoom, event: Event) -> bool:
try:
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
if not await super().__call__(room, event, client):
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
except:
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()
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`
4. First element of splitted `event.body` (after stripping `prefix`)
starts with any of strings in `verbs` list (case-insensitive)
4. First element of splitted `event.body` (after lstripping `prefix`) starts
with any of strings in `verbs` list
Remarks:
- If this filter is satified, 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.
`verbs` list is stored by reference so you can dynamically edit behavior of
this filter.
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 = "!"):
super().__init__()
if type(verbs) is str:
verbs = [verbs]
verbs = [v.lower() for v in verbs]
self._verbs = verbs
self._min_args = min_args
self._max_args = max_args
@@ -144,8 +184,9 @@ class TextCommandFilter(BaseEventFilter):
def __repr__(self):
return f"TextCommandFilter({repr(self._verbs)}, {repr(self._min_args)}, {repr(self._max_args)}, {repr(self._prefix)})"
def __call__(self, room: MatrixRoom, event: Event) -> bool:
try:
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
if not await super().__call__(room, event, client):
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:
@@ -160,5 +201,24 @@ class TextCommandFilter(BaseEventFilter):
setattr(event, "command_args", parts[1:])
return True
return False
except:
class TextRegexFilter(TextFilter):
"""
This filter returns True if the `event.body` passes the regex.
"""
def __init__(self, regex: re.Pattern | str):
if isinstance(regex, str):
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
try:
return self._regex.match(event.body) # type: ignore
except:
self._logger.error(traceback.format_exc())
return False