All middlewares now work through a single DB session

This commit is contained in:
Vladless
2025-12-11 15:49:38 +03:00
parent 985ae18bfb
commit 851e452292
6 changed files with 90 additions and 71 deletions
BIN
View File
Binary file not shown.
+22 -17
View File
@@ -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
+21 -26
View File
@@ -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:
+8 -4
View File
@@ -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
+19 -8
View File
@@ -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)
+20 -16
View File
@@ -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",
]
)