diff --git a/src/aiogram_input/__init__.py b/src/aiogram_input/__init__.py index e333a9f..67993be 100644 --- a/src/aiogram_input/__init__.py +++ b/src/aiogram_input/__init__.py @@ -1,5 +1,6 @@ from importlib.metadata import PackageNotFoundError, version +from .middleware import DEFAULT_DATA_KEY from .setup import setup_input from .storage import InputStorage, MemoryInputStorage from .waiter import InputWaiter @@ -10,6 +11,7 @@ __version__ = "0.0.0" __all__ = ( + "DEFAULT_DATA_KEY", "InputStorage", "InputWaiter", "MemoryInputStorage", diff --git a/src/aiogram_input/middleware.py b/src/aiogram_input/middleware.py index b09758d..ec66322 100644 --- a/src/aiogram_input/middleware.py +++ b/src/aiogram_input/middleware.py @@ -8,7 +8,7 @@ from .session import SessionManager from .waiter import InputWaiter -HANDLER_DATA_KEY = "input" +DEFAULT_DATA_KEY = "input" class InputMiddleware(BaseMiddleware): @@ -18,9 +18,18 @@ class InputMiddleware(BaseMiddleware): Register only on Dispatcher so every router shares one waiter instance. """ - def __init__(self, session: SessionManager, waiter: InputWaiter) -> None: + def __init__( + self, + session: SessionManager, + waiter: InputWaiter, + *, + data_key: str = DEFAULT_DATA_KEY, + ) -> None: + if not data_key: + raise ValueError("data_key must be a non-empty string") self._session = session self._waiter = waiter + self._data_key = data_key async def __call__( self, @@ -28,7 +37,7 @@ async def __call__( event: TelegramObject, data: Dict[str, Any], ) -> Any: - data[HANDLER_DATA_KEY] = self._waiter + data[self._data_key] = self._waiter if isinstance(event, Message): if await self._session.feed(event): return None diff --git a/src/aiogram_input/setup.py b/src/aiogram_input/setup.py index 6158aec..e24cf3b 100644 --- a/src/aiogram_input/setup.py +++ b/src/aiogram_input/setup.py @@ -4,7 +4,7 @@ from aiogram import Dispatcher -from .middleware import InputMiddleware +from .middleware import DEFAULT_DATA_KEY, InputMiddleware from .registry import WaitRegistry from .session import SessionManager from .storage import InputStorage, MemoryInputStorage @@ -16,13 +16,14 @@ def setup_input( /, *, storage: Optional[InputStorage] = None, + data_key: str = DEFAULT_DATA_KEY, ) -> InputWaiter: """ Register input waiting once on a Dispatcher. - Injects ``InputWaiter`` into handler data as ``input`` and consumes matching - messages only while a wait is active. Other messages pass through to FSM - and normal handlers. + Injects ``InputWaiter`` into handler data under ``data_key`` (default: + ``"input"``) and consumes matching messages only while a wait is active. + Other messages pass through to FSM and normal handlers. """ if not isinstance(dispatcher, Dispatcher): raise TypeError( @@ -33,5 +34,5 @@ def setup_input( registry = WaitRegistry() session = SessionManager(storage, registry) waiter = InputWaiter(session) - InputMiddleware(session, waiter).setup(dispatcher) + InputMiddleware(session, waiter, data_key=data_key).setup(dispatcher) return waiter diff --git a/tests/test_filters_and_pass_through.py b/tests/test_filters_and_pass_through.py index a2bd05a..fa34cf2 100644 --- a/tests/test_filters_and_pass_through.py +++ b/tests/test_filters_and_pass_through.py @@ -8,7 +8,7 @@ from aiogram.types import CallbackQuery from aiogram_input import InputWaiter, setup_input -from aiogram_input.middleware import HANDLER_DATA_KEY, InputMiddleware +from aiogram_input.middleware import DEFAULT_DATA_KEY, InputMiddleware from aiogram_input.session import SessionManager from .helpers import make_message @@ -64,7 +64,7 @@ async def test_non_message_event_passes_through_with_injection() -> None: async def handler(evt, data): seen["event"] = evt - seen["input"] = data.get(HANDLER_DATA_KEY) + seen["input"] = data.get(DEFAULT_DATA_KEY) return "cb" assert await middleware(handler, event, {}) == "cb" @@ -83,7 +83,7 @@ async def test_unmatched_message_reaches_handler_fsm_coexistence() -> None: async def handler(event, data): nonlocal reached reached = True - assert data[HANDLER_DATA_KEY] is waiter + assert data[DEFAULT_DATA_KEY] is waiter return "handled" result = await middleware(handler, make_message(999), {}) diff --git a/tests/test_middleware_di.py b/tests/test_middleware_di.py index e1586c1..7e498f6 100644 --- a/tests/test_middleware_di.py +++ b/tests/test_middleware_di.py @@ -8,7 +8,7 @@ from aiogram.types import Message from aiogram_input import InputWaiter, MemoryInputStorage, setup_input -from aiogram_input.middleware import HANDLER_DATA_KEY, InputMiddleware +from aiogram_input.middleware import DEFAULT_DATA_KEY, InputMiddleware def _message(chat_id: int) -> Message: @@ -37,7 +37,7 @@ async def test_middleware_injects_waiter_and_passes_through() -> None: async def handler(event, data): nonlocal called called = True - assert data[HANDLER_DATA_KEY] is waiter + assert data[DEFAULT_DATA_KEY] is waiter return "ok" result = await middleware(handler, _message(1), {}) @@ -45,6 +45,28 @@ async def handler(event, data): assert called is True +@pytest.mark.asyncio +async def test_custom_data_key_avoids_name_collision() -> None: + dp = Dispatcher() + waiter = setup_input(dp, data_key="aiogram_input") + middleware = InputMiddleware(waiter._session, waiter, data_key="aiogram_input") + + async def handler(event, data): + assert data["input"] == "user-dep" + assert data["aiogram_input"] is waiter + return "ok" + + assert await middleware(handler, _message(1), {"input": "user-dep"}) == "ok" + + +@pytest.mark.asyncio +async def test_empty_data_key_rejected() -> None: + dp = Dispatcher() + waiter = setup_input(dp) + with pytest.raises(ValueError, match="data_key"): + InputMiddleware(waiter._session, waiter, data_key="") + + @pytest.mark.asyncio async def test_middleware_consumes_matching_wait() -> None: dp = Dispatcher()