From d8bf1818d7393b594955c290c65289d2a7055425 Mon Sep 17 00:00:00 2001 From: nikita Date: Wed, 2 Sep 2026 18:20:48 +0300 Subject: [PATCH] Updated to v0.1.0, added event filters --- pyproject.toml | 2 +- src/mab/__init__.py | 16 ++++- src/mab/bot.py | 35 ++++++++-- src/mab/filters/base.py | 138 +++++++++++++++++++++++++++++++++++++ src/mab/filters/text.py | 147 ++++++++++++++++++++++++++++++++++++++++ 5 files changed, 330 insertions(+), 8 deletions(-) create mode 100644 src/mab/filters/base.py create mode 100644 src/mab/filters/text.py diff --git a/pyproject.toml b/pyproject.toml index 08f9092..93bb3c0 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta" [project] name = "mab" -version = "0.0.2" +version = "0.1.0" authors = [ { name = "Tyukalov Nikita", email = "nikita@tyukalov.su" } ] diff --git a/src/mab/__init__.py b/src/mab/__init__.py index 369c463..f0fc5f8 100644 --- a/src/mab/__init__.py +++ b/src/mab/__init__.py @@ -5,6 +5,9 @@ from .types import MatrixBotConfig from .bot import MatrixBot +from .filters.base import * +from .filters.text import * + __all__ = [ # module names "bot", @@ -14,5 +17,16 @@ __all__ = [ "MatrixBotConfig", # .bot - "MatrixBot" + "MatrixBot", + + # .filters.base + "BaseEventFilter", + + # .filters.text + "TextFilter", + "FormattedTextFilter", + "TextContainsFilter", + "TextStartsWithFilter", + "TextEndsWithFilter", + "TextCommandFilter", ] \ No newline at end of file diff --git a/src/mab/bot.py b/src/mab/bot.py index ea104a1..19dbbdf 100644 --- a/src/mab/bot.py +++ b/src/mab/bot.py @@ -21,7 +21,9 @@ from nio import OlmUnverifiedDeviceError from nio import MatrixInvitedRoom, InviteMemberEvent from nio import JoinResponse -import nio.events +from .filters.base import BaseEventFilter + +from nio.events.room_events import Event as RoomEvemt from .types import * @@ -244,13 +246,34 @@ class MatrixBot: except: self._logger.error(traceback.format_exc()) + async def _callback_filter_router(self, *args, **kwargs): + if len(args) != 2: + self._logger.debug("Can't process the event, not enough positional args") + self._debug_event_callback(*args, **kwargs) + return + room = args[0] + event = args[1] + for filter in self._filters: + filter_object = filter[0] + filter_callback = filter[1] + filter_stop_after_this = filter[2] + if filter_object(room, event): + self._logger.debug(f"Filter {repr(filter_object)} matched") + try: + await filter_callback(room, event) + except: + self._logger.error(traceback.format_exc()) + if filter_stop_after_this: + self._logger.debug(f"Filter {repr(filter_object)} stops matching") + break + # # LIFECYCLE # def _setup_client_callbacks(self) -> None: """Setup internal client callbacks""" self._client.add_response_callback(self._callback_sync, SyncResponse) # type: ignore - + self._client.add_event_callback(self._callback_filter_router, RoomEvemt) # type: ignore if self._config.auto_join_any_room_on_invite: self._client.add_event_callback(self._callback_autojoin, InviteMemberEvent) # type: ignore @@ -404,6 +427,7 @@ class MatrixBot: self._last_next_batch_dump: float = 0.0 self._last_next_batch: str | None = None self._cb_password = self._default_password_callback + self._filters = [] def start(self) -> None: """Start the bot. @@ -451,12 +475,11 @@ class MatrixBot: result = True return result - def add_event_callback(self, callback: Callable[[Any], Awaitable[None]] | None, event_class: nio.events.Event) -> None: - """Added event callback for events of specified class. - Use `None` instead of callback to print parameter types you need to use in your callback.""" + def add_event_callback(self, callback: Callable[[Any], Awaitable[None]] | None, filter: BaseEventFilter, stop_after_this: bool = True) -> None: + """Add event callback for events that pass the filter.""" if callback is None: callback = self._debug_event_callback - self._client.add_event_callback(callback, event_class) # type: ignore + self._filters.append((filter, callback, stop_after_this)) def get_client(self) -> AsyncClient: """Get AsyncClient in use""" diff --git a/src/mab/filters/base.py b/src/mab/filters/base.py new file mode 100644 index 0000000..ecf07ef --- /dev/null +++ b/src/mab/filters/base.py @@ -0,0 +1,138 @@ +from abc import ABC, abstractmethod +import logging + +from nio import MatrixRoom, Event + +class BaseEventFilter(ABC): + """Base class for all message filters""" + _logger = logging.Logger("EventFilter") + + # AND + def __and__(self, other): + if not isinstance(other, BaseEventFilter): + raise TypeError(f"BaseEventFilter can't be ANDed against {type(other)}") + return CompoundEventFilter( + CompoundEventFilter.OPERATOR_AND, + [self, other] + ) + + def __rand__(self, other): + return self.__and__(other) + + # OR + def __or__(self, other): + if not isinstance(other, BaseEventFilter): + raise TypeError(f"BaseEventFilter can't be ORed against {type(other)}") + return CompoundEventFilter( + CompoundEventFilter.OPERATOR_OR, + [self, other] + ) + + def __ror__(self, other): + return self.__ror__(other) + + # XOR + def __xor__(self, other): + if not isinstance(other, BaseEventFilter): + raise TypeError(f"BaseEventFilter can't be XORed against {type(other)}") + return CompoundEventFilter( + CompoundEventFilter.OPERATOR_XOR, + [self, other] + ) + + def __rxor__(self, other): + return self.__xor__(other) + + # INVERT + def __invert__(self): + return CompoundEventFilter( + CompoundEventFilter.OPERATOR_INVERT, + [self] + ) + + # PAYLOAD + @abstractmethod + def __repr__(self) -> str: + """This method must be redefined in derived classes to improve + debugging experience. + """ + pass + + @abstractmethod + 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 + +class CompoundEventFilter(BaseEventFilter): + """Event filter that consists of multiple filters""" + OPERATOR_AND = "and" + OPERATOR_OR = "or" + OPERATOR_XOR = "xor" + OPERATOR_INVERT = "invert" + + @staticmethod + def _is_operator_valid(op: str) -> bool: + """Returns True if operator `op` is a valid operator""" + return op in [ + CompoundEventFilter.OPERATOR_AND, + CompoundEventFilter.OPERATOR_OR, + CompoundEventFilter.OPERATOR_XOR, + CompoundEventFilter.OPERATOR_INVERT + ] + + @staticmethod + def _is_elements_count_valid(op: str, count: int) -> bool: + """Returns True if operator `op` may take `count` arguments""" + return count in { + CompoundEventFilter.OPERATOR_AND: [2], + CompoundEventFilter.OPERATOR_OR: [2], + CompoundEventFilter.OPERATOR_XOR: [2], + CompoundEventFilter.OPERATOR_INVERT: [1], + }[op] + + def __init__(self, operator: str, arguments: list[BaseEventFilter]): + super().__init__() + if not self._is_operator_valid(operator): + raise RuntimeError(f"Invalid operator `{operator}`") + if not self._is_elements_count_valid(operator, len(arguments)): + raise RuntimeError(f"Operator `{operator}` does not take `{len(arguments)}` arguments") + self._operator = operator + self._arguments = list(arguments) + + def __repr__(self) -> str: + expression = "False" + if self._operator == CompoundEventFilter.OPERATOR_AND: + expression = " & ".join(self._arguments) + elif self._operator == CompoundEventFilter.OPERATOR_OR: + expression = " | ".join(self._arguments) + elif self._operator == CompoundEventFilter.OPERATOR_XOR: + expression = " ^ ".join(self._arguments) + elif self._operator == CompoundEventFilter.OPERATOR_INVERT: + expression = f"~{self._arguments[0]}" + return f"({expression})" + + def __call__(self, room: MatrixRoom, event: Event) -> bool: + evaluated = [arg(room, event) for arg in self._arguments] + if self._operator == self.OPERATOR_AND: + return all(evaluated) + elif self._operator == self.OPERATOR_OR: + return any(evaluated) + elif self._operator == self.OPERATOR_XOR: + result = evaluated[0] + for v in evaluated[1:]: + result ^= v + return result + elif self._operator == self.OPERATOR_INVERT: + return not evaluated[0] + return False \ No newline at end of file diff --git a/src/mab/filters/text.py b/src/mab/filters/text.py new file mode 100644 index 0000000..363fc7e --- /dev/null +++ b/src/mab/filters/text.py @@ -0,0 +1,147 @@ +from .base import BaseEventFilter + +from nio import MatrixRoom, Event + +class TextFilter(BaseEventFilter): + """ + This filter returns True if the event contains `body` attribute. + `body` attribute contains unformatted text, string. + """ + def __init__(self): + super().__init__() + + def __repr__(self): + return "TextFilter" + + def __call__(self, room: MatrixRoom, event: Event) -> bool: + return hasattr(event, "body") and type(event.body) is str + +class FormattedTextFilter(BaseEventFilter): + """ + This filter returns True if the event contains valid `formatted_body` + attribute. `formatted_body` attribute contains formatted text, string. + """ + def __init__(self): + super().__init__() + + 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 + +class TextContainsFilter(BaseEventFilter): + """ + This filter returns True if the `event.body` contains `needle` + substring (or any of neddle from the list). + """ + def __init__(self, needle: str | list[str]): + super().__init__() + if type(needle) is str: + needle = [needle] + self._needle = needle + + def __repr__(self): + return f"TextContainsFilter({repr(self._neddle)})" + + def __call__(self, room: MatrixRoom, event: Event) -> bool: + try: + for n in self._needle: + if n in event.body: + return True + return False + except: + return False + +class TextStartsWithFilter(BaseEventFilter): + """ + This filter returns True if the `event.body` starts with `substring` (or + any of substrings from the list). + """ + def __init__(self, substring: str | list[str]): + super().__init__() + if type(substring) is str: + substring = [substring] + self._substring = substring + + def __repr__(self): + return f"TextStartsWithFilter({repr(self._substring)})" + + def __call__(self, room: MatrixRoom, event: Event) -> bool: + try: + for s in self._substring: + if event.body.startswith(s): + return True + print(self._substring) + return False + except: + return False + +class TextEndsWithFilter(BaseEventFilter): + """ + This filter returns True if the `event.body` ends with `substring` (or + any of substrings from the list). + """ + def __init__(self, substring: str | list[str]): + super().__init__() + if type(substring) is str: + substring = [substring] + self._substring = substring + + def __repr__(self): + return f"TextEndsWithFilter({repr(self._substring)})" + + def __call__(self, room: MatrixRoom, event: Event) -> bool: + try: + for s in self._substring: + if event.body.endswith(s): + return True + return False + except: + return False + +class TextCommandFilter(BaseEventFilter): + """ + This filter returns True if all conditions are True: + 1. `event.body` contains at least `min_args + 1` words after split() + 2. `event.body` conrains 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) + + 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. + """ + 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 + self._prefix = prefix + + 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: + parts = [p.strip() for p in event.body.split() if p.strip()] + 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 + except: + return False \ No newline at end of file