diff --git a/middlewares.zip b/middlewares.zip new file mode 100644 index 00000000..ebd5e8c6 Binary files /dev/null and b/middlewares.zip differ diff --git a/middlewares/admin.py b/middlewares/admin.py index 123a7228..f0f3271a 100644 --- a/middlewares/admin.py +++ b/middlewares/admin.py @@ -4,19 +4,16 @@ from typing import Any from aiogram import BaseMiddleware from aiogram.types import CallbackQuery, Message, TelegramObject from sqlalchemy import select +from sqlalchemy.ext.asyncio import AsyncSession from config import ADMIN_ID from database.models import Admin class AdminMiddleware(BaseMiddleware): - """Middleware для проверки прав администратора. + """Проверяет, является ли пользователь администратором.""" - Добавляет в data['admin'] = True/False в зависимости от того, - является ли пользователь администратором. - """ - - _admin_ids: set[int] = set(ADMIN_ID) if isinstance(ADMIN_ID, list | tuple) else {ADMIN_ID} + _admin_ids: set[int] = set(ADMIN_ID) if isinstance(ADMIN_ID, (list, tuple)) else {ADMIN_ID} async def __call__( self, @@ -24,20 +21,28 @@ class AdminMiddleware(BaseMiddleware): event: TelegramObject, data: dict[str, Any], ) -> Any: - """Обрабатывает событие и добавляет флаг администратора в data.""" - data["admin"] = await self._check_admin_access(event, data.get("session")) + session: AsyncSession | None = data.get("session") + data["admin"] = await self._check_admin_access(event, session) return await handler(event, data) - async def _check_admin_access(self, event: TelegramObject, session) -> bool: - """Проверяет, имеет ли пользователь права администратора.""" + async def _check_admin_access( + self, + event: TelegramObject, + session: AsyncSession | None, + ) -> bool: try: user_id = None + if isinstance(event, Message): - user_id = event.from_user.id if event.from_user else None + if event.from_user: + user_id = event.from_user.id elif isinstance(event, CallbackQuery): - user_id = event.from_user.id if event.from_user else None + if event.from_user: + user_id = event.from_user.id else: - user_id = getattr(getattr(event, "from_user", None), "id", None) + from_user = getattr(event, "from_user", None) + if from_user: + user_id = getattr(from_user, "id", None) if not user_id: return False @@ -45,10 +50,10 @@ class AdminMiddleware(BaseMiddleware): if user_id in self._admin_ids: return True - if session: - result = await session.execute(select(Admin).where(Admin.tg_id == user_id)) - return result.scalar_one_or_none() is not None + if not session: + return False - return False + result = await session.execute(select(Admin).where(Admin.tg_id == user_id)) + return result.scalar_one_or_none() is not None except Exception: return False diff --git a/middlewares/ban_checker.py b/middlewares/ban_checker.py index 7e3f67e8..f749a3c6 100644 --- a/middlewares/ban_checker.py +++ b/middlewares/ban_checker.py @@ -19,7 +19,7 @@ _ban_cache: dict[int, tuple[float, dict | None]] = {} class BanCheckerMiddleware(BaseMiddleware): - def __init__(self, session_factory: Callable[[], AsyncSession]) -> None: + def __init__(self, session_factory: Callable[[], AsyncSession] | None = None) -> None: self.session_factory = session_factory async def __call__( @@ -50,32 +50,27 @@ class BanCheckerMiddleware(BaseMiddleware): if cached and cached[0] > now_ts: ban_info = cached[1] else: - session: AsyncSession | None = ( - data.get("session") if isinstance(data.get("session"), AsyncSession) else None - ) - created_here = False - if session is None: - session = self.session_factory() - created_here = True - try: - q = ( - select(ManualBan.reason, ManualBan.until) - .where( - ManualBan.tg_id == tg_id, - (ManualBan.until.is_(None)) | (ManualBan.until > datetime.utcnow()), - ) - .limit(1) + session = data.get("session") + if not isinstance(session, AsyncSession): + logger.error("[BanChecker] session отсутствует в data") + return await handler(event, data) + + query = ( + select(ManualBan.reason, ManualBan.until) + .where( + ManualBan.tg_id == tg_id, + (ManualBan.until.is_(None)) | (ManualBan.until > datetime.utcnow()), ) - res = await session.execute(q) - row = res.first() - if row: - reason, until = row - ban_info = {"reason": reason or "не указана", "until": until} - else: - ban_info = None - finally: - if created_here: - await session.close() + .limit(1) + ) + result = await session.execute(query) + row = result.first() + if row: + reason, until = row + ban_info = {"reason": reason or "не указана", "until": until} + else: + ban_info = None + _ban_cache[tg_id] = (now_ts + _BAN_CACHE_TTL, ban_info) if not ban_info: diff --git a/middlewares/direct_start_blocker.py b/middlewares/direct_start_blocker.py index e843d4b9..6ba8975b 100644 --- a/middlewares/direct_start_blocker.py +++ b/middlewares/direct_start_blocker.py @@ -1,14 +1,14 @@ import time - from collections.abc import Awaitable, Callable from typing import Any from aiogram import BaseMiddleware from aiogram.types import Message, Update +from sqlalchemy.ext.asyncio import AsyncSession from config import DISABLE_DIRECT_START from core.bootstrap import MODES_CONFIG -from database import async_session_maker, check_user_exists +from database import check_user_exists from logger import logger @@ -40,6 +40,11 @@ class DirectStartBlockerMiddleware(BaseMiddleware): if current_state: return await handler(event, data) + 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() @@ -48,8 +53,7 @@ class DirectStartBlockerMiddleware(BaseMiddleware): cached = _cache_user_exists.get(tg_id) if cached and cached[0] > now: return cached[1] - async with async_session_maker() as session: - exists = await check_user_exists(session, tg_id) + 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 e1477a33..e8f04045 100644 --- a/middlewares/maintenance.py +++ b/middlewares/maintenance.py @@ -3,14 +3,20 @@ 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 import async_session_maker 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]], @@ -23,20 +29,25 @@ class MaintenanceModeMiddleware(BaseMiddleware): user_id = None if isinstance(event, Message): - user_id = event.from_user.id + if event.from_user: + user_id = event.from_user.id elif isinstance(event, CallbackQuery): - user_id = event.from_user.id + if event.from_user: + user_id = event.from_user.id if not user_id: return - if user_id in ADMIN_ID: + if user_id in self._admin_ids: return await handler(event, data) - async with async_session_maker() as session: - db_admin = await session.get(Admin, user_id) - if db_admin: - 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: + return await handler(event, data) if isinstance(event, CallbackQuery): await event.answer("⚙️ Бот временно недоступен. Ведутся технические работы.", show_alert=True) diff --git a/middlewares/user.py b/middlewares/user.py index 47e638b6..b5fb7b2f 100644 --- a/middlewares/user.py +++ b/middlewares/user.py @@ -4,6 +4,7 @@ from typing import Any from aiogram import BaseMiddleware from aiogram.types import TelegramObject, User +from sqlalchemy.ext.asyncio import AsyncSession from database import upsert_user from logger import logger @@ -24,22 +25,23 @@ class UserMiddleware(BaseMiddleware): user: User | None = data.get("event_from_user") if user and not user.is_bot: session = data.get("session") - db_user = await self._process_user(user, session) - if db_user: - data["user"] = db_user + if isinstance(session, AsyncSession): + db_user = await self._process_user(user, session) + if db_user: + data["user"] = db_user except Exception as e: logger.error(f"Ошибка при обработке пользователя: {e}") return await handler(event, data) - async def _process_user(self, user: User, session: Any = None) -> dict | None: + async def _process_user(self, user: User, session: AsyncSession) -> dict | None: uid = user.id - fp = self._fingerprint(user) + fingerprint = self._fingerprint(user) now = monotonic() cached = self._cache.get(uid) if cached: - cached_fp, ts, cached_db_user = cached - if fp == cached_fp and now - ts < self._debounce: + cached_fingerprint, ts, cached_db_user = cached + if fingerprint == cached_fingerprint and now - ts < self._debounce: return cached_db_user logger.debug(f"Обработка пользователя: {uid}") @@ -53,17 +55,19 @@ class UserMiddleware(BaseMiddleware): session=session, only_if_exists=True, ) - self._cache[uid] = (fp, now, db_user) + self._cache[uid] = (fingerprint, now, db_user) if db_user: logger.debug(f"Получены данные пользователя из БД: {uid}") return db_user def _fingerprint(self, user: User) -> str: - return "|".join([ - str(user.id), - user.username or "", - user.first_name or "", - user.last_name or "", - user.language_code or "", - "1" if user.is_bot else "0", - ]) + return "|".join( + [ + str(user.id), + user.username or "", + user.first_name or "", + user.last_name or "", + user.language_code or "", + "1" if user.is_bot else "0", + ] + )