From 5a769f8972c7bddfafb0be822eeaeadc531fa700 Mon Sep 17 00:00:00 2001 From: Vladless Date: Tue, 14 Oct 2025 01:05:52 +0300 Subject: [PATCH] optimization of database transactions/fixing middlewares/ cache of images --- database/users.py | 111 ++++++++++++++++++++++++++++------------ handlers/start.py | 82 +++++++++++++++-------------- handlers/utils.py | 90 +++++++++++++++++++++----------- middlewares/__init__.py | 46 +++++++++++------ middlewares/probe.py | 73 ++++++++++++++++++++++++++ 5 files changed, 284 insertions(+), 118 deletions(-) create mode 100644 middlewares/probe.py diff --git a/database/users.py b/database/users.py index aa644281..f3ac2577 100644 --- a/database/users.py +++ b/database/users.py @@ -1,6 +1,6 @@ from datetime import datetime -from sqlalchemy import delete, exists, or_, select, update +from sqlalchemy import delete, exists, or_, select, update, func from sqlalchemy.dialects.postgresql import insert from sqlalchemy.exc import SQLAlchemyError from sqlalchemy.ext.asyncio import AsyncSession @@ -16,6 +16,7 @@ from database.models import ( Referral, TemporaryData, User, + Key, ) from logger import logger @@ -44,7 +45,6 @@ async def add_user( ) .on_conflict_do_nothing(index_elements=[User.tg_id]) ) - await session.execute(stmt) await session.commit() logger.info(f"[DB] Новый пользователь добавлен: {tg_id} (source: {source_code})") @@ -56,12 +56,19 @@ async def add_user( async def update_balance(session: AsyncSession, tg_id: int, amount: float) -> None: try: - result = await session.execute(select(User.balance).where(User.tg_id == tg_id)) - current = result.scalar_one_or_none() or 0 - new_balance = current + amount - await session.execute(update(User).where(User.tg_id == tg_id).values(balance=new_balance)) + res = await session.execute( + update(User) + .where(User.tg_id == tg_id) + .values(balance=func.coalesce(User.balance, 0) + amount) + .returning(User.balance) + ) + new_balance = res.scalar_one_or_none() await session.commit() - logger.info(f"[DB] Баланс пользователя {tg_id} обновлён: {current} → {new_balance}") + if new_balance is not None: + old_balance = new_balance - amount + logger.info(f"[DB] Баланс пользователя {tg_id} обновлён: {old_balance} → {new_balance}") + else: + logger.info(f"[DB] Баланс пользователя {tg_id} не изменён: пользователь не найден") except SQLAlchemyError as e: logger.error(f"[DB] Ошибка при обновлении баланса пользователя {tg_id}: {e}") await session.rollback() @@ -74,14 +81,17 @@ async def check_user_exists(session: AsyncSession, tg_id: int) -> bool: async def get_balance(session: AsyncSession, tg_id: int) -> float: - result = await session.execute(select(User.balance).where(User.tg_id == tg_id)) - balance = result.scalar_one_or_none() - return round(balance, 1) if balance is not None else 0.0 + result = await session.execute( + select(func.coalesce(User.balance, 0.0)).where(User.tg_id == tg_id) + ) + return round(float(result.scalar_one()), 1) async def set_user_balance(session: AsyncSession, tg_id: int, balance: float) -> None: try: - await session.execute(update(User).where(User.tg_id == tg_id).values(balance=balance)) + await session.execute( + update(User).where(User.tg_id == tg_id).values(balance=balance) + ) await session.commit() except SQLAlchemyError as e: logger.error(f"Ошибка при установке баланса для пользователя {tg_id}: {e}") @@ -90,7 +100,9 @@ async def set_user_balance(session: AsyncSession, tg_id: int, balance: float) -> async def update_trial(session: AsyncSession, tg_id: int, status: int): try: - await session.execute(update(User).where(User.tg_id == tg_id).values(trial=status)) + await session.execute( + update(User).where(User.tg_id == tg_id).values(trial=status) + ) await session.commit() logger.info(f"[DB] Триал статус обновлён для пользователя {tg_id}: {status}") except SQLAlchemyError as e: @@ -99,9 +111,10 @@ async def update_trial(session: AsyncSession, tg_id: int, status: int): async def get_trial(session: AsyncSession, tg_id: int) -> int: - result = await session.execute(select(User.trial).where(User.tg_id == tg_id)) - trial = result.scalar_one_or_none() - return trial or 0 + result = await session.execute( + select(func.coalesce(User.trial, 0)).where(User.tg_id == tg_id) + ) + return int(result.scalar_one()) async def upsert_user( @@ -120,7 +133,6 @@ async def upsert_user( user = result.scalar_one_or_none() if not user: return None - await session.execute( update(User) .where(User.tg_id == tg_id) @@ -133,8 +145,11 @@ async def upsert_user( updated_at=datetime.utcnow(), ) ) + await session.commit() + result = await session.execute(select(User).where(User.tg_id == tg_id)) + return dict(result.scalar_one().__dict__) else: - await session.execute( + res = await session.execute( insert(User) .values( tg_id=tg_id, @@ -157,10 +172,13 @@ async def upsert_user( "updated_at": datetime.utcnow(), }, ) + .returning(User) ) - await session.commit() - result = await session.execute(select(User).where(User.tg_id == tg_id)) - return dict(result.scalar_one().__dict__) + obj = res.scalar_one() + await session.commit() + d = obj.__dict__.copy() + d.pop("_sa_instance_state", None) + return d except SQLAlchemyError as e: logger.error(f"[DB] Ошибка при UPSERT пользователя {tg_id}: {e}") await session.rollback() @@ -170,29 +188,28 @@ async def upsert_user( async def delete_user_data(session: AsyncSession, tg_id: int): try: await session.execute(delete(Notification).where(Notification.tg_id == tg_id)) - - result = await session.execute(select(Gift.gift_id).where(Gift.sender_tg_id == tg_id)) - gift_ids = [row[0] for row in result.all()] - if gift_ids: - await session.execute(delete(GiftUsage).where(GiftUsage.gift_id.in_(gift_ids))) - + await session.execute( + delete(GiftUsage).where( + GiftUsage.gift_id.in_(select(Gift.gift_id).where(Gift.sender_tg_id == tg_id)) + ) + ) await session.execute(delete(Gift).where(Gift.sender_tg_id == tg_id)) - - await session.execute(update(Gift).where(Gift.recipient_tg_id == tg_id).values(recipient_tg_id=None)) - + await session.execute( + update(Gift).where(Gift.recipient_tg_id == tg_id).values(recipient_tg_id=None) + ) await session.execute(delete(Payment).where(Payment.tg_id == tg_id)) await session.execute( - delete(Referral).where(or_(Referral.referrer_tg_id == tg_id, Referral.referred_tg_id == tg_id)) + delete(Referral).where( + or_(Referral.referrer_tg_id == tg_id, Referral.referred_tg_id == tg_id) + ) ) await session.execute(delete(CouponUsage).where(CouponUsage.user_id == tg_id)) await delete_key(session, tg_id) await session.execute(delete(TemporaryData).where(TemporaryData.tg_id == tg_id)) await session.execute(delete(BlockedUser).where(BlockedUser.tg_id == tg_id)) await session.execute(delete(User).where(User.tg_id == tg_id)) - await session.commit() logger.info(f"[DB] Данные пользователя {tg_id} полностью удалены") - except SQLAlchemyError as e: await session.rollback() logger.error(f"[DB] Ошибка при удалении данных пользователя {tg_id}: {e}") @@ -202,3 +219,33 @@ async def delete_user_data(session: AsyncSession, tg_id: int): async def mark_trial_extended(tg_id: int, session: AsyncSession): await session.execute(update(User).where(User.tg_id == tg_id).values(trial=-1)) await session.commit() + + +async def get_user_snapshot(session: AsyncSession, tg_id: int) -> tuple[int, int] | None: + res = await session.execute( + select(func.coalesce(User.trial, 0), func.count(Key.client_id)) + .select_from(User) + .join(Key, Key.tg_id == User.tg_id, isouter=True) + .where(User.tg_id == tg_id) + .group_by(User.tg_id, User.trial) + ) + row = res.first() + if row is None: + return None + return int(row[0]), int(row[1]) + + +async def upsert_source_if_empty(session: AsyncSession, tg_id: int, source_code: str) -> None: + if not source_code: + return + stmt = ( + insert(User) + .values(tg_id=tg_id, source_code=source_code) + .on_conflict_do_update( + index_elements=[User.tg_id], + set_={"source_code": insert(User).excluded.source_code}, + where=(User.source_code.is_(None)), + ) + ) + await session.execute(stmt) + await session.commit() diff --git a/handlers/start.py b/handlers/start.py index acb69d34..bfcea265 100644 --- a/handlers/start.py +++ b/handlers/start.py @@ -1,6 +1,4 @@ -import asyncio import os - from typing import Any from aiogram import F, Router @@ -8,7 +6,6 @@ from aiogram.filters import Command from aiogram.fsm.context import FSMContext from aiogram.types import CallbackQuery, InlineKeyboardButton, Message from aiogram.utils.keyboard import InlineKeyboardBuilder -from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession from bot import bot @@ -24,12 +21,10 @@ from config import ( ) from database import ( add_user, - check_user_exists, get_coupon_by_code, - get_key_count, - get_trial, + get_user_snapshot, + upsert_source_if_empty, ) -from database.models import TrackingSource, User from handlers.buttons import ( ABOUT_VPN, BACK, @@ -60,7 +55,8 @@ from logger import logger from .admin.panel.keyboard import AdminPanelCallback from .refferal import handle_referral_link from .utils import edit_or_send_message, extract_user_data - +from sqlalchemy import select +from database.models import TrackingSource router = Router() processing_gifts = set() @@ -73,7 +69,8 @@ async def start_entry( ): message = event.message if isinstance(event, CallbackQuery) else event if CAPTCHA_ENABLE and captcha: - if not await check_user_exists(session, message.chat.id): + exists = await get_user_snapshot(session, message.chat.id) + if exists is None: captcha_data = await generate_captcha(message, state) await edit_or_send_message(message, captcha_data["text"], reply_markup=captcha_data["markup"]) return @@ -111,7 +108,12 @@ async def process_start_logic( user_data = user_data or extract_user_data(message.from_user or message.chat) text = text_to_process or message.text or message.caption if not text: - await show_start_menu(message, admin, session) + trial_key = await get_user_snapshot(session, user_data["tg_id"]) + trial = 0 + key_count = 0 + if trial_key is not None: + trial, key_count = trial_key + await show_start_menu(message, admin, session, trial=trial, key_count=key_count) return if text.startswith("/start "): @@ -122,7 +124,6 @@ async def process_start_logic( gift_detected = False for part in text.split("-"): await run_hooks("start_link", message=message, state=state, session=session, user_data=user_data, part=part) - if "coupons" in part: await handle_coupon_link(part, message, state, session, admin, user_data) continue @@ -139,19 +140,21 @@ async def process_start_logic( if gift_detected: return - if not await check_user_exists(session, user_data["tg_id"]): - await add_user(session=session, **user_data) + await add_user(session=session, **user_data) - trial_status = await get_trial(session, user_data["tg_id"]) - key_count = await get_key_count(session, user_data["tg_id"]) + trial_key = await get_user_snapshot(session, user_data["tg_id"]) + trial = 0 + key_count = 0 + if trial_key is not None: + trial, key_count = trial_key if SHOW_START_MENU_ONCE: - if key_count > 0 or trial_status == 1: + if key_count > 0 or trial == 1: await process_callback_view_profile(message, state, admin, session) else: - await show_start_menu(message, admin, session) + await show_start_menu(message, admin, session, trial=trial, key_count=key_count) else: - await show_start_menu(message, admin, session) + await show_start_menu(message, admin, session, trial=trial, key_count=key_count) async def handle_coupon_link(part, message, state, session, admin, user_data): @@ -159,7 +162,7 @@ async def handle_coupon_link(part, message, state, session, admin, user_data): coupon = await get_coupon_by_code(session, code) if coupon: await activate_coupon(message, state, session, code, admin=admin, user_data=user_data) - if coupon.days: + if getattr(coupon, "days", None): return @@ -169,21 +172,17 @@ async def handle_gift(part, message, state, session, user_data): await message.answer("❌ Неверный формат ссылки на подарок.") await process_callback_view_profile(message, state, False, session) return False - if gift_id in processing_gifts: await message.answer("⏳ Подарок уже обрабатывается, подождите...") await process_callback_view_profile(message, state, False, session) return False - processing_gifts.add(gift_id) try: gift_results = await run_hooks( "gift_activation", gift_id=gift_id, message=message, state=state, session=session, user_data=user_data ) - if gift_results and "SUCCESS" in gift_results: return True - await handle_gift_link(gift_id, message, state, session, user_data=user_data) return True finally: @@ -207,30 +206,34 @@ async def prompt_subscription(callback: CallbackQuery): async def handle_utm_link(utm_code: str, message: Message, state: FSMContext, session: AsyncSession, user_data: dict): - user_id = user_data["tg_id"] - result = await session.execute(select(TrackingSource).where(TrackingSource.code == utm_code)) - if not result.scalar_one_or_none(): + res = await session.execute(select(TrackingSource).where(TrackingSource.code == utm_code)) + if not res.scalar_one_or_none(): await message.answer("❌ UTM ссылка не найдена.") return - - user = (await session.execute(select(User).where(User.tg_id == user_id))).scalar_one_or_none() - if user and not user.source_code: - user.source_code = utm_code - await session.commit() - elif not user: - await add_user(session=session, source_code=utm_code, **user_data) + await upsert_source_if_empty(session, user_data["tg_id"], utm_code) -async def show_start_menu(message: Message, admin: bool, session: AsyncSession): +async def show_start_menu( + message: Message, + admin: bool, + session: AsyncSession, + trial: int | None = None, + key_count: int | None = None, +): image_path = os.path.join("img", "pic.jpg") kb = InlineKeyboardBuilder() - trial_status = await get_trial(session, message.chat.id) if session else None + if trial is None or key_count is None: + snap = await get_user_snapshot(session, message.chat.id) + if snap is None: + trial_status = 0 + else: + trial_status, key_cnt = snap + else: + trial_status = trial show_trial = (trial_status in (-1, 0)) and (not TRIAL_TIME_DISABLE) - show_profile = ((not SHOW_START_MENU_ONCE) or (trial_status not in (-1, 0)) or TRIAL_TIME_DISABLE) and ( - not show_trial - ) + show_profile = ((not SHOW_START_MENU_ONCE) or (trial_status not in (-1, 0)) or TRIAL_TIME_DISABLE) and (not show_trial) if show_trial: kb.row(InlineKeyboardButton(text=TRIAL_SUB, callback_data="create_key")) @@ -262,7 +265,8 @@ async def show_start_menu(message: Message, admin: bool, session: AsyncSession): @router.callback_query(F.data == "about_vpn") async def handle_about_vpn(callback: CallbackQuery, session: AsyncSession): user_id = callback.from_user.id - trial = await get_trial(session, user_id) + snap = await get_user_snapshot(session, user_id) + trial = 0 if snap is None else snap[0] back_target = "profile" if SHOW_START_MENU_ONCE and trial > 0 else "start" kb = InlineKeyboardBuilder() diff --git a/handlers/utils.py b/handlers/utils.py index ee7bab03..c3201115 100644 --- a/handlers/utils.py +++ b/handlers/utils.py @@ -4,6 +4,9 @@ import re import secrets import string +from collections import OrderedDict +import asyncio + from datetime import datetime, timedelta import aiofiles @@ -202,47 +205,74 @@ async def edit_or_send_message( disable_web_page_preview: bool = False, force_text: bool = False, ): - """ - Универсальная функция для редактирования исходного сообщения target_message. - """ + if not hasattr(edit_or_send_message, "cache"): + from collections import OrderedDict + import asyncio + edit_or_send_message.cache = OrderedDict() + edit_or_send_message.lock = asyncio.Lock() + edit_or_send_message.max = 256 + if media_path and os.path.isfile(media_path): + async with edit_or_send_message.lock: + cached_id = edit_or_send_message.cache.get(media_path) + if cached_id: + edit_or_send_message.cache.move_to_end(media_path) + if cached_id: + try: + await target_message.edit_media(InputMediaPhoto(media=cached_id, caption=text), reply_markup=reply_markup) + return + except Exception: + try: + await target_message.answer_photo( + photo=cached_id, + caption=text, + reply_markup=reply_markup, + disable_web_page_preview=disable_web_page_preview, + ) + return + except Exception: + pass + async with aiofiles.open(media_path, "rb") as f: - image_data = await f.read() - media = InputMediaPhoto( - media=BufferedInputFile(image_data, filename=os.path.basename(media_path)), - caption=text, - ) + data = await f.read() + upload = BufferedInputFile(data, filename=os.path.basename(media_path)) try: - await target_message.edit_media(media=media, reply_markup=reply_markup) - return + msg = await target_message.edit_media(InputMediaPhoto(media=upload, caption=text), reply_markup=reply_markup) except Exception: - await target_message.answer_photo( - photo=BufferedInputFile(image_data, filename=os.path.basename(media_path)), + msg = await target_message.answer_photo( + photo=upload, caption=text, reply_markup=reply_markup, disable_web_page_preview=disable_web_page_preview, ) - return - else: - if not force_text and target_message.caption is not None: - try: - await target_message.edit_caption(caption=text, reply_markup=reply_markup) - return - except Exception as e: - logger.error(f"Ошибка редактирования подписи: {e}") + if getattr(msg, "photo", None): + fid = msg.photo[-1].file_id + async with edit_or_send_message.lock: + if media_path not in edit_or_send_message.cache: + edit_or_send_message.cache[media_path] = fid + if len(edit_or_send_message.cache) > edit_or_send_message.max: + edit_or_send_message.cache.popitem(last=False) + return + + if not force_text and target_message.caption is not None: try: - await target_message.edit_text( - text=text, - reply_markup=reply_markup, - disable_web_page_preview=disable_web_page_preview, - ) + await target_message.edit_caption(caption=text, reply_markup=reply_markup) return except Exception: - await target_message.answer( - text=text, - reply_markup=reply_markup, - disable_web_page_preview=disable_web_page_preview, - ) + pass + try: + await target_message.edit_text( + text=text, + reply_markup=reply_markup, + disable_web_page_preview=disable_web_page_preview, + ) + return + except Exception: + await target_message.answer( + text=text, + reply_markup=reply_markup, + disable_web_page_preview=disable_web_page_preview, + ) def convert_to_bytes(value: float, unit: str) -> int: diff --git a/middlewares/__init__.py b/middlewares/__init__.py index 0d8f22cf..46075242 100644 --- a/middlewares/__init__.py +++ b/middlewares/__init__.py @@ -3,10 +3,12 @@ from collections.abc import Iterable from aiogram import Dispatcher from aiogram.dispatcher.middlewares.base import BaseMiddleware -from logger import logger +from config import DISABLE_DIRECT_START, CHANNEL_REQUIRED + from middlewares.ban_checker import BanCheckerMiddleware from middlewares.subscription import SubscriptionMiddleware +from .probe import StreamProbeMiddleware, MiddlewareProbe, TailHandlerProbe from .admin import AdminMiddleware from .answer import CallbackAnswerMiddleware from .direct_start_blocker import DirectStartBlockerMiddleware @@ -17,6 +19,9 @@ from .throttling import ThrottlingMiddleware from .user import UserMiddleware +PROBE_LOGGING = False + + def register_middleware( dispatcher: Dispatcher, middlewares: Iterable[BaseMiddleware | type[BaseMiddleware]] | None = None, @@ -24,13 +29,19 @@ def register_middleware( pool=None, sessionmaker=None, ) -> None: - """Регистрирует middleware в диспетчере.""" + def wrap(mw, name: str): + return MiddlewareProbe(mw, name) if PROBE_LOGGING else mw - dispatcher.update.outer_middleware(DirectStartBlockerMiddleware()) + if PROBE_LOGGING: + dispatcher.update.outer_middleware(StreamProbeMiddleware("global")) + + if DISABLE_DIRECT_START: + dispatcher.update.outer_middleware(wrap(DirectStartBlockerMiddleware(), "direct_start_blocker")) if sessionmaker: - dispatcher.update.outer_middleware(SubscriptionMiddleware()) - dispatcher.update.outer_middleware(BanCheckerMiddleware(sessionmaker)) + if CHANNEL_REQUIRED: + dispatcher.update.outer_middleware(wrap(SubscriptionMiddleware(), "subscription")) + dispatcher.update.outer_middleware(wrap(BanCheckerMiddleware(sessionmaker), "ban_checker")) if middlewares is None: available_middlewares = { @@ -42,19 +53,20 @@ def register_middleware( "user": UserMiddleware(), "answer": CallbackAnswerMiddleware(), } - exclude_set = set(exclude or []) - middlewares = [middleware for name, middleware in available_middlewares.items() if name not in exclude_set] - - handlers = [ - dispatcher.message, - dispatcher.callback_query, - dispatcher.inline_query, - ] + middlewares = [wrap(mw, name) for name, mw in available_middlewares.items() if name not in exclude_set] + else: + wrapped = [] + for mw in middlewares: + inst = mw() if isinstance(mw, type) else mw + wrapped.append(wrap(inst, getattr(inst, "name", inst.__class__.__name__))) + middlewares = wrapped + handlers = [dispatcher.message, dispatcher.callback_query, dispatcher.inline_query] for middleware in middlewares: - if isinstance(middleware, type): - middleware = middleware() + for h in handlers: + h.outer_middleware(middleware) - for handler in handlers: - handler.outer_middleware(middleware) + if PROBE_LOGGING: + for h in handlers: + h.outer_middleware(TailHandlerProbe("handler")) diff --git a/middlewares/probe.py b/middlewares/probe.py new file mode 100644 index 00000000..1959976c --- /dev/null +++ b/middlewares/probe.py @@ -0,0 +1,73 @@ +import time +from aiogram.dispatcher.middlewares.base import BaseMiddleware +from logger import logger + + +class StreamProbeMiddleware(BaseMiddleware): + def __init__(self, name: str = "global"): + self.name = name + + async def __call__(self, handler, event, data): + t0 = data.get("_mw_t0") + if t0 is None: + now = time.perf_counter() + data["_mw_t0"] = now + data["_mw_prev"] = now + try: + return await handler(event, data) + finally: + total = (time.perf_counter() - data["_mw_t0"]) * 1000 + logger.info(f"[mw:{self.name}] {total:.2f} ms") + + +class MiddlewareProbe(BaseMiddleware): + def __init__(self, inner: BaseMiddleware, name: str): + self.inner = inner + self.name = name + + async def __call__(self, handler, event, data): + now = time.perf_counter() + t0 = data.setdefault("_mw_t0", now) + prev = data.setdefault("_mw_prev", now) + logger.info(f"[mw:{self.name}] +{(now - prev)*1000:.2f} ms total {(now - t0)*1000:.2f} ms") + + downstream_ms = 0.0 + + async def timed_handler(event, data): + nonlocal downstream_ms + ts = time.perf_counter() + res = await handler(event, data) + downstream_ms = (time.perf_counter() - ts) * 1000 + return res + + start = time.perf_counter() + try: + return await self.inner(timed_handler, event, data) + finally: + end = time.perf_counter() + data["_mw_prev"] = end + total_ms = (end - start) * 1000 + self_ms = total_ms - downstream_ms + logger.info(f"[mw:{self.name}:self] {self_ms:.2f} ms") + logger.info(f"[mw:{self.name}:down] {downstream_ms:.2f} ms") + logger.info(f"[mw:{self.name}:total] {total_ms:.2f} ms") + + +class TailHandlerProbe(BaseMiddleware): + def __init__(self, name: str = "handler"): + self.name = name + + async def __call__(self, handler, event, data): + now = time.perf_counter() + t0 = data.setdefault("_mw_t0", now) + prev = data.setdefault("_mw_prev", now) + logger.info(f"[mw:{self.name}:enter] +{(now - prev)*1000:.2f} ms total {(now - t0)*1000:.2f} ms") + start = time.perf_counter() + try: + return await handler(event, data) + finally: + end = time.perf_counter() + data["_mw_prev"] = end + handler_ms = (end - start) * 1000 + total = (end - t0) * 1000 + logger.info(f"[mw:{self.name}] {handler_ms:.2f} ms total {total:.2f} ms")