users cache

This commit is contained in:
Vladless
2025-12-11 16:41:49 +03:00
parent 851e452292
commit 4f5878bc4e
4 changed files with 14 additions and 25 deletions
+7 -5
View File
@@ -34,17 +34,19 @@ def register_middleware(
if PROBE_LOGGING: if PROBE_LOGGING:
dispatcher.update.outer_middleware(StreamProbeMiddleware("global")) dispatcher.update.outer_middleware(StreamProbeMiddleware("global"))
if sessionmaker:
dispatcher.update.outer_middleware(wrap(SessionMiddleware(sessionmaker), "session"))
if DISABLE_DIRECT_START: if DISABLE_DIRECT_START:
dispatcher.update.outer_middleware(wrap(DirectStartBlockerMiddleware(), "direct_start_blocker")) dispatcher.update.outer_middleware(wrap(DirectStartBlockerMiddleware(), "direct_start_blocker"))
if sessionmaker: if CHANNEL_REQUIRED:
if CHANNEL_REQUIRED: dispatcher.update.outer_middleware(wrap(SubscriptionMiddleware(), "subscription"))
dispatcher.update.outer_middleware(wrap(SubscriptionMiddleware(), "subscription"))
dispatcher.update.outer_middleware(wrap(BanCheckerMiddleware(sessionmaker), "ban_checker")) dispatcher.update.outer_middleware(wrap(BanCheckerMiddleware(), "ban_checker"))
if middlewares is None: if middlewares is None:
available_middlewares = { available_middlewares = {
"session": (SessionMiddleware(sessionmaker) if sessionmaker else SessionMiddleware()),
"admin": AdminMiddleware(), "admin": AdminMiddleware(),
"maintenance": MaintenanceModeMiddleware(), "maintenance": MaintenanceModeMiddleware(),
"logging": LoggingMiddleware(), "logging": LoggingMiddleware(),
+5 -1
View File
@@ -42,17 +42,21 @@ class DirectStartBlockerMiddleware(BaseMiddleware):
session = data.get("session") session = data.get("session")
if not isinstance(session, AsyncSession): if not isinstance(session, AsyncSession):
logger.error("[DirectStartBlocker] session отсутствует в data")
return await handler(event, data) return await handler(event, data)
tg_id = message.from_user.id tg_id = message.from_user.id
text = message.text.strip() text = message.text.strip()
now = time.time() now = time.time()
user_in_data = bool(data.get("user"))
async def user_exists_cached() -> bool: async def user_exists_cached() -> bool:
if user_in_data:
return True
cached = _cache_user_exists.get(tg_id) cached = _cache_user_exists.get(tg_id)
if cached and cached[0] > now: if cached and cached[0] > now:
return cached[1] return cached[1]
exists = await check_user_exists(session, tg_id) exists = await check_user_exists(session, tg_id)
_cache_user_exists[tg_id] = (now + _TTL, exists) _cache_user_exists[tg_id] = (now + _TTL, exists)
return exists return exists
+1 -18
View File
@@ -3,20 +3,11 @@ from typing import Any
from aiogram import BaseMiddleware from aiogram import BaseMiddleware
from aiogram.types import CallbackQuery, Message, Update 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 core.bootstrap import MANAGEMENT_CONFIG
from database.models import Admin
class MaintenanceModeMiddleware(BaseMiddleware): 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__( async def __call__(
self, self,
handler: Callable[[Update, dict[str, Any]], Awaitable[Any]], handler: Callable[[Update, dict[str, Any]], Awaitable[Any]],
@@ -38,15 +29,7 @@ class MaintenanceModeMiddleware(BaseMiddleware):
if not user_id: if not user_id:
return return
if user_id in self._admin_ids: if data.get("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) return await handler(event, data)
if isinstance(event, CallbackQuery): if isinstance(event, CallbackQuery):
+1 -1
View File
@@ -92,4 +92,4 @@ def get_git_commit_number() -> str:
def get_version() -> 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()}"