From 7d6230881badf537b01f45b37904c1cc15aaa7cf Mon Sep 17 00:00:00 2001 From: nikita Date: Wed, 9 Sep 2026 01:03:51 +0300 Subject: [PATCH] Preparing to update filter system - Filters are asynchronous from now on - Filters are provided with the AsyncClient from now on --- src/mab/bot/_callbacks.py | 2 +- src/mab/filters/base.py | 9 ++++++--- src/mab/filters/text.py | 13 +++++++------ 3 files changed, 14 insertions(+), 10 deletions(-) diff --git a/src/mab/bot/_callbacks.py b/src/mab/bot/_callbacks.py index 6b6ed90..b190e39 100644 --- a/src/mab/bot/_callbacks.py +++ b/src/mab/bot/_callbacks.py @@ -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, diff --git a/src/mab/filters/base.py b/src/mab/filters/base.py index ed938e8..3c5a2b8 100644 --- a/src/mab/filters/base.py +++ b/src/mab/filters/base.py @@ -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: diff --git a/src/mab/filters/text.py b/src/mab/filters/text.py index a25bad0..33fa6ec 100644 --- a/src/mab/filters/text.py +++ b/src/mab/filters/text.py @@ -1,5 +1,6 @@ from .base import BaseEventFilter +from nio import AsyncClient from nio import MatrixRoom, Event class TextFilter(BaseEventFilter): @@ -13,7 +14,7 @@ class TextFilter(BaseEventFilter): def __repr__(self): return "TextFilter()" - def __call__(self, room: MatrixRoom, event: Event) -> bool: + async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool: return hasattr(event, "body") and type(event.body) is str # type: ignore class FormattedTextFilter(BaseEventFilter): @@ -27,7 +28,7 @@ class FormattedTextFilter(BaseEventFilter): def __repr__(self): return "FormattedTextFilter()" - def __call__(self, room: MatrixRoom, event: Event) -> bool: + async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool: return hasattr(event, "formatted_body") and type(event.formatted_body) is str # type: ignore class TextContainsFilter(BaseEventFilter): @@ -49,7 +50,7 @@ class TextContainsFilter(BaseEventFilter): def __repr__(self): return f"TextContainsFilter({repr(self._needle)}, any_case={repr(self._any_case)})" - def __call__(self, room: MatrixRoom, event: Event) -> bool: + async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool: try: body = event.body.lower() if self._any_case else event.body # type: ignore for n in self._needle: @@ -78,7 +79,7 @@ class TextStartsWithFilter(BaseEventFilter): def __repr__(self): return f"TextStartsWithFilter({repr(self._substring)}, any_case={repr(self._any_case)})" - def __call__(self, room: MatrixRoom, event: Event) -> bool: + async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool: try: body = event.body.lower() if self._any_case else event.body # type: ignore for s in self._substring: @@ -107,7 +108,7 @@ class TextEndsWithFilter(BaseEventFilter): def __repr__(self): return f"TextEndsWithFilter({repr(self._substring)}, any_case={repr(self._any_case)})" - def __call__(self, room: MatrixRoom, event: Event) -> bool: + async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool: try: body = event.body.lower() if self._any_case else event.body # type: ignore for s in self._substring: @@ -144,7 +145,7 @@ 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: + async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool: try: parts = [p.strip() for p in event.body.split() if p.strip()] # type: ignore args_count = len(parts) - 1