Add filter and middleware

This commit is contained in:
Zakhar Izmaylov
2025-01-19 08:29:10 +03:00
parent 588c32453d
commit d83da5daaa
3 changed files with 68 additions and 10 deletions
+7 -2
View File
@@ -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):
+8
View File
@@ -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
+53 -8
View File
@@ -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)