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
This commit is contained in:
2026-09-09 17:25:23 +03:00
parent 3d74cb737b
commit d14e110525

View File

@@ -1,3 +1,5 @@
import re
import traceback
from .base import BaseEventFilter from .base import BaseEventFilter
from nio import AsyncClient from nio import AsyncClient
@@ -5,22 +7,38 @@ from nio import MatrixRoom, Event
class TextFilter(BaseEventFilter): class TextFilter(BaseEventFilter):
""" """
This filter returns True if the event contains `body` attribute. This filter returns True if all conditions are met:
`body` attribute contains unformatted text, string. 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): def __init__(self):
super().__init__() super().__init__()
def __repr__(self): def __repr__(self) -> str:
return "TextFilter()" return "TextFilter()"
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool: async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
return hasattr(event, "body") and type(event.body) is str # type: ignore 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): class FormattedTextFilter(BaseEventFilter):
""" """
This filter returns True if the event contains valid `formatted_body` This filter returns True if all conditions are met:
attribute. `formatted_body` attribute contains formatted text, string. 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): def __init__(self):
super().__init__() super().__init__()
@@ -29,114 +47,135 @@ class FormattedTextFilter(BaseEventFilter):
return "FormattedTextFilter()" return "FormattedTextFilter()"
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> 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 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` This filter returns True if `event.body` contains `needle` substring (or any
substring (or any of neddle from the list). The check will be case of neddle from the list). `event.body` will be converted to lower case if
insensetive if `any_case` is True. `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): def __init__(self, needle: str | list[str], *, any_case: bool = True):
super().__init__() super().__init__()
if type(needle) is str: if type(needle) is str:
needle = [needle] needle = [needle]
self._any_case = any_case self._any_case = any_case
if self._any_case: self._needle = needle
self._needle = [s.lower() for s in needle]
else:
self._needle = list(needle)
def __repr__(self): def __repr__(self) -> str:
return f"TextContainsFilter({repr(self._needle)}, any_case={repr(self._any_case)})" return f"TextContainsFilter({repr(self._needle)}, any_case={repr(self._any_case)})"
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool: async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
try: if not await super().__call__(room, event, client):
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 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
class TextStartsWithFilter(BaseEventFilter): class TextStartsWithFilter(TextFilter):
""" """
This filter returns True if the `event.body` starts with `substring` (or This filter returns True if `event.body` starts with `substring` (or any of
any of substrings from the list). The check will be case insensetive if substrings from the list). The check will be case insensetive if `any_case`
`any_case` is True. 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): def __init__(self, substring: str | list[str], *, any_case: bool = True):
super().__init__() super().__init__()
if type(substring) is str: if type(substring) is str:
substring = [substring] substring = [substring]
self._any_case = any_case self._any_case = any_case
if self._any_case: self._substring = substring
self._substring = [s.lower() for s in substring]
else:
self._substring = list(substring)
def __repr__(self): def __repr__(self):
return f"TextStartsWithFilter({repr(self._substring)}, any_case={repr(self._any_case)})" return f"TextStartsWithFilter({repr(self._substring)}, any_case={repr(self._any_case)})"
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool: async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
try: if not await super().__call__(room, event, client):
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 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
class TextEndsWithFilter(BaseEventFilter): class TextEndsWithFilter(TextFilter):
""" """
This filter returns True if the `event.body` ends with `substring` (or This filter returns True if `event.body` ends with `substring` (or any of
any of substrings from the list). The check will be case insensetive if substrings from the list). The check will be case insensetive if `any_case`
`any_case` is True. 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): def __init__(self, substring: str | list[str], *, any_case: bool = True):
super().__init__() super().__init__()
if type(substring) is str: if type(substring) is str:
substring = [substring] substring = [substring]
self._any_case = any_case self._any_case = any_case
if self._any_case: self._substring = substring
self._substring = [s.lower() for s in substring]
else:
self._substring = list(substring)
def __repr__(self): def __repr__(self):
return f"TextEndsWithFilter({repr(self._substring)}, any_case={repr(self._any_case)})" return f"TextEndsWithFilter({repr(self._substring)}, any_case={repr(self._any_case)})"
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool: async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
try: if not await super().__call__(room, event, client):
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 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
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() 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` 3. First element of splitted `event.body` starts with `prefix`
4. First element of splitted `event.body` (after stripping `prefix`) 4. First element of splitted `event.body` (after lstripping `prefix`) starts
starts with any of strings in `verbs` list (case-insensitive) with any of strings in `verbs` list
Remarks: `verbs` list is stored by reference so you can dynamically edit behavior of
- If this filter is satified, then it will set a new attribute for the this filter.
event: `event.command_args: list[str]`. You may use this attribute in
your callback for this event. 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 = "!"): def __init__(self, verbs: str | list[str], min_args: int = 0, max_args: int | None = None, prefix: str = "!"):
super().__init__() super().__init__()
if type(verbs) is str: if type(verbs) is str:
verbs = [verbs] verbs = [verbs]
verbs = [v.lower() for v in verbs]
self._verbs = verbs self._verbs = verbs
self._min_args = min_args self._min_args = min_args
self._max_args = max_args self._max_args = max_args
@@ -146,20 +185,40 @@ class TextCommandFilter(BaseEventFilter):
return f"TextCommandFilter({repr(self._verbs)}, {repr(self._min_args)}, {repr(self._max_args)}, {repr(self._prefix)})" return f"TextCommandFilter({repr(self._verbs)}, {repr(self._min_args)}, {repr(self._max_args)}, {repr(self._prefix)})"
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
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
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
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: try:
parts = [p.strip() for p in event.body.split() if p.strip()] # type: ignore return self._regex.match(event.body) # type: ignore
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: except:
self._logger.error(traceback.format_exc())
return False return False