diff --git a/README.md b/README.md index 801c43b..366f94b 100644 --- a/README.md +++ b/README.md @@ -18,7 +18,7 @@ The package supports the following features: Use `pip` to install this package: ```bash -python -m pip install git+https://git.tyukalov.su/nikita/mab@v0.3.0 +python -m pip install git+https://git.tyukalov.su/nikita/mab@v0.4.0 ``` You should specify package version you want to use, because `main` without tags @@ -27,18 +27,16 @@ contain unstable code. ## Basic usage This is the most simple bot you can create. It would respond to any message -that starts with `!test`, `!hello` or `!hi`. +that contains `hello` and `hi` words. ```python import asyncio +from pathlib import Path from mab import MatrixBot, MatrixBotConfig -from mab import TextCommandFilter +from mab import TextContainsFilter, SenderIsBotFilter from mab.types import RoomEventData -async def on_valid_command(data: RoomEventData) -> None: - # do not respond to ourselves - if event.sender == data.bot.get_client().user_id: - return +async def on_message(data: RoomEventData) -> None: text = f"Your message contains {len(data.event.body)} symbols" await data.bot.send_text_to_room(data.room, text) @@ -51,8 +49,8 @@ async def main() -> None: ) bot = MatrixBot(matrix_bot_config) bot.add_callback( - TextCommandFilter(["test", "hello", "hi"]), - on_valid_command + ~SenderIsBotFilter() & TextContainsFilter(["test", "hello", "hi"]), + on_message ) await bot.start() # wait for Ctrl+C diff --git a/pyproject.toml b/pyproject.toml index 31e453a..775fb43 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta" [project] name = "mab" -version = "0.3.0" +version = "0.4.0" authors = [ { name = "Tyukalov Nikita", email = "nikita@tyukalov.su" } ] diff --git a/src/mab/bot/_client_manager.py b/src/mab/bot/_client_manager.py index e37f4aa..eabb28e 100644 --- a/src/mab/bot/_client_manager.py +++ b/src/mab/bot/_client_manager.py @@ -191,11 +191,13 @@ class ClientManager: raise RuntimeError("The bot was never started") self._background_task.cancel() try: - asyncio.shield(self._background_task) + await asyncio.shield(self._background_task) + except asyncio.CancelledError: + pass except: self._logger.error(traceback.format_exc()) try: - asyncio.shield(self._close_client()) + await asyncio.shield(self._close_client()) except: self._logger.error(traceback.format_exc()) self._background_task = None diff --git a/src/mab/filters/__init__.py b/src/mab/filters/__init__.py new file mode 100644 index 0000000..55cb0eb --- /dev/null +++ b/src/mab/filters/__init__.py @@ -0,0 +1,28 @@ +from .base import * +from .message import * +from .room import * +from .text import * + +__all__ = [ + # base.py + "BaseEventFilter", + "EventTypeFilter", + + # message.py + "NewMessageFilter", + "EditedMessageFilter", + "RedactedMessageFilter", + "SenderIsFilter", + "SenderIsBotFilter", + + # room.py + "RoomEncryptedFilter", + + # text.py + "TextFilter", + "TextContainsFilter", + "TextStartsWithFilter", + "TextEndsWithFilter", + "TextCommandFilter", + "TextRegexFilter" +] \ No newline at end of file diff --git a/src/mab/filters/base.py b/src/mab/filters/base.py index 3c5a2b8..5ebde04 100644 --- a/src/mab/filters/base.py +++ b/src/mab/filters/base.py @@ -1,5 +1,6 @@ from abc import ABC, abstractmethod import logging +from typing import Any, Type from nio import AsyncClient from nio import MatrixRoom, Event @@ -52,30 +53,38 @@ class BaseEventFilter(ABC): ) # PAYLOAD - @abstractmethod def __repr__(self) -> str: - """This method must be redefined in derived classes to improve - debugging experience. """ - pass + This method may be redefined in derived classes to improve debugging + experience. + """ + return str(self.__class__.__name__) - @abstractmethod async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> 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: - room - room the event has happened in - event - the event to check againts this filter - client - the client - - Returns: - True if the event satisfies this filter - False if the event does not satisfy this filter """ - pass + 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: + - room - room the event has happened in + - event - the event to check againts this filter + - client - the client + + Returns: + - True if the event satisfies this filter + - False if the event does not satisfy this filter + """ + return True + +class EventTypeFilter(BaseEventFilter): + """Event filter that checks if the event is an instance of some class""" + def __init__(self, type: Type): + self._type = type + + async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool: + return isinstance(event, self._type) class CompoundEventFilter(BaseEventFilter): """Event filter that consists of multiple filters""" diff --git a/src/mab/filters/message.py b/src/mab/filters/message.py new file mode 100644 index 0000000..c3f7dee --- /dev/null +++ b/src/mab/filters/message.py @@ -0,0 +1,87 @@ +import traceback +from .base import BaseEventFilter, EventTypeFilter + +from nio import AsyncClient +from nio import MatrixRoom, Event + +from nio import RedactionEvent + +class NewMessageFilter(BaseEventFilter): + """ + This filter returns True if the event is a new message. Most filters are + derived from this base class because it ignores events about edited + messages. + """ + def __init__(self, **kwargs): + super().__init__(**kwargs) + + async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool: + if not await super().__call__(room, event, client): + return False + return "m.new_content" not in event.source["content"] + +class EditedMessageFilter(BaseEventFilter): + """ + This filter returns True if the event is an edited message. You may use this + filter to create callbacks that are called if the message gets edited. + """ + def __init__(self, **kwargs): + super().__init__(**kwargs) + + async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool: + if not await super().__call__(room, event, client): + return False + return "m.new_content" in event.source["content"] + +class RedactedMessageFilter(EventTypeFilter): + """ + This filter returns True if the event is a RedactionEvent. + """ + def __init__(self, **kwargs): + super().__init__(RedactionEvent, **kwargs) + + async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool: + return await super().__call__(room, event, client) + +class SenderIsFilter(BaseEventFilter): + """ + This filter returns True if `event.sender` is any of specified senders. + `event.sender` is converted to lower case if `any_case` is True (default). + Supplied sender list is NEVER converted to lower case, so it is your duty to + use lower case if `any_case` is True. + + `senders` list is stored by reference so you can modify behavior of this + filter dynamically. + """ + def __init__(self, sender: list[str] | str, *, any_case: bool = True, **kwargs): + super().__init__(**kwargs) + if isinstance(sender, str): + sender = [sender] + self._sender = sender + self._any_case = any_case + + async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool: + if not await super().__call__(room, event, client): + return False + sender = event.sender.lower() if self._any_case else event.sender + for s in self._sender: + if sender == s: + return True + return False + +class SenderIsBotFilter(BaseEventFilter): + """ + This filter returns True if `event.sender` is the client that has received + the event. You may use this filter to set callbacks for messages sent by + other users by using the following syntax: + ```py + ~SenderIsBotFilter() + ``` + """ + def __init__(self, **kwargs): + super().__init__(**kwargs) + + async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool: + if not await super().__call__(room, event, client): + return False + return client.user_id == event.sender \ No newline at end of file diff --git a/src/mab/filters/room.py b/src/mab/filters/room.py index 9f02f56..f699cf3 100644 --- a/src/mab/filters/room.py +++ b/src/mab/filters/room.py @@ -3,93 +3,6 @@ from .base import BaseEventFilter from nio import AsyncClient from nio import MatrixRoom, Event -class RoomIdContainsFilter(BaseEventFilter): - """ - This filter returns True if the `room.room_id` contains `needle` (or - any of needles from the list). The check will be case insensetive if - `any_case` is True. - """ - def __init__(self, needle: str | list[str], *, any_case: bool = True): - super().__init__() - if type(needle) is str: - 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"RoomIdContainsFilter({repr(self._needle)}, any_case={repr(self._any_case)})" - - async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool: - try: - room_id = room.room_id.lower() if self._any_case else room.room_id - for s in self._needle: - if s in room_id: - return True - return False - except: - return False - -class RoomIdStartsWithFilter(BaseEventFilter): - """ - This filter returns True if the `room.room_id` starts with `substring` - (or any of substrings from the list). The check will be case insensetive - if `any_case` is True. - """ - def __init__(self, substring: str | list[str], *, any_case: bool = True): - super().__init__() - if type(substring) is str: - 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"RoomIdStartsWithFilter({repr(self._substring)}, any_case={repr(self._any_case)})" - - async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool: - try: - room_id = room.room_id.lower() if self._any_case else room.room_id - for s in self._substring: - if room_id.startswith(s): - return True - return False - except: - return False - -class RoomIdEndsWithFilter(BaseEventFilter): - """ - This filter returns True if the `room.room_id` ends with `substring` (or - any of substrings from the list). The check will be case insensetive - if `any_case` is True. - """ - def __init__(self, substring: str | list[str], *, any_case: bool = True): - super().__init__() - if type(substring) is str: - 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"RoomIdEndsWithFilter({repr(self._substring)}, any_case={repr(self._any_case)})" - - async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool: - try: - room_id = room.room_id.lower() if self._any_case else room.room_id - for s in self._substring: - if room_id.endswith(s): - return True - return False - except: - return False - class RoomEncryptedFilter(BaseEventFilter): """ This filter returns True if the room is encrypted. @@ -97,9 +10,6 @@ class RoomEncryptedFilter(BaseEventFilter): def __init__(self): super().__init__() - def __repr__(self): - return f"RoomEncryptedFilter()" - async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool: try: return room.encrypted diff --git a/src/mab/filters/text.py b/src/mab/filters/text.py index e2862f2..afb4f5f 100644 --- a/src/mab/filters/text.py +++ b/src/mab/filters/text.py @@ -1,11 +1,11 @@ import re import traceback -from .base import BaseEventFilter +from .message import NewMessageFilter from nio import AsyncClient from nio import MatrixRoom, Event -class TextFilter(BaseEventFilter): +class TextFilter(NewMessageFilter): """ This filter returns True if all conditions are met: 1. `event` has attribute `body` @@ -23,9 +23,6 @@ class TextFilter(BaseEventFilter): super().__init__(**kwargs) self._ignore_filename_in_body = ignore_filename_in_body - def __repr__(self) -> str: - return "TextFilter()" - async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool: if not hasattr(event, "body"): return False @@ -39,31 +36,6 @@ class TextFilter(BaseEventFilter): return False return True -class FormattedTextFilter(BaseEventFilter): - """ - This filter returns True if all conditions are met: - 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, **kwargs): - super().__init__(**kwargs) - - def __repr__(self): - return "FormattedTextFilter()" - - async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool: - 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(TextFilter): """ This filter returns True if `event.body` contains `needle` substring (or any @@ -84,9 +56,6 @@ class TextContainsFilter(TextFilter): self._any_case = any_case self._needle = needle - def __repr__(self) -> str: - return f"TextContainsFilter({repr(self._needle)}, any_case={repr(self._any_case)})" - async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool: if not await super().__call__(room, event, client): return False @@ -116,9 +85,6 @@ class TextStartsWithFilter(TextFilter): self._any_case = any_case self._substring = substring - def __repr__(self): - return f"TextStartsWithFilter({repr(self._substring)}, any_case={repr(self._any_case)})" - async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool: if not await super().__call__(room, event, client): return False @@ -148,9 +114,6 @@ class TextEndsWithFilter(TextFilter): self._any_case = any_case self._substring = substring - def __repr__(self): - return f"TextEndsWithFilter({repr(self._substring)}, any_case={repr(self._any_case)})" - async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool: if not await super().__call__(room, event, client): return False @@ -190,9 +153,6 @@ class TextCommandFilter(TextFilter): 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)})" - async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool: if not await super().__call__(room, event, client): return False @@ -221,9 +181,6 @@ class TextRegexFilter(TextFilter): 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