optimization of database transactions/fixing middlewares/ cache of images
This commit is contained in:
+79
-32
@@ -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
@@ -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
@@ -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
@@ -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"))
|
||||
|
||||
@@ -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")
|
||||
Reference in New Issue
Block a user