Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions src/aiogram_input/__init__.py
Original file line number Diff line number Diff line change
@@ -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
Expand All @@ -10,6 +11,7 @@
__version__ = "0.0.0"

__all__ = (
"DEFAULT_DATA_KEY",
"InputStorage",
"InputWaiter",
"MemoryInputStorage",
Expand Down
15 changes: 12 additions & 3 deletions src/aiogram_input/middleware.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,7 @@
from .session import SessionManager
from .waiter import InputWaiter

HANDLER_DATA_KEY = "input"
DEFAULT_DATA_KEY = "input"

Copy link
Copy Markdown
Owner Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

critical name change



class InputMiddleware(BaseMiddleware):
Expand All @@ -18,17 +18,26 @@ 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,
handler: Callable[[TelegramObject, Dict[str, Any]], Awaitable[Any]],
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
Expand Down
11 changes: 6 additions & 5 deletions src/aiogram_input/setup.py

Copy link
Copy Markdown
Owner Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

must have a key validitaion

Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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(
Expand All @@ -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
6 changes: 3 additions & 3 deletions tests/test_filters_and_pass_through.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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"
Expand All @@ -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), {})
Expand Down
26 changes: 24 additions & 2 deletions tests/test_middleware_di.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -37,14 +37,36 @@ 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), {})
assert result == "ok"
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()
Expand Down
Loading