diff --git a/src/mab/context.py b/src/mab/context.py index 9ec178f..b424179 100644 --- a/src/mab/context.py +++ b/src/mab/context.py @@ -3,7 +3,7 @@ from dataclasses import dataclass from typing import Any, TYPE_CHECKING -from .types import ContextDataKey +from .types import ContextDataKey, MessageType from nio import MatrixRoom, Event @@ -16,7 +16,7 @@ if TYPE_CHECKING: CTX_BODY = ContextDataKey[str]("CTX_BODY") """Value of `event.body`""" -CTX_MESSAGE_TYPE = ContextDataKey[str]("CTX_MESSAGE_TYPE") +CTX_MESSAGE_TYPE = ContextDataKey[MessageType]("CTX_MESSAGE_TYPE") """Value of `msgtype` for the event""" CTX_SENDER = ContextDataKey[str]("CTX_SENDER") diff --git a/src/mab/filters/message.py b/src/mab/filters/message.py index 8b875a9..a200242 100644 --- a/src/mab/filters/message.py +++ b/src/mab/filters/message.py @@ -33,7 +33,7 @@ class MessageTypeFilter(BaseEventFilter): msgtype = context.event.source["content"]["msgtype"] if not msgtype in [t.value for t in self._types]: return False - context[CTX_MESSAGE_TYPE] = msgtype + context[CTX_MESSAGE_TYPE] = MessageType(msgtype) return True class NewMessageFilter(BaseEventFilter):