From bc3a7500e2f149baaf73a996778c6724295377a1 Mon Sep 17 00:00:00 2001 From: "Nikita Tyukalov, ASUS, Linux" Date: Mon, 7 Sep 2026 01:02:38 +0300 Subject: [PATCH] Filters fixes and updates - Removed leftover `print` in `TextStartsWithFilter.__call__` - Fixed broken `CompoundEventFilter.__repr__` - Fixed typo in `TextContainsFilter.__repr__` - Fixed text returned by `TextFilter.__repr__`, `FormattedTextFilter.__repr__` - Fixed type error in `TextFilter.__call__`, `FormattedTextFilter.__call__` - Added `any_case` setting for `TextContainsFilter`, `TextStartsWithFilter`, `TextEndsWithFilter` --- src/mab/filters/base.py | 9 ++++--- src/mab/filters/text.py | 59 ++++++++++++++++++++++++++--------------- 2 files changed, 43 insertions(+), 25 deletions(-) diff --git a/src/mab/filters/base.py b/src/mab/filters/base.py index ecf07ef..ed938e8 100644 --- a/src/mab/filters/base.py +++ b/src/mab/filters/base.py @@ -112,14 +112,15 @@ class CompoundEventFilter(BaseEventFilter): def __repr__(self) -> str: expression = "False" + reprs = [repr(a) for a in self._arguments] if self._operator == CompoundEventFilter.OPERATOR_AND: - expression = " & ".join(self._arguments) + expression = " & ".join(reprs) elif self._operator == CompoundEventFilter.OPERATOR_OR: - expression = " | ".join(self._arguments) + expression = " | ".join(reprs) elif self._operator == CompoundEventFilter.OPERATOR_XOR: - expression = " ^ ".join(self._arguments) + expression = " ^ ".join(reprs) elif self._operator == CompoundEventFilter.OPERATOR_INVERT: - expression = f"~{self._arguments[0]}" + expression = f"~{reprs[0]}" return f"({expression})" def __call__(self, room: MatrixRoom, event: Event) -> bool: diff --git a/src/mab/filters/text.py b/src/mab/filters/text.py index 363fc7e..a25bad0 100644 --- a/src/mab/filters/text.py +++ b/src/mab/filters/text.py @@ -11,10 +11,10 @@ class TextFilter(BaseEventFilter): super().__init__() def __repr__(self): - return "TextFilter" + return "TextFilter()" def __call__(self, room: MatrixRoom, event: Event) -> bool: - return hasattr(event, "body") and type(event.body) is str + return hasattr(event, "body") and type(event.body) is str # type: ignore class FormattedTextFilter(BaseEventFilter): """ @@ -25,29 +25,35 @@ class FormattedTextFilter(BaseEventFilter): super().__init__() def __repr__(self): - return "FormattedTextFilter" + return "FormattedTextFilter()" def __call__(self, room: MatrixRoom, event: Event) -> bool: - return hasattr(event, "formatted_body") and type(event.formatted_body) is str + return hasattr(event, "formatted_body") and type(event.formatted_body) is str # type: ignore class TextContainsFilter(BaseEventFilter): """ This filter returns True if the `event.body` contains `needle` - substring (or any of neddle from the list). + substring (or any of neddle from the list). The check will be case + insensetive if `any_case` is True. """ - def __init__(self, needle: str | list[str]): + def __init__(self, needle: str | list[str], *, any_case: bool = True): super().__init__() if type(needle) is str: needle = [needle] - self._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"TextContainsFilter({repr(self._neddle)})" + return f"TextContainsFilter({repr(self._needle)}, any_case={repr(self._any_case)})" 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 event.body: + if n in body: return True return False except: @@ -56,23 +62,28 @@ class TextContainsFilter(BaseEventFilter): class TextStartsWithFilter(BaseEventFilter): """ This filter returns True if the `event.body` starts with `substring` (or - any of substrings from the list). + any of substrings from the list). The check will be case insensetive if + `any_case` is True. """ - def __init__(self, substring: str | list[str]): + def __init__(self, substring: str | list[str], *, any_case: bool = True): super().__init__() if type(substring) is str: substring = [substring] - self._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"TextStartsWithFilter({repr(self._substring)})" + return f"TextStartsWithFilter({repr(self._substring)}, any_case={repr(self._any_case)})" 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 event.body.startswith(s): + if body.startswith(s): return True - print(self._substring) return False except: return False @@ -80,21 +91,27 @@ class TextStartsWithFilter(BaseEventFilter): class TextEndsWithFilter(BaseEventFilter): """ This filter returns True if the `event.body` ends with `substring` (or - any of substrings from the list). + any of substrings from the list). The check will be case insensetive if + `any_case` is True. """ - def __init__(self, substring: str | list[str]): + def __init__(self, substring: str | list[str], *, any_case: bool = True): super().__init__() if type(substring) is str: substring = [substring] - self._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"TextEndsWithFilter({repr(self._substring)})" + return f"TextEndsWithFilter({repr(self._substring)}, any_case={repr(self._any_case)})" 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 event.body.endswith(s): + if body.endswith(s): return True return False except: @@ -129,7 +146,7 @@ class TextCommandFilter(BaseEventFilter): def __call__(self, room: MatrixRoom, event: Event) -> bool: try: - parts = [p.strip() for p in event.body.split() if p.strip()] + 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