diff --git a/backend/bot/app/controllers/dispatcher_controller.py b/backend/bot/app/controllers/dispatcher_controller.py index 14deecf..97c112c 100644 --- a/backend/bot/app/controllers/dispatcher_controller.py +++ b/backend/bot/app/controllers/dispatcher_controller.py @@ -17,6 +17,7 @@ from bot.middlewares.channel_subscription import ChannelSubscriptionMiddleware from bot.middlewares.db_session import DBSessionMiddleware from bot.middlewares.i18n import I18nMiddleware, get_i18n_instance from bot.middlewares.profile_sync import ProfileSyncMiddleware +from bot.middlewares.update_antiflood import UpdateAntiFloodMiddleware from config.settings import Settings @@ -38,6 +39,7 @@ def build_dispatcher( dp["i18n_instance"] = i18n_instance dp["async_session_factory"] = async_session_factory + dp.update.outer_middleware(UpdateAntiFloodMiddleware(settings=settings)) dp.update.outer_middleware(DBSessionMiddleware(async_session_factory)) dp.update.outer_middleware(I18nMiddleware(i18n=i18n_instance, settings=settings)) dp.update.outer_middleware(ProfileSyncMiddleware()) diff --git a/backend/bot/middlewares/update_antiflood.py b/backend/bot/middlewares/update_antiflood.py new file mode 100644 index 0000000..08e763d --- /dev/null +++ b/backend/bot/middlewares/update_antiflood.py @@ -0,0 +1,134 @@ +import asyncio +import logging +import time +from collections import defaultdict, deque +from dataclasses import dataclass +from typing import Any, Awaitable, Callable, Deque, Dict, Optional + +from aiogram import BaseMiddleware +from aiogram.types import Update + +from bot.infra.redis import get_redis, redis_key +from config.settings import Settings + +logger = logging.getLogger(__name__) + +DEFAULT_WINDOW_SECONDS = 60 +DEFAULT_MAX_UPDATES_PER_WINDOW = 180 + + +@dataclass(frozen=True) +class RateLimitRule: + window_seconds: int + max_events: int + + +class UpdateAntiFloodMiddleware(BaseMiddleware): + """Drop extreme update floods before DB-backed middleware runs.""" + + def __init__( + self, + settings: Settings, + *, + default_rule: Optional[RateLimitRule] = None, + ) -> None: + super().__init__() + self.settings = settings + self.default_rule = default_rule or RateLimitRule( + window_seconds=int( + getattr(settings, "TELEGRAM_ANTIFLOOD_WINDOW_SECONDS", DEFAULT_WINDOW_SECONDS) + or DEFAULT_WINDOW_SECONDS + ), + max_events=int( + getattr( + settings, + "TELEGRAM_ANTIFLOOD_MAX_UPDATES_PER_WINDOW", + DEFAULT_MAX_UPDATES_PER_WINDOW, + ) + or DEFAULT_MAX_UPDATES_PER_WINDOW + ), + ) + self._local_buckets: Dict[str, Deque[float]] = defaultdict(deque) + self._local_lock = asyncio.Lock() + + async def __call__( + self, + handler: Callable[[Update, Dict[str, Any]], Awaitable[Any]], + event: Update, + data: Dict[str, Any], + ) -> Any: + if not bool(getattr(self.settings, "TELEGRAM_ANTIFLOOD_ENABLED", True)): + return await handler(event, data) + + actor_key = _update_actor_key(event) + if not actor_key: + return await handler(event, data) + + if await self._is_limited(actor_key, self.default_rule): + logger.warning( + "Telegram update dropped by anti-flood: actor=%s update_type=%s", + actor_key, + getattr(event, "event_type", "unknown"), + ) + data["antiflood_dropped"] = True + return None + + return await handler(event, data) + + async def _is_limited(self, actor_key: str, rule: RateLimitRule) -> bool: + if rule.window_seconds <= 0 or rule.max_events <= 0: + return False + + try: + redis = await get_redis(self.settings) + if redis is not None: + key = redis_key( + self.settings, + "rate-limit", + "telegram", + "updates", + actor_key, + ) + current = int(await redis.incr(key)) + if current == 1: + await redis.expire(key, rule.window_seconds) + return current > rule.max_events + except Exception as exc: + logger.warning("Redis telegram anti-flood unavailable; using local fallback: %s", exc) + + return await self._is_limited_local(actor_key, rule) + + async def _is_limited_local(self, actor_key: str, rule: RateLimitRule) -> bool: + now = time.monotonic() + cutoff = now - rule.window_seconds + async with self._local_lock: + bucket = self._local_buckets[actor_key] + while bucket and bucket[0] <= cutoff: + bucket.popleft() + bucket.append(now) + if len(bucket) > rule.max_events: + return True + if not bucket: + self._local_buckets.pop(actor_key, None) + return False + + +def _update_actor_key(update: Update) -> Optional[str]: + user_id = None + chat_id = None + + if update.message: + user_id = update.message.from_user.id if update.message.from_user else None + chat_id = update.message.chat.id if update.message.chat else None + elif update.callback_query: + user_id = update.callback_query.from_user.id if update.callback_query.from_user else None + if update.callback_query.message and update.callback_query.message.chat: + chat_id = update.callback_query.message.chat.id + elif update.inline_query: + user_id = update.inline_query.from_user.id if update.inline_query.from_user else None + + if user_id is not None: + return f"user:{int(user_id)}" + if chat_id is not None: + return f"chat:{int(chat_id)}" + return None diff --git a/backend/config/settings.py b/backend/config/settings.py index 43795a3..b31c441 100644 --- a/backend/config/settings.py +++ b/backend/config/settings.py @@ -220,6 +220,9 @@ class Settings(BaseSettings): PANEL_SYNC_LIFETIME_TRAFFIC_MIN_DELTA_BYTES: int = Field(default=104857600) WEBAPP_RATE_LIMIT_TTL_SECONDS: int = Field(default=60) WEBAPP_RATE_LIMIT_MAX_REQUESTS: int = Field(default=30) + TELEGRAM_ANTIFLOOD_ENABLED: bool = Field(default=True) + TELEGRAM_ANTIFLOOD_WINDOW_SECONDS: int = Field(default=60) + TELEGRAM_ANTIFLOOD_MAX_UPDATES_PER_WINDOW: int = Field(default=180) WEBHOOK_QUEUE_NAME: str = Field(default="webhook-events") WEBHOOK_QUEUE_CONCURRENCY: int = Field(default=4) WORKER_PANEL_SYNC_INTERVAL_SECONDS: int = Field(default=900) diff --git a/tests/test_update_antiflood_middleware.py b/tests/test_update_antiflood_middleware.py new file mode 100644 index 0000000..cc6bec3 --- /dev/null +++ b/tests/test_update_antiflood_middleware.py @@ -0,0 +1,63 @@ +import unittest +from types import SimpleNamespace +from unittest.mock import AsyncMock, patch + +from bot.middlewares.update_antiflood import RateLimitRule, UpdateAntiFloodMiddleware + + +def _settings(**overrides): + base = { + "REDIS_URL": None, + "REDIS_KEY_PREFIX": "test-shop", + "TELEGRAM_ANTIFLOOD_ENABLED": True, + "TELEGRAM_ANTIFLOOD_WINDOW_SECONDS": 60, + "TELEGRAM_ANTIFLOOD_MAX_UPDATES_PER_WINDOW": 180, + } + base.update(overrides) + return SimpleNamespace(**base) + + +def _message_update(*, user_id=42, chat_id=42, chat_type="private", text="hello"): + return SimpleNamespace( + event_type="message", + message=SimpleNamespace( + from_user=SimpleNamespace(id=user_id), + chat=SimpleNamespace(id=chat_id, type=chat_type), + text=text, + ), + callback_query=None, + inline_query=None, + ) + + +class UpdateAntiFloodMiddlewareTests(unittest.IsolatedAsyncioTestCase): + async def test_extreme_update_flood_is_dropped_before_handler(self): + middleware = UpdateAntiFloodMiddleware( + _settings(), + default_rule=RateLimitRule(window_seconds=60, max_events=2), + ) + handler = AsyncMock(return_value="ok") + event = _message_update() + + with patch("bot.middlewares.update_antiflood.get_redis", AsyncMock(return_value=None)): + self.assertEqual(await middleware(handler, event, {}), "ok") + self.assertEqual(await middleware(handler, event, {}), "ok") + self.assertIsNone(await middleware(handler, event, {})) + + self.assertEqual(handler.await_count, 2) + + async def test_antiflood_can_be_disabled(self): + middleware = UpdateAntiFloodMiddleware( + _settings(TELEGRAM_ANTIFLOOD_ENABLED=False), + default_rule=RateLimitRule(window_seconds=60, max_events=0), + ) + handler = AsyncMock(return_value="ok") + + with patch("bot.middlewares.update_antiflood.get_redis", AsyncMock(return_value=None)): + self.assertEqual(await middleware(handler, _message_update(), {}), "ok") + + handler.assert_awaited_once() + + +if __name__ == "__main__": + unittest.main()