diff --git a/middlewares/__init__.py b/middlewares/__init__.py index 077bc247..968890a6 100644 --- a/middlewares/__init__.py +++ b/middlewares/__init__.py @@ -34,17 +34,19 @@ def register_middleware( if PROBE_LOGGING: dispatcher.update.outer_middleware(StreamProbeMiddleware("global")) + if sessionmaker: + dispatcher.update.outer_middleware(wrap(SessionMiddleware(sessionmaker), "session")) + if DISABLE_DIRECT_START: dispatcher.update.outer_middleware(wrap(DirectStartBlockerMiddleware(), "direct_start_blocker")) - if sessionmaker: - if CHANNEL_REQUIRED: - dispatcher.update.outer_middleware(wrap(SubscriptionMiddleware(), "subscription")) - dispatcher.update.outer_middleware(wrap(BanCheckerMiddleware(sessionmaker), "ban_checker")) + if CHANNEL_REQUIRED: + dispatcher.update.outer_middleware(wrap(SubscriptionMiddleware(), "subscription")) + + dispatcher.update.outer_middleware(wrap(BanCheckerMiddleware(), "ban_checker")) if middlewares is None: available_middlewares = { - "session": (SessionMiddleware(sessionmaker) if sessionmaker else SessionMiddleware()), "admin": AdminMiddleware(), "maintenance": MaintenanceModeMiddleware(), "logging": LoggingMiddleware(), diff --git a/middlewares/direct_start_blocker.py b/middlewares/direct_start_blocker.py index 6ba8975b..b798c829 100644 --- a/middlewares/direct_start_blocker.py +++ b/middlewares/direct_start_blocker.py @@ -42,17 +42,21 @@ class DirectStartBlockerMiddleware(BaseMiddleware): session = data.get("session") if not isinstance(session, AsyncSession): - logger.error("[DirectStartBlocker] session отсутствует в data") return await handler(event, data) tg_id = message.from_user.id text = message.text.strip() now = time.time() + user_in_data = bool(data.get("user")) async def user_exists_cached() -> bool: + if user_in_data: + return True + cached = _cache_user_exists.get(tg_id) if cached and cached[0] > now: return cached[1] + exists = await check_user_exists(session, tg_id) _cache_user_exists[tg_id] = (now + _TTL, exists) return exists diff --git a/middlewares/maintenance.py b/middlewares/maintenance.py index e8f04045..daf31e49 100644 --- a/middlewares/maintenance.py +++ b/middlewares/maintenance.py @@ -3,20 +3,11 @@ from typing import Any from aiogram import BaseMiddleware from aiogram.types import CallbackQuery, Message, Update -from sqlalchemy.ext.asyncio import AsyncSession -from config import ADMIN_ID from core.bootstrap import MANAGEMENT_CONFIG -from database.models import Admin class MaintenanceModeMiddleware(BaseMiddleware): - def __init__(self) -> None: - if isinstance(ADMIN_ID, (list, tuple, set)): - self._admin_ids = set(ADMIN_ID) - else: - self._admin_ids = {ADMIN_ID} - async def __call__( self, handler: Callable[[Update, dict[str, Any]], Awaitable[Any]], @@ -38,15 +29,7 @@ class MaintenanceModeMiddleware(BaseMiddleware): if not user_id: return - if user_id in self._admin_ids: - return await handler(event, data) - - session = data.get("session") - if not isinstance(session, AsyncSession): - return - - db_admin = await session.get(Admin, user_id) - if db_admin: + if data.get("admin"): return await handler(event, data) if isinstance(event, CallbackQuery): diff --git a/utils/versioning.py b/utils/versioning.py index acdca0a6..994b2bd6 100644 --- a/utils/versioning.py +++ b/utils/versioning.py @@ -92,4 +92,4 @@ def get_git_commit_number() -> str: def get_version() -> str: - return f"v.5.1-a101227 {get_git_commit_number()}" + return f"v.5.1-a111227 {get_git_commit_number()}"