optimization of database transactions/fixing middlewares/ cache of images

This commit is contained in:
Vladless
2025-10-14 01:05:52 +03:00
parent a1347dff04
commit 5a769f8972
5 changed files with 284 additions and 118 deletions
+79 -32
View File
@@ -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()
+43 -39
View File
@@ -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()
+60 -30
View File
@@ -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:
+29 -17
View File
@@ -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"))
+73
View File
@@ -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")