Compare commits
3 Commits
v0.0.1
...
b62473c468
| 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
|
## 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
|
```python
|
||||||
import asyncio
|
import asyncio
|
||||||
from mab import MatrixBot, MatrixBotConfig
|
from mab import MatrixBot, MatrixBotConfig
|
||||||
|
from mab import TextCommandFilter
|
||||||
|
|
||||||
from nio import MatrixRoom, MatrixMessageText
|
from nio import MatrixRoom, MatrixMessageText
|
||||||
|
|
||||||
bot: MatrixBot
|
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
|
global bot
|
||||||
# do not respond to ourselves
|
# do not respond to ourselves
|
||||||
if event.sender == bot.get_client().user_id:
|
if event.sender == bot.get_client().user_id:
|
||||||
return
|
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)
|
await bot.send_text_to_room(room.room_id, text)
|
||||||
|
|
||||||
|
|
||||||
async def main() -> None:
|
async def main() -> None:
|
||||||
global bot
|
global bot
|
||||||
# create and start the bot
|
# create and start the bot
|
||||||
@@ -43,6 +45,10 @@ async def main() -> None:
|
|||||||
storage_directory=Path("storage_nagibator666")
|
storage_directory=Path("storage_nagibator666")
|
||||||
)
|
)
|
||||||
bot = MatrixBot(matrix_bot_config)
|
bot = MatrixBot(matrix_bot_config)
|
||||||
|
bot.add_event_callback(
|
||||||
|
on_valid_command,
|
||||||
|
TextCommandFilter(["test", "hello", "hi"])
|
||||||
|
)
|
||||||
bot.start()
|
bot.start()
|
||||||
# wait for Ctrl+C
|
# wait for Ctrl+C
|
||||||
try:
|
try:
|
||||||
|
|||||||
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
|
|||||||
|
|
||||||
[project]
|
[project]
|
||||||
name = "mab"
|
name = "mab"
|
||||||
version = "0.0.1"
|
version = "0.1.0"
|
||||||
authors = [
|
authors = [
|
||||||
{ name = "Tyukalov Nikita", email = "nikita@tyukalov.su" }
|
{ name = "Tyukalov Nikita", email = "nikita@tyukalov.su" }
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -5,6 +5,9 @@ from .types import MatrixBotConfig
|
|||||||
|
|
||||||
from .bot import MatrixBot
|
from .bot import MatrixBot
|
||||||
|
|
||||||
|
from .filters.base import *
|
||||||
|
from .filters.text import *
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
# module names
|
# module names
|
||||||
"bot",
|
"bot",
|
||||||
@@ -14,5 +17,16 @@ __all__ = [
|
|||||||
"MatrixBotConfig",
|
"MatrixBotConfig",
|
||||||
|
|
||||||
# .bot
|
# .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 MatrixInvitedRoom, InviteMemberEvent
|
||||||
from nio import JoinResponse
|
from nio import JoinResponse
|
||||||
|
|
||||||
import nio.events
|
from .filters.base import BaseEventFilter
|
||||||
|
|
||||||
|
from nio.events.room_events import Event as RoomEvemt
|
||||||
|
|
||||||
from .types import *
|
from .types import *
|
||||||
|
|
||||||
@@ -244,13 +246,34 @@ class MatrixBot:
|
|||||||
except:
|
except:
|
||||||
self._logger.error(traceback.format_exc())
|
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
|
# LIFECYCLE
|
||||||
#
|
#
|
||||||
def _setup_client_callbacks(self) -> None:
|
def _setup_client_callbacks(self) -> None:
|
||||||
"""Setup internal client callbacks"""
|
"""Setup internal client callbacks"""
|
||||||
self._client.add_response_callback(self._callback_sync, SyncResponse) # type: ignore
|
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:
|
if self._config.auto_join_any_room_on_invite:
|
||||||
self._client.add_event_callback(self._callback_autojoin, InviteMemberEvent) # type: ignore
|
self._client.add_event_callback(self._callback_autojoin, InviteMemberEvent) # type: ignore
|
||||||
|
|
||||||
@@ -263,6 +286,7 @@ class MatrixBot:
|
|||||||
user_id=username,
|
user_id=username,
|
||||||
**session_data
|
**session_data
|
||||||
)
|
)
|
||||||
|
self._client.load_store()
|
||||||
result = await self._client.whoami()
|
result = await self._client.whoami()
|
||||||
if type(result) is WhoamiError:
|
if type(result) is WhoamiError:
|
||||||
self._logger.error(f"Can't log in using stored session data: '{result.message}'")
|
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(
|
sync_task = asyncio.create_task(
|
||||||
self._client_cancellable_sync_forever(
|
self._client_cancellable_sync_forever(
|
||||||
timeout=self.MATRIX_SYNC_PERIOD,
|
timeout=self.MATRIX_SYNC_PERIOD,
|
||||||
since=(await self._read_next_batch())
|
since=(await self._read_next_batch()),
|
||||||
|
full_state=True
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
try:
|
try:
|
||||||
@@ -402,6 +427,7 @@ class MatrixBot:
|
|||||||
self._last_next_batch_dump: float = 0.0
|
self._last_next_batch_dump: float = 0.0
|
||||||
self._last_next_batch: str | None = None
|
self._last_next_batch: str | None = None
|
||||||
self._cb_password = self._default_password_callback
|
self._cb_password = self._default_password_callback
|
||||||
|
self._filters = []
|
||||||
|
|
||||||
def start(self) -> None:
|
def start(self) -> None:
|
||||||
"""Start the bot.
|
"""Start the bot.
|
||||||
@@ -449,12 +475,11 @@ class MatrixBot:
|
|||||||
result = True
|
result = True
|
||||||
return result
|
return result
|
||||||
|
|
||||||
def add_event_callback(self, callback: Callable[[Any], Awaitable[None]] | None, event_class: nio.events.Event) -> None:
|
def add_event_callback(self, callback: Callable[[Any], Awaitable[None]] | None, filter: BaseEventFilter, stop_after_this: bool = True) -> None:
|
||||||
"""Added event callback for events of specified class.
|
"""Add event callback for events that pass the filter."""
|
||||||
Use `None` instead of callback to print parameter types you need to use in your callback."""
|
|
||||||
if callback is None:
|
if callback is None:
|
||||||
callback = self._debug_event_callback
|
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:
|
def get_client(self) -> AsyncClient:
|
||||||
"""Get AsyncClient in use"""
|
"""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