From d83da5daaa9717629b9921402505744e3c56dbfa Mon Sep 17 00:00:00 2001 From: Zakhar Izmaylov Date: Sun, 19 Jan 2025 08:29:10 +0300 Subject: [PATCH] Add filter and middleware --- bot.py | 9 ++++-- filters/private.py | 8 +++++ middlewares/throttling.py | 61 ++++++++++++++++++++++++++++++++++----- 3 files changed, 68 insertions(+), 10 deletions(-) create mode 100644 filters/private.py diff --git a/bot.py b/bot.py index 715fb7b2..7a708d80 100644 --- a/bot.py +++ b/bot.py @@ -13,6 +13,8 @@ from middlewares.database import DatabaseMiddleware from middlewares.delete import DeleteMessageMiddleware from middlewares.logging import LoggingMiddleware from middlewares.user import UserMiddleware +from middlewares.throttling import ThrottlingMiddleware +from filters.private import IsPrivate bot = Bot(token=API_TOKEN, default=DefaultBotProperties(parse_mode=ParseMode.HTML)) storage = MemoryStorage() @@ -29,12 +31,15 @@ dp.callback_query.middleware(UserMiddleware()) dp.message.middleware(DatabaseMiddleware()) dp.callback_query.middleware(DatabaseMiddleware()) -# dp.message.middleware(ThrottlingMiddleware(limit=1)) -# dp.callback_query.middleware(ThrottlingMiddleware(limit=1)) +dp.message.middleware(ThrottlingMiddleware()) +dp.callback_query.middleware(ThrottlingMiddleware()) dp.message.outer_middleware(DeleteMessageMiddleware()) dp.callback_query.outer_middleware(DeleteMessageMiddleware()) +dp.message.filter(IsPrivate()) +dp.callback_query.filter(IsPrivate()) + @dp.error() async def error_handler(event: ErrorEvent): diff --git a/filters/private.py b/filters/private.py new file mode 100644 index 00000000..988db95c --- /dev/null +++ b/filters/private.py @@ -0,0 +1,8 @@ +from aiogram.enums import ChatType +from aiogram.filters import BaseFilter +from aiogram.types import Chat, TelegramObject + + +class IsPrivate(BaseFilter): + async def __call__(self, event: TelegramObject, event_chat: Chat) -> bool: + return event_chat.type == ChatType.PRIVATE \ No newline at end of file diff --git a/middlewares/throttling.py b/middlewares/throttling.py index 2761920c..af753651 100644 --- a/middlewares/throttling.py +++ b/middlewares/throttling.py @@ -1,13 +1,29 @@ -from collections.abc import Awaitable, Callable -from typing import Any +from typing import Any, Awaitable, Callable, MutableMapping from aiogram import BaseMiddleware -from aiogram.types import TelegramObject +from aiogram.dispatcher.flags import get_flag +from aiogram.types import TelegramObject, Update, User +from cachetools import TTLCache +from logger import logger -class ThrottleMiddleware(BaseMiddleware): - def __init__(self, limit: int): - self.limit = limit +class ThrottlingMiddleware(BaseMiddleware): + def __init__( + self, + *, + default_key: str | None = "default", + default_ttl: float = 0.5, + **ttl_map: float, + ) -> None: + if default_key: + ttl_map[default_key] = default_ttl + + self.default_key = default_key + self.caches: dict[str, MutableMapping[int, None]] = {} + + for name, ttl in ttl_map.items(): + self.caches[name] = TTLCache(maxsize=10_000, ttl=ttl) + logger.debug("ThrottlingMiddleware initialized.") async def __call__( self, @@ -15,5 +31,34 @@ class ThrottleMiddleware(BaseMiddleware): event: TelegramObject, data: dict[str, Any], ) -> Any: - #todo - return await handler(event, data) + if not isinstance(event, Update): + logger.debug(f"Received event of type {type(event)}, skipping throttling.") + return await handler(event, data) + + if event.pre_checkout_query: + logger.debug("Pre-checkout query event, skipping throttling.") + return await handler(event, data) + + if event.message and event.message.successful_payment: + logger.debug("Successful payment event, skipping throttling.") + return await handler(event, data) + + user: User | None = data.get("event_from_user", None) + + if user is not None: + key = get_flag(data, "throttling_key", default=self.default_key) + + if key: + if user.id in self.caches[key]: + logger.warning(f"User {user.id} is being throttled with key: {key}") + return None + logger.debug( + f"User {user.id} is allowed to proceed, adding to cache with key: {key}", + ) + self.caches[key][user.id] = None + else: + logger.debug( + f"No throttling key provided for user {user.id}, proceeding without throttle." + ) + + return await handler(event, data) \ No newline at end of file