Filters refactoring and small improvement

- Removed `filters/__init__.py` (filters are to be imported from `mab` directly)
- Renamed `Text` filters to `Body` filters
- Fixed `base.py` filters not propagating **kwargs to base classes
- Added `__init__` for BaseEventFilter so that keyword argument that remain unused by its children get logged using `logging.critical`
This commit is contained in:
2026-09-12 20:14:28 +03:00
parent 4f0792b9aa
commit 1fe4434e16
4 changed files with 40 additions and 46 deletions

View File

@@ -1,12 +1,13 @@
from . import bot from . import bot
from . import types from . import types
from .types import MatrixBotConfig from .types import MatrixBotConfig, RoomEventData, MessageType
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 +16,29 @@ __all__ = [
# .types # .types
"MatrixBotConfig", "MatrixBotConfig",
"RoomEventData",
"MessageType",
# .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",
] ]

View File

@@ -1,28 +0,0 @@
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

@@ -9,6 +9,10 @@ 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):
@@ -80,7 +84,8 @@ class BaseEventFilter(ABC):
class EventTypeFilter(BaseEventFilter): class EventTypeFilter(BaseEventFilter):
"""Event filter that checks if the event is an instance of some class""" """Event filter that checks if the event is an instance of some class"""
def __init__(self, type: Type): def __init__(self, type: Type, **kwargs):
super().__init__(**kwargs)
self._type = type self._type = type
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool: async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
@@ -113,8 +118,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)):

View File

@@ -5,7 +5,7 @@ from .message import NewMessageFilter
from nio import AsyncClient from nio import AsyncClient
from nio import MatrixRoom, Event from nio import MatrixRoom, Event
class TextFilter(NewMessageFilter): 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,12 +18,17 @@ class TextFilter(NewMessageFilter):
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.
""" """
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
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool: async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
if not await super().__call__(room, event, client):
return False
if not hasattr(event, "body"): if not hasattr(event, "body"):
return False return False
if not isinstance(event.body, str): # type: ignore if not isinstance(event.body, str): # type: ignore
@@ -36,7 +41,7 @@ class TextFilter(NewMessageFilter):
return False return False
return True return True
class TextContainsFilter(TextFilter): class BodyContainsFilter(BodyExistsFilter):
""" """
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
@@ -65,7 +70,7 @@ class TextContainsFilter(TextFilter):
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`
@@ -94,7 +99,7 @@ class TextStartsWithFilter(TextFilter):
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`
@@ -123,7 +128,7 @@ class TextEndsWithFilter(TextFilter):
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()
@@ -171,7 +176,7 @@ class TextCommandFilter(TextFilter):
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.
""" """