From 4287a0d20c92c15086a4c096b6b7820d31c9b315 Mon Sep 17 00:00:00 2001 From: nikita Date: Sat, 12 Sep 2026 22:19:28 +0300 Subject: [PATCH] Improved filters system - `RoomEventData` is renamed to `EventContext` and moved to context.py - `EventContext.filter` is removed - Added context variables system which improves type hints and simplifies callbacks code - `BodyCommandFilter` sets context variables from now on - Added `super().__call__` invocation to filters implements in base.py - Examples are updated to include required changes --- examples/echo_bot.py | 14 ++----- examples/image_gen_bot.py | 16 ++------ src/mab/__init__.py | 13 ++++++- src/mab/bot/_callbacks.py | 20 +++++----- src/mab/bot/bot.py | 5 ++- src/mab/context.py | 76 ++++++++++++++++++++++++++++++++++++++ src/mab/filters/base.py | 22 +++++++---- src/mab/filters/body.py | 65 +++++++++++++++++++------------- src/mab/filters/message.py | 49 ++++++++++++++---------- src/mab/types.py | 42 +++++++++------------ 10 files changed, 207 insertions(+), 115 deletions(-) create mode 100644 src/mab/context.py diff --git a/examples/echo_bot.py b/examples/echo_bot.py index 3734ea4..3a2b075 100644 --- a/examples/echo_bot.py +++ b/examples/echo_bot.py @@ -9,21 +9,13 @@ import asyncio import os import logging -from mab import ( - MatrixBot, - MatrixBotConfig, - RoomEventData, - BodyExistsFilter, - MessageTypeFilter, - SenderIsBotFilter, - MessageType -) +from mab import * from _environment import check_environment -async def on_text_message(data: RoomEventData) -> None: +async def on_text_message(ctx: EventContext) -> None: """This callback is called when a text message arrives.""" - await data.bot.send_text(data.room, data.event.body) # type: ignore + await ctx.bot.send_text(ctx.room, ctx[CTX_BODY]) async def main() -> None: """Application entry point""" diff --git a/examples/image_gen_bot.py b/examples/image_gen_bot.py index 7a85830..dede3c9 100644 --- a/examples/image_gen_bot.py +++ b/examples/image_gen_bot.py @@ -13,23 +13,15 @@ import logging import random from PIL import Image -from mab import ( - MatrixBot, - MatrixBotConfig, - RoomEventData, - MessageTypeFilter, - BodyCommandFilter, - SenderIsBotFilter, - MessageType -) +from mab import * from _environment import check_environment -async def on_gen_command(data: RoomEventData) -> None: +async def on_gen_command(data: EventContext) -> None: """This callback is called when `!gen R G B` command is received.""" # convert R, G and B to floats try: - r, g, b = [float(v) for v in data.event.command_args] # type: ignore + r, g, b = [float(v) for v in data[CTX_CMD_ARGS]] except: await data.bot.send_text(data.room, "Invalid arguments") return @@ -53,7 +45,7 @@ async def on_gen_command(data: RoomEventData) -> None: # send await data.bot.send_image_bytes(data.room, buf, "noise.png") -async def on_wrong_message(data: RoomEventData) -> None: +async def on_wrong_message(data: EventContext) -> None: """This callback is called when a wrong message is received.""" await data.bot.send_text( data.room, diff --git a/src/mab/__init__.py b/src/mab/__init__.py index c120ece..40d85a0 100644 --- a/src/mab/__init__.py +++ b/src/mab/__init__.py @@ -1,7 +1,8 @@ from . import bot from . import types -from .types import MatrixBotConfig, RoomEventData, MessageType +from .types import MatrixBotConfig, MessageType +from .context import * from .bot import MatrixBot @@ -16,9 +17,17 @@ __all__ = [ # .types "MatrixBotConfig", - "RoomEventData", "MessageType", + # .context + "EventContext", + "CTX_BODY", + "CTX_MESSAGE_TYPE", + "CTX_SENDER", + "CTX_CMD_PREFIX", + "CTX_CMD_VERB", + "CTX_CMD_ARGS", + # .bot "MatrixBot", diff --git a/src/mab/bot/_callbacks.py b/src/mab/bot/_callbacks.py index 761f333..15ef831 100644 --- a/src/mab/bot/_callbacks.py +++ b/src/mab/bot/_callbacks.py @@ -10,7 +10,8 @@ from nio import MatrixInvitedRoom, InviteMemberEvent, JoinResponse from nio.events.room_events import Event as RoomEvent from ._storage import Storage -from ..types import MatrixBotConfig, RoomEventData +from ..types import MatrixBotConfig +from ..context import EventContext from ..filters.base import BaseEventFilter if TYPE_CHECKING: @@ -31,7 +32,7 @@ class Callbacks: filter: BaseEventFilter """Filter to use for matching""" - callback: Callable[[RoomEventData], Coroutine[Any, Any, None]] | None + callback: Callable[[EventContext], Coroutine[Any, Any, None]] | None """Callback that will be called if the filter matches""" stop_matching: bool @@ -51,20 +52,19 @@ class Callbacks: for callback_info in self._filters: if not isinstance(callback_info, self._FilterBasedCallback): continue + event_data = EventContext( + room=room, + event=event, + bot=self._matrix_bot + ) try: - if not await callback_info.filter(room, event, self._client): + if not await callback_info.filter(event_data): continue except asyncio.CancelledError: raise except: self._logger.error(traceback.format_exc()) continue - event_data = RoomEventData( - room=room, - event=event, - filter=callback_info.filter, - bot=self._matrix_bot - ) try: # dump argument types if callback_info.callback is None: @@ -144,7 +144,7 @@ class Callbacks: def add_room_event_callback( self, filter: BaseEventFilter, - callback: Callable[[RoomEventData], Coroutine[Any, Any, None]] | None, + callback: Callable[[EventContext], Coroutine[Any, Any, None]] | None, *, stop_matching: bool = True) -> None: """ diff --git a/src/mab/bot/bot.py b/src/mab/bot/bot.py index 4b472a2..fda11b2 100644 --- a/src/mab/bot/bot.py +++ b/src/mab/bot/bot.py @@ -3,10 +3,11 @@ import logging from typing import Callable, Coroutine, Any -from nio import AsyncClient +from nio import AsyncClient, MatrixRoom from ..filters.base import BaseEventFilter from ..types import * +from ..context import EventContext from ._validation import Validator from ._storage import Storage @@ -44,7 +45,7 @@ class MatrixBot: def add_callback(self, filter: BaseEventFilter, - callback: Callable[[RoomEventData], Coroutine[Any, Any, None]] | None, + callback: Callable[[EventContext], Coroutine[Any, Any, None]] | None, *, stop_matching: bool = True) -> None: """ diff --git a/src/mab/context.py b/src/mab/context.py new file mode 100644 index 0000000..9ec178f --- /dev/null +++ b/src/mab/context.py @@ -0,0 +1,76 @@ +"""This module implements logic for event context""" + +from dataclasses import dataclass +from typing import Any, TYPE_CHECKING + +from .types import ContextDataKey + +from nio import MatrixRoom, Event + +if TYPE_CHECKING: + from .bot import MatrixBot + +# +# Possible context variables +# +CTX_BODY = ContextDataKey[str]("CTX_BODY") +"""Value of `event.body`""" + +CTX_MESSAGE_TYPE = ContextDataKey[str]("CTX_MESSAGE_TYPE") +"""Value of `msgtype` for the event""" + +CTX_SENDER = ContextDataKey[str]("CTX_SENDER") +"""Value of `event.sender`""" + +CTX_CMD_PREFIX = ContextDataKey[str]("CTX_CMD_PREFIX") +"""Command prefix that was used when matching""" + +CTX_CMD_VERB = ContextDataKey[str]("CTX_CMD_VERB") +"""The verb that was used to execute the command""" + +CTX_CMD_ARGS = ContextDataKey[list[str]]("CTX_CMD_ARGS") +"""Arguments that were passed with the command""" + + +# +# EventContext implementation +# +@dataclass +class EventContext: + """The class holding information about an event that happened in the room""" + + room: MatrixRoom + """The room the event has happened in""" + + event: Event + """The event that has happened in the room""" + + bot: "MatrixBot" + """The bot that is the source of the event""" + + def __setitem__[T](self, key: ContextDataKey[T], value: T | None) -> None: + """Set a value inside the context data storage. `None` removes it""" + if not hasattr(self, "_datastore"): + self._datastore: dict[ContextDataKey, Any] = {} + if value is None: + del self._datastore[key] + else: + self._datastore[key] = value + + def __getitem__[T](self, key: ContextDataKey[T]) -> T: + """ + Get a value inside the context data storage. + + Raises RuntimeError if the value is not present. + """ + if not hasattr(self, "_datastore") or key not in self._datastore: + raise RuntimeError(f"Context does not contain {repr(key)}") + return self._datastore[key] + + def __contains__[T](self, key: ContextDataKey[T]) -> bool: + """Check if context data storage contains the value""" + if not hasattr(self, "_datastore"): + return False + if key not in self._datastore: + return False + return True \ No newline at end of file diff --git a/src/mab/filters/base.py b/src/mab/filters/base.py index 9622e8a..549a63c 100644 --- a/src/mab/filters/base.py +++ b/src/mab/filters/base.py @@ -5,6 +5,8 @@ from typing import Any, Type from nio import AsyncClient from nio import MatrixRoom, Event +from ..context import EventContext + class BaseEventFilter(ABC): """Base class for all message filters""" _logger = logging.Logger("EventFilter") @@ -64,7 +66,7 @@ class BaseEventFilter(ABC): """ return str(self.__class__.__name__) - async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool: + async def __call__(self, context: EventContext) -> bool: """ This abstract method must be redefined in derived classes so that the filter operates according to its description. This method must not raise @@ -72,9 +74,9 @@ class BaseEventFilter(ABC): and return False Args: - - room - room the event has happened in - - event - the event to check againts this filter - - client - the client + - context - event context; your derived classes may add variables + to it (see `message.MessageTypeFilter` implementation + for reference) Returns: - True if the event satisfies this filter @@ -88,8 +90,10 @@ class EventTypeFilter(BaseEventFilter): super().__init__(**kwargs) self._type = type - async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool: - return isinstance(event, self._type) + async def __call__(self, context: EventContext) -> bool: + if not await super().__call__(context): + return False + return isinstance(context.event, self._type) class CompoundEventFilter(BaseEventFilter): """Event filter that consists of multiple filters""" @@ -140,8 +144,10 @@ class CompoundEventFilter(BaseEventFilter): expression = f"~{reprs[0]}" return f"({expression})" - async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool: - evaluated = [await arg(room, event, client) for arg in self._arguments] + async def __call__(self, context: EventContext) -> bool: + if not await super().__call__(context): + return False + evaluated = [await arg(context) for arg in self._arguments] if self._operator == self.OPERATOR_AND: return all(evaluated) elif self._operator == self.OPERATOR_OR: diff --git a/src/mab/filters/body.py b/src/mab/filters/body.py index 4501384..bf3879f 100644 --- a/src/mab/filters/body.py +++ b/src/mab/filters/body.py @@ -1,6 +1,13 @@ import re import traceback from .message import NewMessageFilter +from ..context import EventContext +from ..context import ( + CTX_BODY, + CTX_CMD_PREFIX, + CTX_CMD_VERB, + CTX_CMD_ARGS +) from nio import AsyncClient from nio import MatrixRoom, Event @@ -21,24 +28,27 @@ class BodyExistsFilter(NewMessageFilter): This filter will match any message that has `body` in it, including images, videos, files, etc. + + This filter sets `CTX_BODY` context variable. """ def __init__(self, *, ignore_filename_in_body: bool = True, **kwargs): super().__init__(**kwargs) self._ignore_filename_in_body = ignore_filename_in_body - async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool: - if not await super().__call__(room, event, client): + async def __call__(self, context: EventContext) -> bool: + if not await super().__call__(context): return False - if not hasattr(event, "body"): + if not hasattr(context.event, "body"): return False - if not isinstance(event.body, str): # type: ignore + if not isinstance(context.event.body, str): # type: ignore return False - if not event.body.strip(): # type: ignore + if not context.event.body.strip(): # type: ignore return False if self._ignore_filename_in_body: - content = event.source["content"] - if "filename" in content and content["filename"] == event.body: # type: ignore + content = context.event.source["content"] + if "filename" in content and content["filename"] == context.event.body: # type: ignore return False + context[CTX_BODY] = context.event.body # type: ignore return True class BodyContainsFilter(BodyExistsFilter): @@ -61,10 +71,10 @@ class BodyContainsFilter(BodyExistsFilter): self._any_case = any_case self._needle = needle - async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool: - if not await super().__call__(room, event, client): + async def __call__(self, context: EventContext) -> bool: + if not await super().__call__(context): return False - body = event.body.lower() if self._any_case else event.body # type: ignore + body = context.event.body.lower() if self._any_case else context.event.body # type: ignore for n in self._needle: if n in body: return True @@ -90,10 +100,10 @@ class BodyStartsWithFilter(BodyExistsFilter): self._any_case = any_case self._substring = substring - async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool: - if not await super().__call__(room, event, client): + async def __call__(self, context: EventContext) -> bool: + if not await super().__call__(context): return False - body = event.body.lower() if self._any_case else event.body # type: ignore + body = context.event.body.lower() if self._any_case else context.event.body # type: ignore for s in self._substring: if body.startswith(s): return True @@ -119,10 +129,10 @@ class BodyEndsWithFilter(BodyExistsFilter): self._any_case = any_case self._substring = substring - async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool: - if not await super().__call__(room, event, client): + async def __call__(self, context: EventContext) -> bool: + if not await super().__call__(context): return False - body = event.body.lower() if self._any_case else event.body # type: ignore + body = context.event.body.lower() if self._any_case else context.event.body # type: ignore for s in self._substring: if body.endswith(s): return True @@ -145,9 +155,10 @@ class BodyCommandFilter(BodyExistsFilter): 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. + This filter sets the following context variables: + - `CTX_CMD_PREFIX` - prefix that was used + - `CTX_CMD_VERB` - verb that was used + - `CTX_CMD_ARGS` - arguments that were passed """ def __init__(self, verbs: str | list[str], min_args: int = 0, max_args: int | None = None, prefix: str = "!", **kwargs): super().__init__(**kwargs) @@ -158,10 +169,10 @@ class BodyCommandFilter(BodyExistsFilter): self._max_args = max_args self._prefix = prefix - async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool: - if not await super().__call__(room, event, client): + async def __call__(self, context: EventContext) -> bool: + if not await super().__call__(context): return False - parts = [p.strip() for p in event.body.split() if p.strip()] # type: ignore + parts = [p.strip() for p in context.event.body.split() if p.strip()] # type: ignore args_count = len(parts) - 1 if args_count < self._min_args: return False @@ -172,7 +183,9 @@ class BodyCommandFilter(BodyExistsFilter): cmd = parts[0][len(self._prefix):].lower() for verb in self._verbs: if cmd == verb: - setattr(event, "command_args", parts[1:]) + context[CTX_CMD_PREFIX] = self._prefix + context[CTX_CMD_VERB] = verb + context[CTX_CMD_ARGS] = parts[1:] return True return False @@ -186,11 +199,11 @@ class BodyRegexFilter(BodyExistsFilter): regex = re.compile(regex) self._regex = regex - async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool: - if not await super().__call__(room, event, client): + async def __call__(self, context: EventContext) -> bool: + if not await super().__call__(context): return False try: - return self._regex.match(event.body) is not None # type: ignore + return self._regex.match(context.event.body) is not None # type: ignore except: self._logger.error(traceback.format_exc()) return False \ No newline at end of file diff --git a/src/mab/filters/message.py b/src/mab/filters/message.py index f3deaa7..8b875a9 100644 --- a/src/mab/filters/message.py +++ b/src/mab/filters/message.py @@ -1,6 +1,7 @@ import traceback from .base import BaseEventFilter, EventTypeFilter from ..types import MessageType +from ..context import EventContext, CTX_MESSAGE_TYPE, CTX_SENDER from nio import AsyncClient from nio import MatrixRoom, Event @@ -15,6 +16,8 @@ class MessageTypeFilter(BaseEventFilter): `types` list is stored by reference so you may modify the behavior of this filter dynamically. + + This filter sets `CTX_MESSAGE_TYPE` variable in the context. """ def __init__(self, types: list[MessageType] | MessageType, **kwargs): super().__init__(**kwargs) @@ -22,14 +25,16 @@ class MessageTypeFilter(BaseEventFilter): types = [types] self._types = types - async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool: - if not await super().__call__(room, event, client): + async def __call__(self, context: EventContext) -> bool: + if not await super().__call__(context): return False - if "msgtype" not in event.source["content"]: + if "msgtype" not in context.event.source["content"]: return False - return ( - event.source["content"]["msgtype"] in [t.value for t in self._types] - ) + msgtype = context.event.source["content"]["msgtype"] + if not msgtype in [t.value for t in self._types]: + return False + context[CTX_MESSAGE_TYPE] = msgtype + return True class NewMessageFilter(BaseEventFilter): """ @@ -40,10 +45,10 @@ class NewMessageFilter(BaseEventFilter): 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): + async def __call__(self, context: EventContext) -> bool: + if not await super().__call__(context): return False - return "m.new_content" not in event.source["content"] + return "m.new_content" not in context.event.source["content"] class EditedMessageFilter(BaseEventFilter): """ @@ -53,10 +58,10 @@ class EditedMessageFilter(BaseEventFilter): 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): + async def __call__(self, context: EventContext) -> bool: + if not await super().__call__(context): return False - return "m.new_content" in event.source["content"] + return "m.new_content" in context.event.source["content"] class RedactedMessageFilter(EventTypeFilter): """ @@ -64,9 +69,10 @@ class RedactedMessageFilter(EventTypeFilter): """ def __init__(self, **kwargs): super().__init__(RedactionEvent, **kwargs) + raise NotImplementedError() - async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool: - return await super().__call__(room, event, client) + async def __call__(self, context: EventContext) -> bool: + raise NotImplementedError() class SenderIsFilter(BaseEventFilter): """ @@ -77,6 +83,8 @@ class SenderIsFilter(BaseEventFilter): `senders` list is stored by reference so you can modify behavior of this filter dynamically. + + This filter sets `CTX_SENDER` variable in the context. """ def __init__(self, sender: list[str] | str, *, any_case: bool = True, **kwargs): super().__init__(**kwargs) @@ -85,12 +93,13 @@ class SenderIsFilter(BaseEventFilter): 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): + async def __call__(self, context: EventContext) -> bool: + if not await super().__call__(context): return False - sender = event.sender.lower() if self._any_case else event.sender + sender = context.event.sender.lower() if self._any_case else context.event.sender for s in self._sender: if sender == s: + context[CTX_SENDER] = sender return True return False @@ -106,7 +115,7 @@ class SenderIsBotFilter(BaseEventFilter): 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): + async def __call__(self, context: EventContext) -> bool: + if not await super().__call__(context): return False - return client.user_id == event.sender \ No newline at end of file + return context.bot.get_client().user_id == context.event.sender \ No newline at end of file diff --git a/src/mab/types.py b/src/mab/types.py index 4597db4..d3751dd 100644 --- a/src/mab/types.py +++ b/src/mab/types.py @@ -4,15 +4,8 @@ from pathlib import Path from dataclasses import dataclass from enum import Enum -from nio import MatrixRoom, Event from nio import UploadResponse -from .filters.base import BaseEventFilter - -from typing import TYPE_CHECKING -if TYPE_CHECKING: - from .bot import MatrixBot - @dataclass class MatrixBotConfig: """Configuration for MatrixBot""" @@ -68,22 +61,6 @@ class VideoFileProperties: thumbnail: Path | str | bytes | None = None """Path to the thumbnail or the raw JPEG thumbnail data""" -@dataclass -class RoomEventData: - """Dataclass that hold information about event that happened in the room""" - - room: MatrixRoom - """The room the event has happened in""" - - event: Event - """The event that has happened in the room""" - - filter: BaseEventFilter - """The filter that invoked this event""" - - bot: "MatrixBot" - """The bot that is the source of the event""" - @dataclass class UploadResult: """Result of data upload""" @@ -108,4 +85,21 @@ class MessageType(Enum): FILE = "m.file" AUDIO = "m.audio" LOCATION = "m.location" - VIDEO = "m.video" \ No newline at end of file + VIDEO = "m.video" + +class ContextDataKey[T]: + """ + Instances of this class represent a single possible data key that can be + stored inside EventContext. + """ + def __init__(self, name: str) -> None: + """ + Initialize a ContextDataKey + + Args: + - name - name that will be used internally + """ + self._name = name + + def __repr__(self) -> str: + return f"ContextDataKey[{type(T)}]({repr(self._name)})" \ No newline at end of file