All middlewares now work through a single DB session
This commit is contained in:
Binary file not shown.
+22
-17
@@ -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
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
@@ -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",
|
||||
]
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user