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
This commit is contained in:
@@ -9,21 +9,13 @@ import asyncio
|
|||||||
import os
|
import os
|
||||||
import logging
|
import logging
|
||||||
|
|
||||||
from mab import (
|
from mab import *
|
||||||
MatrixBot,
|
|
||||||
MatrixBotConfig,
|
|
||||||
RoomEventData,
|
|
||||||
BodyExistsFilter,
|
|
||||||
MessageTypeFilter,
|
|
||||||
SenderIsBotFilter,
|
|
||||||
MessageType
|
|
||||||
)
|
|
||||||
|
|
||||||
from _environment import check_environment
|
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."""
|
"""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:
|
async def main() -> None:
|
||||||
"""Application entry point"""
|
"""Application entry point"""
|
||||||
|
|||||||
@@ -13,23 +13,15 @@ import logging
|
|||||||
import random
|
import random
|
||||||
from PIL import Image
|
from PIL import Image
|
||||||
|
|
||||||
from mab import (
|
from mab import *
|
||||||
MatrixBot,
|
|
||||||
MatrixBotConfig,
|
|
||||||
RoomEventData,
|
|
||||||
MessageTypeFilter,
|
|
||||||
BodyCommandFilter,
|
|
||||||
SenderIsBotFilter,
|
|
||||||
MessageType
|
|
||||||
)
|
|
||||||
|
|
||||||
from _environment import check_environment
|
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."""
|
"""This callback is called when `!gen R G B` command is received."""
|
||||||
# convert R, G and B to floats
|
# convert R, G and B to floats
|
||||||
try:
|
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:
|
except:
|
||||||
await data.bot.send_text(data.room, "Invalid arguments")
|
await data.bot.send_text(data.room, "Invalid arguments")
|
||||||
return
|
return
|
||||||
@@ -53,7 +45,7 @@ async def on_gen_command(data: RoomEventData) -> None:
|
|||||||
# send
|
# send
|
||||||
await data.bot.send_image_bytes(data.room, buf, "noise.png")
|
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."""
|
"""This callback is called when a wrong message is received."""
|
||||||
await data.bot.send_text(
|
await data.bot.send_text(
|
||||||
data.room,
|
data.room,
|
||||||
|
|||||||
@@ -1,7 +1,8 @@
|
|||||||
from . import bot
|
from . import bot
|
||||||
from . import types
|
from . import types
|
||||||
|
|
||||||
from .types import MatrixBotConfig, RoomEventData, MessageType
|
from .types import MatrixBotConfig, MessageType
|
||||||
|
from .context import *
|
||||||
|
|
||||||
from .bot import MatrixBot
|
from .bot import MatrixBot
|
||||||
|
|
||||||
@@ -16,9 +17,17 @@ __all__ = [
|
|||||||
|
|
||||||
# .types
|
# .types
|
||||||
"MatrixBotConfig",
|
"MatrixBotConfig",
|
||||||
"RoomEventData",
|
|
||||||
"MessageType",
|
"MessageType",
|
||||||
|
|
||||||
|
# .context
|
||||||
|
"EventContext",
|
||||||
|
"CTX_BODY",
|
||||||
|
"CTX_MESSAGE_TYPE",
|
||||||
|
"CTX_SENDER",
|
||||||
|
"CTX_CMD_PREFIX",
|
||||||
|
"CTX_CMD_VERB",
|
||||||
|
"CTX_CMD_ARGS",
|
||||||
|
|
||||||
# .bot
|
# .bot
|
||||||
"MatrixBot",
|
"MatrixBot",
|
||||||
|
|
||||||
|
|||||||
@@ -10,7 +10,8 @@ from nio import MatrixInvitedRoom, InviteMemberEvent, JoinResponse
|
|||||||
from nio.events.room_events import Event as RoomEvent
|
from nio.events.room_events import Event as RoomEvent
|
||||||
|
|
||||||
from ._storage import Storage
|
from ._storage import Storage
|
||||||
from ..types import MatrixBotConfig, RoomEventData
|
from ..types import MatrixBotConfig
|
||||||
|
from ..context import EventContext
|
||||||
from ..filters.base import BaseEventFilter
|
from ..filters.base import BaseEventFilter
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
@@ -31,7 +32,7 @@ class Callbacks:
|
|||||||
filter: BaseEventFilter
|
filter: BaseEventFilter
|
||||||
"""Filter to use for matching"""
|
"""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"""
|
"""Callback that will be called if the filter matches"""
|
||||||
|
|
||||||
stop_matching: bool
|
stop_matching: bool
|
||||||
@@ -51,20 +52,19 @@ class Callbacks:
|
|||||||
for callback_info in self._filters:
|
for callback_info in self._filters:
|
||||||
if not isinstance(callback_info, self._FilterBasedCallback):
|
if not isinstance(callback_info, self._FilterBasedCallback):
|
||||||
continue
|
continue
|
||||||
|
event_data = EventContext(
|
||||||
|
room=room,
|
||||||
|
event=event,
|
||||||
|
bot=self._matrix_bot
|
||||||
|
)
|
||||||
try:
|
try:
|
||||||
if not await callback_info.filter(room, event, self._client):
|
if not await callback_info.filter(event_data):
|
||||||
continue
|
continue
|
||||||
except asyncio.CancelledError:
|
except asyncio.CancelledError:
|
||||||
raise
|
raise
|
||||||
except:
|
except:
|
||||||
self._logger.error(traceback.format_exc())
|
self._logger.error(traceback.format_exc())
|
||||||
continue
|
continue
|
||||||
event_data = RoomEventData(
|
|
||||||
room=room,
|
|
||||||
event=event,
|
|
||||||
filter=callback_info.filter,
|
|
||||||
bot=self._matrix_bot
|
|
||||||
)
|
|
||||||
try:
|
try:
|
||||||
# dump argument types
|
# dump argument types
|
||||||
if callback_info.callback is None:
|
if callback_info.callback is None:
|
||||||
@@ -144,7 +144,7 @@ class Callbacks:
|
|||||||
def add_room_event_callback(
|
def add_room_event_callback(
|
||||||
self,
|
self,
|
||||||
filter: BaseEventFilter,
|
filter: BaseEventFilter,
|
||||||
callback: Callable[[RoomEventData], Coroutine[Any, Any, None]] | None,
|
callback: Callable[[EventContext], Coroutine[Any, Any, None]] | None,
|
||||||
*,
|
*,
|
||||||
stop_matching: bool = True) -> None:
|
stop_matching: bool = True) -> None:
|
||||||
"""
|
"""
|
||||||
|
|||||||
@@ -3,10 +3,11 @@ import logging
|
|||||||
|
|
||||||
from typing import Callable, Coroutine, Any
|
from typing import Callable, Coroutine, Any
|
||||||
|
|
||||||
from nio import AsyncClient
|
from nio import AsyncClient, MatrixRoom
|
||||||
|
|
||||||
from ..filters.base import BaseEventFilter
|
from ..filters.base import BaseEventFilter
|
||||||
from ..types import *
|
from ..types import *
|
||||||
|
from ..context import EventContext
|
||||||
|
|
||||||
from ._validation import Validator
|
from ._validation import Validator
|
||||||
from ._storage import Storage
|
from ._storage import Storage
|
||||||
@@ -44,7 +45,7 @@ class MatrixBot:
|
|||||||
|
|
||||||
def add_callback(self,
|
def add_callback(self,
|
||||||
filter: BaseEventFilter,
|
filter: BaseEventFilter,
|
||||||
callback: Callable[[RoomEventData], Coroutine[Any, Any, None]] | None,
|
callback: Callable[[EventContext], Coroutine[Any, Any, None]] | None,
|
||||||
*,
|
*,
|
||||||
stop_matching: bool = True) -> None:
|
stop_matching: bool = True) -> None:
|
||||||
"""
|
"""
|
||||||
|
|||||||
76
src/mab/context.py
Normal file
76
src/mab/context.py
Normal file
@@ -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
|
||||||
@@ -5,6 +5,8 @@ from typing import Any, Type
|
|||||||
from nio import AsyncClient
|
from nio import AsyncClient
|
||||||
from nio import MatrixRoom, Event
|
from nio import MatrixRoom, Event
|
||||||
|
|
||||||
|
from ..context import EventContext
|
||||||
|
|
||||||
class BaseEventFilter(ABC):
|
class BaseEventFilter(ABC):
|
||||||
"""Base class for all message filters"""
|
"""Base class for all message filters"""
|
||||||
_logger = logging.Logger("EventFilter")
|
_logger = logging.Logger("EventFilter")
|
||||||
@@ -64,7 +66,7 @@ class BaseEventFilter(ABC):
|
|||||||
"""
|
"""
|
||||||
return str(self.__class__.__name__)
|
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
|
This abstract method must be redefined in derived classes so that the
|
||||||
filter operates according to its description. This method must not raise
|
filter operates according to its description. This method must not raise
|
||||||
@@ -72,9 +74,9 @@ class BaseEventFilter(ABC):
|
|||||||
and return False
|
and return False
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
- room - room the event has happened in
|
- context - event context; your derived classes may add variables
|
||||||
- event - the event to check againts this filter
|
to it (see `message.MessageTypeFilter` implementation
|
||||||
- client - the client
|
for reference)
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
- True if the event satisfies this filter
|
- True if the event satisfies this filter
|
||||||
@@ -88,8 +90,10 @@ class EventTypeFilter(BaseEventFilter):
|
|||||||
super().__init__(**kwargs)
|
super().__init__(**kwargs)
|
||||||
self._type = type
|
self._type = type
|
||||||
|
|
||||||
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
|
async def __call__(self, context: EventContext) -> bool:
|
||||||
return isinstance(event, self._type)
|
if not await super().__call__(context):
|
||||||
|
return False
|
||||||
|
return isinstance(context.event, self._type)
|
||||||
|
|
||||||
class CompoundEventFilter(BaseEventFilter):
|
class CompoundEventFilter(BaseEventFilter):
|
||||||
"""Event filter that consists of multiple filters"""
|
"""Event filter that consists of multiple filters"""
|
||||||
@@ -140,8 +144,10 @@ class CompoundEventFilter(BaseEventFilter):
|
|||||||
expression = f"~{reprs[0]}"
|
expression = f"~{reprs[0]}"
|
||||||
return f"({expression})"
|
return f"({expression})"
|
||||||
|
|
||||||
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
|
async def __call__(self, context: EventContext) -> bool:
|
||||||
evaluated = [await arg(room, event, client) for arg in self._arguments]
|
if not await super().__call__(context):
|
||||||
|
return False
|
||||||
|
evaluated = [await arg(context) for arg in self._arguments]
|
||||||
if self._operator == self.OPERATOR_AND:
|
if self._operator == self.OPERATOR_AND:
|
||||||
return all(evaluated)
|
return all(evaluated)
|
||||||
elif self._operator == self.OPERATOR_OR:
|
elif self._operator == self.OPERATOR_OR:
|
||||||
|
|||||||
@@ -1,6 +1,13 @@
|
|||||||
import re
|
import re
|
||||||
import traceback
|
import traceback
|
||||||
from .message import NewMessageFilter
|
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 AsyncClient
|
||||||
from nio import MatrixRoom, Event
|
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,
|
This filter will match any message that has `body` in it, including images,
|
||||||
videos, files, etc.
|
videos, files, etc.
|
||||||
|
|
||||||
|
This filter sets `CTX_BODY` context variable.
|
||||||
"""
|
"""
|
||||||
def __init__(self, *, ignore_filename_in_body: bool = True, **kwargs):
|
def __init__(self, *, ignore_filename_in_body: bool = True, **kwargs):
|
||||||
super().__init__(**kwargs)
|
super().__init__(**kwargs)
|
||||||
self._ignore_filename_in_body = ignore_filename_in_body
|
self._ignore_filename_in_body = ignore_filename_in_body
|
||||||
|
|
||||||
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
|
async def __call__(self, context: EventContext) -> bool:
|
||||||
if not await super().__call__(room, event, client):
|
if not await super().__call__(context):
|
||||||
return False
|
return False
|
||||||
if not hasattr(event, "body"):
|
if not hasattr(context.event, "body"):
|
||||||
return False
|
return False
|
||||||
if not isinstance(event.body, str): # type: ignore
|
if not isinstance(context.event.body, str): # type: ignore
|
||||||
return False
|
return False
|
||||||
if not event.body.strip(): # type: ignore
|
if not context.event.body.strip(): # type: ignore
|
||||||
return False
|
return False
|
||||||
if self._ignore_filename_in_body:
|
if self._ignore_filename_in_body:
|
||||||
content = event.source["content"]
|
content = context.event.source["content"]
|
||||||
if "filename" in content and content["filename"] == event.body: # type: ignore
|
if "filename" in content and content["filename"] == context.event.body: # type: ignore
|
||||||
return False
|
return False
|
||||||
|
context[CTX_BODY] = context.event.body # type: ignore
|
||||||
return True
|
return True
|
||||||
|
|
||||||
class BodyContainsFilter(BodyExistsFilter):
|
class BodyContainsFilter(BodyExistsFilter):
|
||||||
@@ -61,10 +71,10 @@ class BodyContainsFilter(BodyExistsFilter):
|
|||||||
self._any_case = any_case
|
self._any_case = any_case
|
||||||
self._needle = needle
|
self._needle = needle
|
||||||
|
|
||||||
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
|
async def __call__(self, context: EventContext) -> bool:
|
||||||
if not await super().__call__(room, event, client):
|
if not await super().__call__(context):
|
||||||
return False
|
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:
|
for n in self._needle:
|
||||||
if n in body:
|
if n in body:
|
||||||
return True
|
return True
|
||||||
@@ -90,10 +100,10 @@ class BodyStartsWithFilter(BodyExistsFilter):
|
|||||||
self._any_case = any_case
|
self._any_case = any_case
|
||||||
self._substring = substring
|
self._substring = substring
|
||||||
|
|
||||||
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
|
async def __call__(self, context: EventContext) -> bool:
|
||||||
if not await super().__call__(room, event, client):
|
if not await super().__call__(context):
|
||||||
return False
|
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:
|
for s in self._substring:
|
||||||
if body.startswith(s):
|
if body.startswith(s):
|
||||||
return True
|
return True
|
||||||
@@ -119,10 +129,10 @@ class BodyEndsWithFilter(BodyExistsFilter):
|
|||||||
self._any_case = any_case
|
self._any_case = any_case
|
||||||
self._substring = substring
|
self._substring = substring
|
||||||
|
|
||||||
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
|
async def __call__(self, context: EventContext) -> bool:
|
||||||
if not await super().__call__(room, event, client):
|
if not await super().__call__(context):
|
||||||
return False
|
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:
|
for s in self._substring:
|
||||||
if body.endswith(s):
|
if body.endswith(s):
|
||||||
return True
|
return True
|
||||||
@@ -145,9 +155,10 @@ class BodyCommandFilter(BodyExistsFilter):
|
|||||||
store all verbs in lower case. This filter will not match any verbs that
|
store all verbs in lower case. This filter will not match any verbs that
|
||||||
use mixed case of upper case.
|
use mixed case of upper case.
|
||||||
|
|
||||||
If this filter is matched, then it will set a new attribute for the event:
|
This filter sets the following context variables:
|
||||||
`event.command_args: list[str]`. You may use this attribute in your callback
|
- `CTX_CMD_PREFIX` - prefix that was used
|
||||||
for this event.
|
- `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):
|
def __init__(self, verbs: str | list[str], min_args: int = 0, max_args: int | None = None, prefix: str = "!", **kwargs):
|
||||||
super().__init__(**kwargs)
|
super().__init__(**kwargs)
|
||||||
@@ -158,10 +169,10 @@ class BodyCommandFilter(BodyExistsFilter):
|
|||||||
self._max_args = max_args
|
self._max_args = max_args
|
||||||
self._prefix = prefix
|
self._prefix = prefix
|
||||||
|
|
||||||
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
|
async def __call__(self, context: EventContext) -> bool:
|
||||||
if not await super().__call__(room, event, client):
|
if not await super().__call__(context):
|
||||||
return False
|
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
|
args_count = len(parts) - 1
|
||||||
if args_count < self._min_args:
|
if args_count < self._min_args:
|
||||||
return False
|
return False
|
||||||
@@ -172,7 +183,9 @@ class BodyCommandFilter(BodyExistsFilter):
|
|||||||
cmd = parts[0][len(self._prefix):].lower()
|
cmd = parts[0][len(self._prefix):].lower()
|
||||||
for verb in self._verbs:
|
for verb in self._verbs:
|
||||||
if cmd == verb:
|
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 True
|
||||||
return False
|
return False
|
||||||
|
|
||||||
@@ -186,11 +199,11 @@ class BodyRegexFilter(BodyExistsFilter):
|
|||||||
regex = re.compile(regex)
|
regex = re.compile(regex)
|
||||||
self._regex = regex
|
self._regex = regex
|
||||||
|
|
||||||
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
|
async def __call__(self, context: EventContext) -> bool:
|
||||||
if not await super().__call__(room, event, client):
|
if not await super().__call__(context):
|
||||||
return False
|
return False
|
||||||
try:
|
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:
|
except:
|
||||||
self._logger.error(traceback.format_exc())
|
self._logger.error(traceback.format_exc())
|
||||||
return False
|
return False
|
||||||
@@ -1,6 +1,7 @@
|
|||||||
import traceback
|
import traceback
|
||||||
from .base import BaseEventFilter, EventTypeFilter
|
from .base import BaseEventFilter, EventTypeFilter
|
||||||
from ..types import MessageType
|
from ..types import MessageType
|
||||||
|
from ..context import EventContext, CTX_MESSAGE_TYPE, CTX_SENDER
|
||||||
|
|
||||||
from nio import AsyncClient
|
from nio import AsyncClient
|
||||||
from nio import MatrixRoom, Event
|
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
|
`types` list is stored by reference so you may modify the behavior of this
|
||||||
filter dynamically.
|
filter dynamically.
|
||||||
|
|
||||||
|
This filter sets `CTX_MESSAGE_TYPE` variable in the context.
|
||||||
"""
|
"""
|
||||||
def __init__(self, types: list[MessageType] | MessageType, **kwargs):
|
def __init__(self, types: list[MessageType] | MessageType, **kwargs):
|
||||||
super().__init__(**kwargs)
|
super().__init__(**kwargs)
|
||||||
@@ -22,14 +25,16 @@ class MessageTypeFilter(BaseEventFilter):
|
|||||||
types = [types]
|
types = [types]
|
||||||
self._types = types
|
self._types = types
|
||||||
|
|
||||||
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
|
async def __call__(self, context: EventContext) -> bool:
|
||||||
if not await super().__call__(room, event, client):
|
if not await super().__call__(context):
|
||||||
return False
|
return False
|
||||||
if "msgtype" not in event.source["content"]:
|
if "msgtype" not in context.event.source["content"]:
|
||||||
return False
|
return False
|
||||||
return (
|
msgtype = context.event.source["content"]["msgtype"]
|
||||||
event.source["content"]["msgtype"] in [t.value for t in self._types]
|
if not msgtype in [t.value for t in self._types]:
|
||||||
)
|
return False
|
||||||
|
context[CTX_MESSAGE_TYPE] = msgtype
|
||||||
|
return True
|
||||||
|
|
||||||
class NewMessageFilter(BaseEventFilter):
|
class NewMessageFilter(BaseEventFilter):
|
||||||
"""
|
"""
|
||||||
@@ -40,10 +45,10 @@ class NewMessageFilter(BaseEventFilter):
|
|||||||
def __init__(self, **kwargs):
|
def __init__(self, **kwargs):
|
||||||
super().__init__(**kwargs)
|
super().__init__(**kwargs)
|
||||||
|
|
||||||
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
|
async def __call__(self, context: EventContext) -> bool:
|
||||||
if not await super().__call__(room, event, client):
|
if not await super().__call__(context):
|
||||||
return False
|
return False
|
||||||
return "m.new_content" not in event.source["content"]
|
return "m.new_content" not in context.event.source["content"]
|
||||||
|
|
||||||
class EditedMessageFilter(BaseEventFilter):
|
class EditedMessageFilter(BaseEventFilter):
|
||||||
"""
|
"""
|
||||||
@@ -53,10 +58,10 @@ class EditedMessageFilter(BaseEventFilter):
|
|||||||
def __init__(self, **kwargs):
|
def __init__(self, **kwargs):
|
||||||
super().__init__(**kwargs)
|
super().__init__(**kwargs)
|
||||||
|
|
||||||
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
|
async def __call__(self, context: EventContext) -> bool:
|
||||||
if not await super().__call__(room, event, client):
|
if not await super().__call__(context):
|
||||||
return False
|
return False
|
||||||
return "m.new_content" in event.source["content"]
|
return "m.new_content" in context.event.source["content"]
|
||||||
|
|
||||||
class RedactedMessageFilter(EventTypeFilter):
|
class RedactedMessageFilter(EventTypeFilter):
|
||||||
"""
|
"""
|
||||||
@@ -64,9 +69,10 @@ class RedactedMessageFilter(EventTypeFilter):
|
|||||||
"""
|
"""
|
||||||
def __init__(self, **kwargs):
|
def __init__(self, **kwargs):
|
||||||
super().__init__(RedactionEvent, **kwargs)
|
super().__init__(RedactionEvent, **kwargs)
|
||||||
|
raise NotImplementedError()
|
||||||
|
|
||||||
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
|
async def __call__(self, context: EventContext) -> bool:
|
||||||
return await super().__call__(room, event, client)
|
raise NotImplementedError()
|
||||||
|
|
||||||
class SenderIsFilter(BaseEventFilter):
|
class SenderIsFilter(BaseEventFilter):
|
||||||
"""
|
"""
|
||||||
@@ -77,6 +83,8 @@ class SenderIsFilter(BaseEventFilter):
|
|||||||
|
|
||||||
`senders` list is stored by reference so you can modify behavior of this
|
`senders` list is stored by reference so you can modify behavior of this
|
||||||
filter dynamically.
|
filter dynamically.
|
||||||
|
|
||||||
|
This filter sets `CTX_SENDER` variable in the context.
|
||||||
"""
|
"""
|
||||||
def __init__(self, sender: list[str] | str, *, any_case: bool = True, **kwargs):
|
def __init__(self, sender: list[str] | str, *, any_case: bool = True, **kwargs):
|
||||||
super().__init__(**kwargs)
|
super().__init__(**kwargs)
|
||||||
@@ -85,12 +93,13 @@ class SenderIsFilter(BaseEventFilter):
|
|||||||
self._sender = sender
|
self._sender = sender
|
||||||
self._any_case = any_case
|
self._any_case = any_case
|
||||||
|
|
||||||
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
|
async def __call__(self, context: EventContext) -> bool:
|
||||||
if not await super().__call__(room, event, client):
|
if not await super().__call__(context):
|
||||||
return False
|
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:
|
for s in self._sender:
|
||||||
if sender == s:
|
if sender == s:
|
||||||
|
context[CTX_SENDER] = sender
|
||||||
return True
|
return True
|
||||||
return False
|
return False
|
||||||
|
|
||||||
@@ -106,7 +115,7 @@ class SenderIsBotFilter(BaseEventFilter):
|
|||||||
def __init__(self, **kwargs):
|
def __init__(self, **kwargs):
|
||||||
super().__init__(**kwargs)
|
super().__init__(**kwargs)
|
||||||
|
|
||||||
async def __call__(self, room: MatrixRoom, event: Event, client: AsyncClient) -> bool:
|
async def __call__(self, context: EventContext) -> bool:
|
||||||
if not await super().__call__(room, event, client):
|
if not await super().__call__(context):
|
||||||
return False
|
return False
|
||||||
return client.user_id == event.sender
|
return context.bot.get_client().user_id == context.event.sender
|
||||||
@@ -4,15 +4,8 @@ from pathlib import Path
|
|||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from enum import Enum
|
from enum import Enum
|
||||||
|
|
||||||
from nio import MatrixRoom, Event
|
|
||||||
from nio import UploadResponse
|
from nio import UploadResponse
|
||||||
|
|
||||||
from .filters.base import BaseEventFilter
|
|
||||||
|
|
||||||
from typing import TYPE_CHECKING
|
|
||||||
if TYPE_CHECKING:
|
|
||||||
from .bot import MatrixBot
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class MatrixBotConfig:
|
class MatrixBotConfig:
|
||||||
"""Configuration for MatrixBot"""
|
"""Configuration for MatrixBot"""
|
||||||
@@ -68,22 +61,6 @@ class VideoFileProperties:
|
|||||||
thumbnail: Path | str | bytes | None = None
|
thumbnail: Path | str | bytes | None = None
|
||||||
"""Path to the thumbnail or the raw JPEG thumbnail data"""
|
"""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
|
@dataclass
|
||||||
class UploadResult:
|
class UploadResult:
|
||||||
"""Result of data upload"""
|
"""Result of data upload"""
|
||||||
@@ -109,3 +86,20 @@ class MessageType(Enum):
|
|||||||
AUDIO = "m.audio"
|
AUDIO = "m.audio"
|
||||||
LOCATION = "m.location"
|
LOCATION = "m.location"
|
||||||
VIDEO = "m.video"
|
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)})"
|
||||||
Reference in New Issue
Block a user