Compare commits
3 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| b62473c468 | |||
| d8bf1818d7 | |||
| d41b588adc |
14
README.md
14
README.md
@@ -16,24 +16,26 @@ python -m pip install git+https://git.tyukalov.su/nikita/mab
|
||||
|
||||
## Basic usage
|
||||
|
||||
This is the most simple bot you can create
|
||||
This is the most simple bot you can create. It would respond to any message
|
||||
that starts with `!test`, `!hello` or `!hi`.
|
||||
|
||||
```python
|
||||
import asyncio
|
||||
from mab import MatrixBot, MatrixBotConfig
|
||||
from mab import TextCommandFilter
|
||||
|
||||
from nio import MatrixRoom, MatrixMessageText
|
||||
|
||||
bot: MatrixBot
|
||||
|
||||
async def on_room_message_text(room: MatrixRoom, event: RoomMessageText) -> None:
|
||||
async def on_valid_command(room: MatrixRoom, event: RoomMessageText) -> None:
|
||||
global bot
|
||||
# do not respond to ourselves
|
||||
if event.sender == bot.get_client().user_id:
|
||||
return
|
||||
text = f"You message contains {len(event.body)} symbols"
|
||||
text = f"Your message contains {len(event.body)} symbols"
|
||||
await bot.send_text_to_room(room.room_id, text)
|
||||
|
||||
|
||||
async def main() -> None:
|
||||
global bot
|
||||
# create and start the bot
|
||||
@@ -43,6 +45,10 @@ async def main() -> None:
|
||||
storage_directory=Path("storage_nagibator666")
|
||||
)
|
||||
bot = MatrixBot(matrix_bot_config)
|
||||
bot.add_event_callback(
|
||||
on_valid_command,
|
||||
TextCommandFilter(["test", "hello", "hi"])
|
||||
)
|
||||
bot.start()
|
||||
# wait for Ctrl+C
|
||||
try:
|
||||
|
||||
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
|
||||
|
||||
[project]
|
||||
name = "mab"
|
||||
version = "0.0.1"
|
||||
version = "0.1.0"
|
||||
authors = [
|
||||
{ name = "Tyukalov Nikita", email = "nikita@tyukalov.su" }
|
||||
]
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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
|
||||
|
||||
@@ -263,6 +286,7 @@ class MatrixBot:
|
||||
user_id=username,
|
||||
**session_data
|
||||
)
|
||||
self._client.load_store()
|
||||
result = await self._client.whoami()
|
||||
if type(result) is WhoamiError:
|
||||
self._logger.error(f"Can't log in using stored session data: '{result.message}'")
|
||||
@@ -369,7 +393,8 @@ class MatrixBot:
|
||||
sync_task = asyncio.create_task(
|
||||
self._client_cancellable_sync_forever(
|
||||
timeout=self.MATRIX_SYNC_PERIOD,
|
||||
since=(await self._read_next_batch())
|
||||
since=(await self._read_next_batch()),
|
||||
full_state=True
|
||||
)
|
||||
)
|
||||
try:
|
||||
@@ -402,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.
|
||||
@@ -449,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"""
|
||||
|
||||
138
src/mab/filters/base.py
Normal file
138
src/mab/filters/base.py
Normal file
@@ -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
|
||||
147
src/mab/filters/text.py
Normal file
147
src/mab/filters/text.py
Normal file
@@ -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
|
||||
Reference in New Issue
Block a user