diff --git a/core/app.cpython-312-x86_64-linux-gnu.so b/core/app.cpython-312-x86_64-linux-gnu.so index b1061e6c..bf827d64 100644 Binary files a/core/app.cpython-312-x86_64-linux-gnu.so and b/core/app.cpython-312-x86_64-linux-gnu.so differ diff --git a/core/cache_config.py b/core/cache_config.py index e546830c..979bc7ca 100644 --- a/core/cache_config.py +++ b/core/cache_config.py @@ -1,7 +1,7 @@ UPDATE_STALE_AGE_SEC = 60 CONCURRENCY_MAX_WAIT_SEC = 300 -CONCURRENCY_LIMIT = 25 +CONCURRENCY_LIMIT = 200 SUBSCRIPTION_CACHE_SUBSCRIBED_MAXSIZE = 200_000 SUBSCRIPTION_CACHE_SUBSCRIBED_TTL_SEC = 300 @@ -25,7 +25,7 @@ USER_EXISTS_CACHE_MAXSIZE = 150_000 USER_EXISTS_CACHE_TTL_SEC = 60 BAN_CACHE_MAXSIZE = 50_000 -BAN_CACHE_TTL_SEC = 30 +BAN_CACHE_TTL_SEC = 60 DIRECT_START_USER_EXISTS_CACHE_MAXSIZE = 50_000 DIRECT_START_USER_EXISTS_CACHE_TTL_SEC = 20 @@ -37,7 +37,8 @@ REMNAWAVE_SERVER_CACHE_MAXSIZE = 50_000 REMNAWAVE_SERVER_CACHE_TTL_SEC = 300 REMNAWAVE_PROFILE_CACHE_MAXSIZE = 200_000 REMNAWAVE_PROFILE_CACHE_TTL_SEC = 20 -REMNAWAVE_PROFILE_TIMEOUT_SEC = 3.0 +REMNAWAVE_PROFILE_TIMEOUT_SEC = 10.0 +REMNAWAVE_PROFILE_ERROR_CACHE_TTL_SEC = 45 REMNAWAVE_ACTION_TIMEOUT_SEC = 5.0 REMNAWAVE_MAX_CONCURRENCY = 20 @@ -51,7 +52,19 @@ SERVERS_CACHE_TTL_SEC = 60 TARIFF_BY_ID_CACHE_TTL_SEC = 120 TARIFFS_FOR_CLUSTER_CACHE_TTL_SEC = 120 +KEY_DETAILS_CACHE_TTL_SEC = 45 +KEYS_LIST_CACHE_TTL_SEC = 25 +KEY_COUNT_CACHE_TTL_SEC = 25 + +BALANCE_CACHE_TTL_SEC = 25 +PROFILE_DATA_CACHE_TTL_SEC = 25 + +PAYMENT_PENDING_CACHE_TTL_SEC = 3600 ERROR_THROTTLE_WINDOW_SEC = 60 ERROR_THROTTLE_MAX_KEYS = 500 ERROR_THROTTLE_MESSAGE_MAX_LEN = 120 + +WEBHOOK_ABUSE_FAIL_THRESHOLD = 10 +WEBHOOK_ABUSE_FAIL_WINDOW_SEC = 60 +WEBHOOK_ABUSE_BLOCK_TTL_SEC = 300 diff --git a/core/executor.py b/core/executor.py index 9c7e881c..5682890c 100644 --- a/core/executor.py +++ b/core/executor.py @@ -77,6 +77,20 @@ def shutdown_process_pool() -> None: logger.debug("[Executor] Пул процессов остановлен") +def should_run_heavy_tasks_separately() -> bool: + """ + True, если есть запас по ядрам/потокам — тогда рассылка и уведомления + можно выносить в отдельный поток/ядро. + """ + try: + from config import EXECUTOR_POOL_SIZE + pool_size = max(1, int(EXECUTOR_POOL_SIZE)) + except Exception: + pool_size = 1 + cpu_count = multiprocessing.cpu_count() or 1 + return cpu_count >= 2 or pool_size >= 2 + + async def run_io(fn: Callable[..., T], *args: object) -> T: """Выполняет fn(*args) в пуле потоков (I/O). Один вызов для всех блокирующих операций.""" loop = asyncio.get_running_loop() diff --git a/core/webhook_abuse.py b/core/webhook_abuse.py new file mode 100644 index 00000000..dd3d862f --- /dev/null +++ b/core/webhook_abuse.py @@ -0,0 +1,52 @@ +from aiohttp import web + +from core.cache_config import ( + WEBHOOK_ABUSE_BLOCK_TTL_SEC, + WEBHOOK_ABUSE_FAIL_THRESHOLD, + WEBHOOK_ABUSE_FAIL_WINDOW_SEC, +) +from core.redis_cache import cache_delete, cache_get, cache_incr, cache_key, cache_set + + +def get_webhook_client_ip(request: web.Request) -> str: + """IP клиента: X-Forwarded-For (первый) или X-Real-IP, иначе request.remote.""" + forwarded = request.headers.get("X-Forwarded-For") + if forwarded: + return forwarded.split(",")[0].strip() or "unknown" + real = request.headers.get("X-Real-IP") + if real: + return real.strip() or "unknown" + if request.remote: + s = str(request.remote) + if "%" in s: + s = s.split("%")[0] + if ":" in s: + s = s.rsplit(":", 1)[0] + return s or "unknown" + return "unknown" + + +async def is_webhook_ip_blocked(ip: str) -> bool: + """True, если IP временно заблокирован из‑за множества невалидных подписей.""" + if not ip or ip == "unknown": + return False + try: + block_key = cache_key("webhook_abuse_block", ip) + return (await cache_get(block_key)) is not None + except Exception: + return False + + +async def record_webhook_signature_failure(ip: str) -> None: + """Увеличивает счётчик неудачных проверок подписи для IP; при превышении порога блокирует IP.""" + if not ip or ip == "unknown": + return + try: + fail_key = cache_key("webhook_abuse_fail", ip) + count = await cache_incr(fail_key, WEBHOOK_ABUSE_FAIL_WINDOW_SEC) + if count >= WEBHOOK_ABUSE_FAIL_THRESHOLD: + block_key = cache_key("webhook_abuse_block", ip) + await cache_set(block_key, 1, WEBHOOK_ABUSE_BLOCK_TTL_SEC) + await cache_delete(fail_key) + except Exception: + pass diff --git a/database/keys.py b/database/keys.py index 15f39dfb..0b4964a6 100644 --- a/database/keys.py +++ b/database/keys.py @@ -1,14 +1,47 @@ from datetime import datetime +from types import SimpleNamespace from sqlalchemy import delete, func, select, text, update from sqlalchemy.exc import SQLAlchemyError from sqlalchemy.ext.asyncio import AsyncSession +from core.cache_config import ( + KEY_COUNT_CACHE_TTL_SEC, + KEY_DETAILS_CACHE_TTL_SEC, + KEYS_LIST_CACHE_TTL_SEC, +) +from core.redis_cache import cache_delete, cache_get, cache_key, cache_set from database.models import Key, User -from database.users import invalidate_user_snapshot +from database.users import invalidate_profile_cache, invalidate_user_snapshot from logger import logger +async def invalidate_key_details(email: str) -> None: + await cache_delete(cache_key("key_details", email)) + + +async def invalidate_key_email(client_id: str) -> None: + await cache_delete(cache_key("key_email", client_id)) + + +async def invalidate_keys_list(tg_id: int) -> None: + await cache_delete(cache_key("keys_list", tg_id)) + await cache_delete(cache_key("key_count", tg_id)) + await invalidate_profile_cache(tg_id) + + +async def invalidate_key_details_by_client_id(session: AsyncSession, client_id: str) -> None: + email = await cache_get(cache_key("key_email", client_id)) + await cache_delete(cache_key("key_email", client_id)) + if email: + await invalidate_key_details(str(email)) + else: + res = await session.execute(select(Key.email).where(Key.client_id == client_id).limit(1)) + row = res.scalar_one_or_none() + if row is not None: + await invalidate_key_details(str(row)) + + async def store_key( session: AsyncSession, tg_id: int, @@ -83,6 +116,8 @@ async def store_key( await session.commit() invalidate_user_snapshot(tg_id) + await invalidate_keys_list(tg_id) + await invalidate_key_details(email) except SQLAlchemyError as e: logger.error(f"❌ Ошибка при сохранении ключа: {e}") @@ -90,9 +125,29 @@ async def store_key( raise +def _key_to_cache_dict(k: Key) -> dict: + return { + "email": k.email, + "alias": k.alias, + "client_id": k.client_id, + "expiry_time": int(k.expiry_time) if k.expiry_time is not None else 0, + "created_at": int(k.created_at) if k.created_at is not None else 0, + "tariff_id": k.tariff_id, + "server_id": k.server_id, + "is_frozen": bool(k.is_frozen) if k.is_frozen is not None else False, + } + + async def get_keys(session: AsyncSession, tg_id: int): + ckey = cache_key("keys_list", tg_id) + cached = await cache_get(ckey) + if isinstance(cached, list): + return [SimpleNamespace(**d) for d in cached] result = await session.execute(select(Key).where(Key.tg_id == tg_id)) - return result.scalars().all() + rows = result.scalars().all() + serialized = [_key_to_cache_dict(k) for k in rows] + await cache_set(ckey, serialized, KEYS_LIST_CACHE_TTL_SEC) + return rows async def get_all_keys(session: AsyncSession): @@ -107,7 +162,12 @@ async def get_key_by_server(session: AsyncSession, tg_id: int, client_id: str): async def get_key_details(session: AsyncSession, email: str) -> dict | None: - """Возвращает подробную информацию о ключе по email.""" + """Возвращает подробную информацию о ключе по email. Горячие данные кэшируются в Redis.""" + ckey = cache_key("key_details", email) + cached = await cache_get(ckey) + if isinstance(cached, dict): + return cached + stmt = select(Key, User).join(User, Key.tg_id == User.tg_id).where(Key.email == email) result = await session.execute(stmt) row = result.first() @@ -127,7 +187,7 @@ async def get_key_details(session: AsyncSession, email: str) -> dict | None: hours_left = time_left.seconds // 3600 days_left_message = f"Осталось часов: {hours_left}" - return { + out = { "key": key.key, "remnawave_link": key.remnawave_link, "server_id": key.server_id, @@ -151,18 +211,36 @@ async def get_key_details(session: AsyncSession, email: str) -> dict | None: "current_device_limit": key.current_device_limit, "current_traffic_limit": key.current_traffic_limit, } + await cache_set(ckey, out, KEY_DETAILS_CACHE_TTL_SEC) + if key.client_id: + await cache_set(cache_key("key_email", key.client_id), email, KEY_DETAILS_CACHE_TTL_SEC) + return out async def get_key_count(session: AsyncSession, tg_id: int) -> int: + cached = await cache_get(cache_key("key_count", tg_id)) + if cached is not None: + try: + return int(cached) + except (TypeError, ValueError): + pass result = await session.execute(select(func.count()).select_from(Key).where(Key.tg_id == tg_id)) - return result.scalar() or 0 + count = result.scalar() or 0 + await cache_set(cache_key("key_count", tg_id), count, KEY_COUNT_CACHE_TTL_SEC) + return count async def delete_key(session: AsyncSession, identifier: int | str, commit: bool = True): tg_id_for_cache = None + email_for_cache = None if isinstance(identifier, str): - res = await session.execute(select(Key.tg_id).where(Key.client_id == identifier).limit(1)) - tg_id_for_cache = res.scalar_one_or_none() + res = await session.execute( + select(Key.tg_id, Key.email).where(Key.client_id == identifier).limit(1) + ) + row = res.first() + if row: + tg_id_for_cache, email_for_cache = row[0], row[1] + await cache_delete(cache_key("key_email", identifier)) else: tg_id_for_cache = identifier stmt = delete(Key).where(Key.tg_id == identifier if isinstance(identifier, int) else Key.client_id == identifier) @@ -171,12 +249,16 @@ async def delete_key(session: AsyncSession, identifier: int | str, commit: bool await session.commit() if tg_id_for_cache is not None: invalidate_user_snapshot(tg_id_for_cache) + await invalidate_keys_list(tg_id_for_cache) + if email_for_cache is not None: + await invalidate_key_details(str(email_for_cache)) logger.info(f"Ключ с идентификатором {identifier} удалён") async def update_key_expiry(session: AsyncSession, client_id: str, new_expiry_time: int): await session.execute(update(Key).where(Key.client_id == client_id).values(expiry_time=new_expiry_time)) await session.commit() + await invalidate_key_details_by_client_id(session, client_id) logger.info(f"Срок действия ключа {client_id} обновлён до {new_expiry_time}") @@ -188,6 +270,8 @@ async def get_client_id_by_email(session: AsyncSession, email: str): async def update_key_notified(session: AsyncSession, tg_id: int, client_id: str): await session.execute(update(Key).where(Key.tg_id == tg_id, Key.client_id == client_id).values(notified=True)) await session.commit() + await invalidate_keys_list(tg_id) + await invalidate_key_details_by_client_id(session, client_id) async def mark_key_as_frozen(session: AsyncSession, tg_id: int, client_id: str, time_left: int): @@ -203,6 +287,8 @@ async def mark_key_as_frozen(session: AsyncSession, tg_id: int, client_id: str, ), {"expiry": time_left, "tg_id": tg_id, "client_id": client_id}, ) + await invalidate_keys_list(tg_id) + await invalidate_key_details_by_client_id(session, client_id) async def mark_key_as_unfrozen( @@ -223,11 +309,14 @@ async def mark_key_as_unfrozen( ), {"expiry": new_expiry_time, "tg_id": tg_id, "client_id": client_id}, ) + await invalidate_keys_list(tg_id) + await invalidate_key_details_by_client_id(session, client_id) async def update_key_tariff(session: AsyncSession, client_id: str, tariff_id: int): await session.execute(update(Key).where(Key.client_id == client_id).values(tariff_id=tariff_id)) await session.commit() + await invalidate_key_details_by_client_id(session, client_id) logger.info(f"Тариф ключа {client_id} обновлён на {tariff_id}") @@ -239,6 +328,7 @@ async def get_subscription_link(session: AsyncSession, email: str) -> str | None async def update_key_client_id(session: AsyncSession, email: str, new_client_id: str): await session.execute(update(Key).where(Key.email == email).values(client_id=new_client_id)) await session.commit() + await invalidate_key_details(email) logger.info(f"client_id обновлён для {email} -> {new_client_id}") @@ -246,7 +336,10 @@ async def update_key_link(session: AsyncSession, email: str, link: str) -> bool: q = update(Key).where(Key.email == email).values(key=link).returning(Key.client_id) res = await session.execute(q) await session.commit() - return res.scalar_one_or_none() is not None + ok = res.scalar_one_or_none() is not None + if ok: + await invalidate_key_details(email) + return ok async def save_key_config_with_mode( @@ -280,6 +373,7 @@ async def save_key_config_with_mode( return await session.execute(update(Key).where(Key.email == email).values(**values)) + await invalidate_key_details(email) async def reset_key_current_limits_to_selected(session: AsyncSession, client_id: str): @@ -296,4 +390,5 @@ async def reset_key_current_limits_to_selected(session: AsyncSession, client_id: {"client_id": client_id}, ) await session.commit() + await invalidate_key_details_by_client_id(session, client_id) logger.info(f"Текущие лимиты ключа {client_id} сброшены к выбранным") diff --git a/database/notifications.py b/database/notifications.py index c61de549..103ed4ef 100644 --- a/database/notifications.py +++ b/database/notifications.py @@ -110,6 +110,30 @@ async def get_last_notification_time(session: AsyncSession, tg_id: int, notifica return None +async def get_last_notification_times_bulk( + session: AsyncSession, pairs: list[tuple[int, str]] +) -> dict[tuple[int, str], int]: + """ + Один запрос: последние времена уведомлений для списка (tg_id, notification_type). + Возвращает dict[(tg_id, notification_type)] -> timestamp_ms. + """ + if not pairs: + return {} + from sqlalchemy import tuple_ + + stmt = select( + Notification.tg_id, + Notification.notification_type, + Notification.last_notification_time, + ).where(tuple_(Notification.tg_id, Notification.notification_type).in_(pairs)) + result = await session.execute(stmt) + out = {} + for tg_id, ntype, last_time in result.all(): + if last_time: + out[(tg_id, ntype)] = int(last_time.timestamp() * 1000) + return out + + async def check_hot_lead_discount(session: AsyncSession, tg_id: int) -> dict: try: result = await session.execute( diff --git a/database/payments.py b/database/payments.py index 7ec35396..19dbdbb6 100644 --- a/database/payments.py +++ b/database/payments.py @@ -5,6 +5,8 @@ from sqlalchemy import and_, insert, select, update from sqlalchemy.exc import SQLAlchemyError from sqlalchemy.ext.asyncio import AsyncSession +from core.cache_config import PAYMENT_PENDING_CACHE_TTL_SEC +from core.redis_cache import cache_delete, cache_get, cache_key, cache_set from database.models import Payment from logger import logger @@ -12,6 +14,42 @@ from logger import logger MOSCOW_TZ = timezone("Europe/Moscow") +def _payment_cache_key(pid: str) -> str: + return cache_key("payment_pending", pid) + + +async def register_pending_payment( + payment_id: str, + tg_id: int, + amount: float, + payment_system: str, + *, + currency: str = "RUB", + metadata: dict | None = None, + original_amount: float | None = None, +) -> bool: + """Регистрирует ожидающий платёж только в Redis. В БД пишем при success/fail из вебхука.""" + data = { + "tg_id": tg_id, + "amount": amount, + "currency": currency, + "status": "pending", + "payment_system": payment_system, + "payment_id": payment_id, + "metadata": metadata, + "original_amount": original_amount, + } + ok = await cache_set(_payment_cache_key(payment_id), data, PAYMENT_PENDING_CACHE_TTL_SEC) + if ok: + logger.debug(f"[Payments] Pending в кэше: payment_id={payment_id}, tg_id={tg_id}") + return ok + + +async def invalidate_payment_cache(payment_id: str) -> None: + """Вызвать после сохранения платежа в БД (success/fail) из вебхука.""" + await cache_delete(_payment_cache_key(payment_id)) + + async def add_payment( session: AsyncSession, tg_id: int, @@ -142,6 +180,21 @@ async def update_payment_status( async def get_payment_by_payment_id(session: AsyncSession, pid: str) -> dict | None: + """Сначала Redis (pending), затем БД. Из кэша возвращается запись без id — вебхук делает add_payment.""" + cached = await cache_get(_payment_cache_key(pid)) + if cached is not None: + return { + "id": None, + "tg_id": cached["tg_id"], + "amount": cached["amount"], + "currency": cached.get("currency", "RUB"), + "status": cached.get("status", "pending"), + "payment_system": cached["payment_system"], + "payment_id": cached["payment_id"], + "created_at": None, + "metadata": cached.get("metadata"), + "original_amount": cached.get("original_amount"), + } try: result = await session.execute(select(Payment).where(Payment.payment_id == pid).limit(1)) payment = result.scalar_one_or_none() diff --git a/database/tariffs.py b/database/tariffs.py index c590d44f..bbe1b7ed 100644 --- a/database/tariffs.py +++ b/database/tariffs.py @@ -99,6 +99,35 @@ async def get_tariffs( return [] +async def get_tariff_names_groups_subgroups_durations( + session: AsyncSession, tariff_ids: list[int] +) -> tuple[dict[int, str], dict[int, str], dict[int, str | None], dict[int, int]]: + """Один запрос: id, name, group_code, subgroup_title, duration_days → четыре словаря.""" + if not tariff_ids: + return {}, {}, {}, {} + + result = await session.execute( + select( + Tariff.id, + Tariff.name, + Tariff.group_code, + Tariff.subgroup_title, + Tariff.duration_days, + ).where(Tariff.id.in_(tariff_ids)) + ) + rows = result.all() + names = {} + groups = {} + subgroups = {} + durations = {} + for tid, name, group_code, subgroup_title, duration_days in rows: + names[tid] = name + groups[tid] = group_code + subgroups[tid] = subgroup_title + durations[tid] = duration_days + return names, groups, subgroups, durations + + async def get_tariff_by_id(session: AsyncSession, tariff_id: int): key = cache_key("tariff", tariff_id) cached = await cache_get(key) diff --git a/database/users.py b/database/users.py index 990a138b..780a91a1 100644 --- a/database/users.py +++ b/database/users.py @@ -6,6 +6,7 @@ from sqlalchemy.exc import SQLAlchemyError from sqlalchemy.ext.asyncio import AsyncSession from core.cache_config import ( + BALANCE_CACHE_TTL_SEC, USER_EXISTS_CACHE_TTL_SEC, USER_SNAPSHOT_CACHE_TTL_SEC, ) @@ -76,6 +77,14 @@ async def add_user( raise +async def invalidate_balance_cache(tg_id: int) -> None: + await cache_delete(cache_key("balance", tg_id)) + + +async def invalidate_profile_cache(tg_id: int) -> None: + await cache_delete(cache_key("profile_data", tg_id)) + + async def update_balance(session: AsyncSession, tg_id: int, amount: float) -> None: try: res = await session.execute( @@ -91,6 +100,8 @@ async def update_balance(session: AsyncSession, tg_id: int, amount: float) -> No logger.info(f"[DB] Баланс пользователя {tg_id} обновлён: {old_balance} → {new_balance}") else: logger.info(f"[DB] Баланс пользователя {tg_id} не изменён: пользователь не найден") + await invalidate_balance_cache(tg_id) + await invalidate_profile_cache(tg_id) except SQLAlchemyError as e: logger.error(f"[DB] Ошибка при обновлении баланса пользователя {tg_id}: {e}") await session.rollback() @@ -108,15 +119,25 @@ async def check_user_exists(session: AsyncSession, tg_id: int) -> bool: async def get_balance(session: AsyncSession, tg_id: int) -> float: + cached = await cache_get(cache_key("balance", tg_id)) + if cached is not None: + try: + return round(float(cached), 1) + except (TypeError, ValueError): + pass result = await session.execute(select(func.coalesce(User.balance, 0.0)).where(User.tg_id == tg_id)) balance = result.scalar_one_or_none() - return round(float(balance or 0.0), 1) + value = round(float(balance or 0.0), 1) + await cache_set(cache_key("balance", tg_id), value, BALANCE_CACHE_TTL_SEC) + return value 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.commit() + await invalidate_balance_cache(tg_id) + await invalidate_profile_cache(tg_id) except SQLAlchemyError as e: logger.error(f"Ошибка при установке баланса для пользователя {tg_id}: {e}") await session.rollback() @@ -127,6 +148,7 @@ 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.commit() + await invalidate_profile_cache(tg_id) invalidate_user_snapshot(tg_id) logger.info(f"[DB] Триал статус обновлён для пользователя {tg_id}: {status}") except SQLAlchemyError as e: @@ -141,6 +163,45 @@ async def get_trial(session: AsyncSession, tg_id: int) -> int: return int(trial or 0) +async def get_balance_and_trial(session: AsyncSession, tg_id: int) -> tuple[float, int]: + """Один запрос к БД для баланса и триала (профиль при промахе кэша).""" + result = await session.execute( + select( + func.coalesce(User.balance, 0.0), + func.coalesce(User.trial, 0), + ).where(User.tg_id == tg_id) + ) + row = result.one_or_none() + if row is None: + return 0.0, 0 + balance, trial = row + return round(float(balance or 0.0), 1), int(trial or 0) + + +async def get_balance_trial_key_count(session: AsyncSession, tg_id: int) -> tuple[float, int, int]: + """ + Один запрос: баланс, триал и число ключей пользователя (для профиля при промахе кэша). + Возвращает (balance_rub, trial_status, key_count). + """ + key_count_subq = select(func.count()).select_from(Key).where(Key.tg_id == User.tg_id).scalar_subquery() + result = await session.execute( + select( + func.coalesce(User.balance, 0.0), + func.coalesce(User.trial, 0), + key_count_subq, + ).where(User.tg_id == tg_id) + ) + row = result.one_or_none() + if row is None: + return 0.0, 0, 0 + balance, trial, key_count = row + return ( + round(float(balance or 0.0), 1), + int(trial or 0), + int(key_count or 0), + ) + + async def upsert_user( session: AsyncSession, tg_id: int, diff --git a/handlers/admin/bans/bans_handler.py b/handlers/admin/bans/bans_handler.py index 62b5cafd..1307b0ec 100644 --- a/handlers/admin/bans/bans_handler.py +++ b/handlers/admin/bans/bans_handler.py @@ -16,6 +16,7 @@ from database import delete_user_data from database.models import BlockedUser, Key, ManualBan from filters.admin import IsAdminFilter from logger import logger +from middlewares.ban_checker import invalidate_ban_cache from ..panel.keyboard import AdminPanelCallback from .keyboard import ( @@ -245,8 +246,12 @@ async def handle_clear_shadow_bans(callback_query: CallbackQuery, session: Async ) return + tg_ids_result = await session.execute(select(ManualBan.tg_id).where(ManualBan.reason == "shadow")) + tg_ids_to_invalidate = [r[0] for r in tg_ids_result.all()] await session.execute(delete(ManualBan).where(ManualBan.reason == "shadow")) await session.commit() + for uid in tg_ids_to_invalidate: + await invalidate_ban_cache(uid) await callback_query.message.answer( text=f"🗑️ Очищено {total_count} записей теневых банов из базы данных.", @@ -279,8 +284,14 @@ async def handle_clear_manual_bans(callback_query: CallbackQuery, session: Async ) return + tg_ids_result = await session.execute( + select(ManualBan.tg_id).where(or_(ManualBan.reason != "shadow", ManualBan.reason.is_(None))) + ) + tg_ids_to_invalidate = [r[0] for r in tg_ids_result.all()] await session.execute(delete(ManualBan).where(or_(ManualBan.reason != "shadow", ManualBan.reason.is_(None)))) await session.commit() + for uid in tg_ids_to_invalidate: + await invalidate_ban_cache(uid) await callback_query.message.answer( text=f"🗑️ Очищено {total_count} записей ручных банов из базы данных.", @@ -358,6 +369,8 @@ async def handle_preemptive_ids_input(message: Message, state: FSMContext, sessi await session.execute(stmt) await session.commit() + for uid in tg_ids: + await invalidate_ban_cache(uid) await message.answer( f"✅ Успешно добавлено в теневой бан: {len(tg_ids)} пользователей.", diff --git a/handlers/admin/sender/sender_handler.py b/handlers/admin/sender/sender_handler.py index 894e361d..d83c6438 100644 --- a/handlers/admin/sender/sender_handler.py +++ b/handlers/admin/sender/sender_handler.py @@ -1,3 +1,5 @@ +import asyncio + from aiogram import F, Router from aiogram.exceptions import TelegramBadRequest from aiogram.fsm.context import FSMContext @@ -5,13 +7,15 @@ from aiogram.types import CallbackQuery, InlineKeyboardButton, InlineKeyboardMar from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession +from config import API_TOKEN +from core.executor import run_io, should_run_heavy_tasks_separately from database.models import Server from filters.admin import IsAdminFilter from logger import logger from ..panel.keyboard import AdminPanelCallback, build_admin_back_kb from .keyboard import AdminSenderCallback, build_clusters_kb, build_sender_kb -from .sender_service import BroadcastService +from .sender_service import BroadcastService, run_broadcast_in_thread from .sender_states import AdminSender from .sender_utils import get_recipients, parse_message_buttons @@ -192,33 +196,65 @@ async def handle_broadcast_confirm(callback_query: CallbackQuery, state: FSMCont _broadcast_progress_text(0, total_users_for_bar, 0, 0), ) - messages = [] - for tg_id in tg_ids: - message_data = {"tg_id": tg_id, "text": text_message, "photo": photo, "keyboard": keyboard} - messages.append(message_data) - bot = callback_query.bot + state_keyboard_data = data.get("keyboard") - async def on_progress(completed: int, total: int, sent: int, failed: int) -> None: - text = _broadcast_progress_text(completed, total, sent, failed) - try: - await bot.edit_message_text( - chat_id=status_message.chat.id, - message_id=status_message.message_id, - text=text, + if should_run_heavy_tasks_separately(): + main_loop = asyncio.get_running_loop() + + async def _edit_progress(completed: int, total: int, sent: int, failed: int) -> None: + text = _broadcast_progress_text(completed, total, sent, failed) + try: + await bot.edit_message_text( + chat_id=status_message.chat.id, + message_id=status_message.message_id, + text=text, + ) + except TelegramBadRequest as e: + if "message is not modified" not in str(e).lower(): + logger.debug(f"[Sender] Обновление прогресса: {e}") + + def progress_cb(completed: int, total: int, sent: int, failed: int) -> None: + main_loop.call_soon_threadsafe( + lambda c=completed, t=total, s=sent, f=failed: asyncio.ensure_future( + _edit_progress(c, t, s, f), loop=main_loop + ) ) - except TelegramBadRequest as e: - if "message is not modified" not in str(e).lower(): - logger.debug(f"[Sender] Обновление прогресса: {e}") - broadcast_service = BroadcastService(bot=bot, session=session, messages_per_second=35) + stats = await run_io( + run_broadcast_in_thread, + API_TOKEN, + tg_ids, + text_message, + photo, + state_keyboard_data, + progress_cb, + ) + else: + messages = [] + for tg_id in tg_ids: + message_data = {"tg_id": tg_id, "text": text_message, "photo": photo, "keyboard": keyboard} + messages.append(message_data) - stats = await broadcast_service.broadcast( - messages, - workers=5, - on_progress=on_progress, - progress_interval=2.0, - ) + async def on_progress(completed: int, total: int, sent: int, failed: int) -> None: + text = _broadcast_progress_text(completed, total, sent, failed) + try: + await bot.edit_message_text( + chat_id=status_message.chat.id, + message_id=status_message.message_id, + text=text, + ) + except TelegramBadRequest as e: + if "message is not modified" not in str(e).lower(): + logger.debug(f"[Sender] Обновление прогресса: {e}") + + broadcast_service = BroadcastService(bot=bot, session=session, messages_per_second=35) + stats = await broadcast_service.broadcast( + messages, + workers=5, + on_progress=on_progress, + progress_interval=2.0, + ) duration_minutes = int(stats["total_duration"] // 60) duration_seconds = int(stats["total_duration"] % 60) diff --git a/handlers/admin/sender/sender_service.py b/handlers/admin/sender/sender_service.py index 30c61544..7093676a 100644 --- a/handlers/admin/sender/sender_service.py +++ b/handlers/admin/sender/sender_service.py @@ -6,13 +6,60 @@ from collections import deque from typing import Any from aiogram import Bot +from aiogram.client.default import DefaultBotProperties +from aiogram.enums import ParseMode from aiogram.exceptions import TelegramBadRequest, TelegramForbiddenError, TelegramRetryAfter +from aiogram.types import InlineKeyboardMarkup from sqlalchemy.ext.asyncio import AsyncSession from database import async_session_maker from logger import logger +def run_broadcast_in_thread( + api_token: str, + tg_ids: list[int], + text_message: str, + photo: str | None, + keyboard_data: dict | None, + progress_cb: Callable[[int, int, int, int], None] | None = None, +) -> dict: + """ + Синхронная обёртка: запускает рассылку в отдельном event loop в текущем потоке. + """ + loop = asyncio.new_event_loop() + asyncio.set_event_loop(loop) + bot = None + try: + bot = Bot(token=api_token, default=DefaultBotProperties(parse_mode=ParseMode.HTML)) + keyboard = InlineKeyboardMarkup.model_validate(keyboard_data) if keyboard_data else None + messages = [ + {"tg_id": tg_id, "text": text_message, "photo": photo, "keyboard": keyboard} + for tg_id in tg_ids + ] + service = BroadcastService(bot=bot, session=None, messages_per_second=35) + + async def on_progress(completed: int, total: int, sent: int, failed: int) -> None: + if progress_cb: + progress_cb(completed, total, sent, failed) + + return loop.run_until_complete( + service.broadcast( + messages, + workers=5, + on_progress=on_progress, + progress_interval=2.0, + ) + ) + finally: + if bot is not None and bot.session is not None: + try: + loop.run_until_complete(bot.session.close()) + except Exception: + pass + loop.close() + + class BroadcastMessage: def __init__(self, tg_id: int, text: str, photo: str | None = None, keyboard: Any = None) -> None: self.tg_id = tg_id diff --git a/handlers/admin/stats/stats_handler.py b/handlers/admin/stats/stats_handler.py index cdb15119..2bdc712a 100644 --- a/handlers/admin/stats/stats_handler.py +++ b/handlers/admin/stats/stats_handler.py @@ -22,10 +22,7 @@ from database import ( count_users_registered_since, count_users_updated_today, get_tariff_distribution, - get_tariff_durations, - get_tariff_groups, - get_tariff_names, - get_tariff_subgroups, + get_tariff_names_groups_subgroups_durations, sum_payments_between, sum_payments_since, sum_total_payments, @@ -54,21 +51,15 @@ async def handle_stats(callback_query: CallbackQuery, session: AsyncSession): now = datetime.now(moscow_tz) today = now.date() - total_users = await count_total_users(session) today_start = moscow_tz.localize(datetime.combine(today, datetime.min.time())) today_start_utc = today_start.astimezone(pytz.UTC).replace(tzinfo=None) - users_updated_today = await count_users_updated_today(session, today_start_utc) - registrations_today = await count_users_registered_since(session, today_start_utc) - yesterday_date = today - timedelta(days=1) yesterday_start = moscow_tz.localize(datetime.combine(yesterday_date, datetime.min.time())) yesterday_end = moscow_tz.localize(datetime.combine(today, datetime.min.time())) yesterday_start_utc = yesterday_start.astimezone(pytz.UTC).replace(tzinfo=None) yesterday_end_utc = yesterday_end.astimezone(pytz.UTC).replace(tzinfo=None) - registrations_yesterday = await count_users_registered_between(session, yesterday_start_utc, yesterday_end_utc) - week_start_date = today - timedelta(days=today.weekday()) week_start = moscow_tz.localize(datetime.combine(week_start_date, datetime.min.time())) week_start_utc = week_start.astimezone(pytz.UTC).replace(tzinfo=None) @@ -77,32 +68,35 @@ async def handle_stats(callback_query: CallbackQuery, session: AsyncSession): month_start = moscow_tz.localize(datetime.combine(month_start_date, datetime.min.time())) month_start_utc = month_start.astimezone(pytz.UTC).replace(tzinfo=None) - registrations_week = await count_users_registered_since(session, week_start_utc) - registrations_month = await count_users_registered_since(session, month_start_utc) - last_month_start_date = (today.replace(day=1) - timedelta(days=1)).replace(day=1) this_month_start_date = today.replace(day=1) - last_month_start = moscow_tz.localize(datetime.combine(last_month_start_date, datetime.min.time())) last_month_end = moscow_tz.localize(datetime.combine(this_month_start_date, datetime.min.time())) last_month_start_utc = last_month_start.astimezone(pytz.UTC).replace(tzinfo=None) last_month_end_utc = last_month_end.astimezone(pytz.UTC).replace(tzinfo=None) + total_users = await count_total_users(session) + users_updated_today = await count_users_updated_today(session, today_start_utc) + registrations_today = await count_users_registered_since(session, today_start_utc) + registrations_yesterday = await count_users_registered_between( + session, yesterday_start_utc, yesterday_end_utc + ) + registrations_week = await count_users_registered_since(session, week_start_utc) + registrations_month = await count_users_registered_since(session, month_start_utc) registrations_last_month = await count_users_registered_between( session, last_month_start_utc, last_month_end_utc ) - total_keys = await count_total_keys(session) active_keys = await count_active_keys(session) active_paid_keys = await count_active_paid_keys(session) active_trial_keys = await count_active_trial_keys(session) - expired_keys = total_keys - active_keys - tariff_counts, no_tariff_keys = await get_tariff_distribution(session, include_unbound=True) - tariff_names = await get_tariff_names(session, [tid for tid, _ in tariff_counts]) - tariff_groups = await get_tariff_groups(session, [tid for tid, _ in tariff_counts]) - tariff_subgroups = await get_tariff_subgroups(session, [tid for tid, _ in tariff_counts]) - tariff_durations = await get_tariff_durations(session, [tid for tid, _ in tariff_counts]) + + expired_keys = total_keys - active_keys + tariff_ids = [tid for tid, _ in tariff_counts] + tariff_names, tariff_groups, tariff_subgroups, tariff_durations = ( + await get_tariff_names_groups_subgroups_durations(session, tariff_ids) + ) grouped_tariffs = {} for tid, count in tariff_counts: diff --git a/handlers/admin/users/users_bans.py b/handlers/admin/users/users_bans.py index ca1a224c..da9d506d 100644 --- a/handlers/admin/users/users_bans.py +++ b/handlers/admin/users/users_bans.py @@ -11,6 +11,7 @@ from handlers.buttons import BACK from database.models import ManualBan from filters.admin import IsAdminFilter +from middlewares.ban_checker import invalidate_ban_cache from .keyboard import AdminUserEditorCallback, build_editor_btn, build_editor_kb, build_user_ban_type_kb from .users_states import BanUserStates @@ -81,6 +82,7 @@ async def handle_ban_forever_reason_input(message: Message, state: FSMContext, s await session.execute(stmt) await session.commit() + await invalidate_ban_cache(tg_id) await state.clear() await message.answer( @@ -161,6 +163,7 @@ async def handle_ban_duration_input(message: Message, state: FSMContext, session await session.execute(stmt) await session.commit() + await invalidate_ban_cache(tg_id) text = ( f"✅ Пользователь {tg_id} временно забанен до {until:%Y-%m-%d %H:%M} по UTC." @@ -200,6 +203,7 @@ async def handle_ban_shadow(callback: CallbackQuery, callback_data: AdminUserEdi ) await session.execute(stmt) await session.commit() + await invalidate_ban_cache(callback_data.tg_id) await callback.message.edit_text( text=f"👻 Пользователь {callback_data.tg_id} получил теневой бан.", @@ -218,6 +222,7 @@ async def handle_user_unban( ): await session.execute(delete(ManualBan).where(ManualBan.tg_id == callback_data.tg_id)) await session.commit() + await invalidate_ban_cache(callback_data.tg_id) text = ( f"✅ Пользователь {callback_data.tg_id} разблокирован. Нажмите кнопку ниже для возврата в профиль." diff --git a/handlers/instructions/instructions.py b/handlers/instructions/instructions.py index 58d15728..e3a35665 100644 --- a/handlers/instructions/instructions.py +++ b/handlers/instructions/instructions.py @@ -42,6 +42,7 @@ from handlers.texts import ( ROUTER_MESSAGE, SUBSCRIPTION_DETAILS_TEXT, ) +from handlers.keys.utils import key_owned_by_user from handlers.utils import edit_or_send_message, is_full_remnawave_cluster from hooks.processors import process_remnawave_webapp_override @@ -75,6 +76,10 @@ async def send_instructions(callback_query_or_message: CallbackQuery | Message): @router.callback_query(F.data.startswith("connect_pc|")) async def process_connect_pc(callback_query: CallbackQuery, session: Any): key_name = callback_query.data.split("|")[1] + record = await get_key_details(session, key_name) + if not key_owned_by_user(record, callback_query.from_user.id): + await callback_query.answer("Доступ запрещён.", show_alert=True) + return key_link = await get_subscription_link(session, key_name) if not key_link: builder = InlineKeyboardBuilder() @@ -103,6 +108,10 @@ async def process_connect_pc(callback_query: CallbackQuery, session: Any): @router.callback_query(F.data.startswith("windows_menu|")) async def process_windows_menu(callback_query: CallbackQuery, session: Any): key_name = callback_query.data.split("|")[1] + record = await get_key_details(session, key_name) + if not key_owned_by_user(record, callback_query.from_user.id): + await callback_query.answer("Доступ запрещён.", show_alert=True) + return key_link = await get_subscription_link(session, key_name) if not key_link: await callback_query.message.answer("❌ Ошибка: ключ не найден.") @@ -136,6 +145,10 @@ async def process_windows_menu(callback_query: CallbackQuery, session: Any): @router.callback_query(F.data.startswith("macos_menu|")) async def process_macos_menu(callback_query: CallbackQuery, session: Any): key_name = callback_query.data.split("|")[1] + record = await get_key_details(session, key_name) + if not key_owned_by_user(record, callback_query.from_user.id): + await callback_query.answer("Доступ запрещён.", show_alert=True) + return key_link = await get_subscription_link(session, key_name) if not key_link: await callback_query.message.answer("❌ Ошибка: ключ не найден.") @@ -171,6 +184,9 @@ async def process_connect_tv(callback_query: CallbackQuery, session: Any): key_name = callback_query.data.split("|", 1)[1] record = await get_key_details(session, key_name) + if not key_owned_by_user(record, callback_query.from_user.id): + await callback_query.answer("Доступ запрещён.", show_alert=True) + return final_link = None is_full_remnawave = False use_webapp = False @@ -221,6 +237,10 @@ async def process_connect_tv(callback_query: CallbackQuery, session: Any): @router.callback_query(F.data.startswith("continue_tv|")) async def process_continue_tv(callback_query: CallbackQuery, session: Any): key_name = callback_query.data.split("|")[1] + record = await get_key_details(session, key_name) + if not key_owned_by_user(record, callback_query.from_user.id): + await callback_query.answer("Доступ запрещён.", show_alert=True) + return key_link = await get_subscription_link(session, key_name) message_text = SUBSCRIPTION_DETAILS_TEXT.format(subscription_link=key_link) @@ -239,6 +259,10 @@ async def process_continue_tv(callback_query: CallbackQuery, session: Any): @router.callback_query(F.data.startswith("connect_router|")) async def process_connect_router(callback_query: CallbackQuery, session: Any): key_name = callback_query.data.split("|")[1] + record = await get_key_details(session, key_name) + if not key_owned_by_user(record, callback_query.from_user.id): + await callback_query.answer("Доступ запрещён.", show_alert=True) + return key_link = await get_subscription_link(session, key_name) if not key_link: builder = InlineKeyboardBuilder() diff --git a/handlers/keys/key_connect.py b/handlers/keys/key_connect.py index 3deeb3d9..90bc7d94 100644 --- a/handlers/keys/key_connect.py +++ b/handlers/keys/key_connect.py @@ -19,7 +19,7 @@ from config import ( DOWNLOAD_IOS, INSTRUCTIONS_BUTTON, ) -from database import Key, get_subscription_link +from database import Key, get_key_details, get_subscription_link from handlers.buttons import ( ANDROID, BACK, @@ -39,6 +39,7 @@ from handlers.texts import ( IOS_DESCRIPTION_TEMPLATE, SUBSCRIPTION_DESCRIPTION, ) +from handlers.keys.utils import key_owned_by_user from handlers.utils import edit_or_send_message from hooks.hook_buttons import insert_hook_buttons from hooks.processors import process_connect_device_menu @@ -67,6 +68,10 @@ def generate_key_qr_file(qr_data: str, email: str) -> str: async def handle_connect_device(callback_query: CallbackQuery, session: AsyncSession): try: key_name = callback_query.data.split("|")[1] + record = await get_key_details(session, key_name) + if not key_owned_by_user(record, callback_query.from_user.id): + await callback_query.answer("Доступ запрещён.", show_alert=True) + return builder = InlineKeyboardBuilder() builder.row(InlineKeyboardButton(text=IPHONE, callback_data=f"connect_ios|{key_name}")) @@ -102,6 +107,10 @@ async def process_callback_connect_phone(callback_query: CallbackQuery, session: email = callback_query.data.split("|")[1] try: + record = await get_key_details(session, email) + if not key_owned_by_user(record, callback_query.from_user.id): + await callback_query.answer("Доступ запрещён.", show_alert=True) + return key_link = await get_subscription_link(session, email) if not key_link: await callback_query.message.answer("❌ Ошибка: ключ не найден.") @@ -148,6 +157,10 @@ async def process_callback_connect_ios(callback_query: CallbackQuery, session: A email = callback_query.data.split("|")[1] try: + record = await get_key_details(session, email) + if not key_owned_by_user(record, callback_query.from_user.id): + await callback_query.answer("Доступ запрещён.", show_alert=True) + return key_link = await get_subscription_link(session, email) if not key_link: await callback_query.message.answer("❌ Ошибка: ключ не найден.") @@ -188,6 +201,10 @@ async def process_callback_connect_android(callback_query: CallbackQuery, sessio email = callback_query.data.split("|")[1] try: + record = await get_key_details(session, email) + if not key_owned_by_user(record, callback_query.from_user.id): + await callback_query.answer("Доступ запрещён.", show_alert=True) + return key_link = await get_subscription_link(session, email) if not key_link: await callback_query.message.answer("❌ Ошибка: ключ не найден.") @@ -235,6 +252,9 @@ async def show_qr_code(callback_query: types.CallbackQuery, session: AsyncSessio if not record: await callback_query.message.answer("❌ Подписка не найдена.") return + if record.tg_id != callback_query.from_user.id: + await callback_query.answer("Доступ запрещён.", show_alert=True) + return qr_data = record.key or record.remnawave_link if not qr_data: diff --git a/handlers/keys/key_freeze.py b/handlers/keys/key_freeze.py index ea41d6b3..ecf487f6 100644 --- a/handlers/keys/key_freeze.py +++ b/handlers/keys/key_freeze.py @@ -21,6 +21,7 @@ from handlers.texts import ( SUBSCRIPTION_UNFROZEN_MSG, UNFREEZE_SUBSCRIPTION_CONFIRM_MSG, ) +from handlers.keys.utils import key_owned_by_user from handlers.utils import edit_or_send_message, handle_error from middlewares.session import release_session_early from logger import logger @@ -32,6 +33,10 @@ router = Router() @router.callback_query(F.data.startswith("unfreeze_subscription|")) async def process_callback_unfreeze_subscription(callback_query: CallbackQuery, session: Any): key_name = callback_query.data.split("|")[1] + record = await get_key_details(session, key_name) + if not key_owned_by_user(record, callback_query.from_user.id): + await callback_query.answer("Доступ запрещён.", show_alert=True) + return confirm_text = UNFREEZE_SUBSCRIPTION_CONFIRM_MSG builder = InlineKeyboardBuilder() @@ -66,6 +71,9 @@ async def process_callback_unfreeze_subscription_confirm(callback_query: Callbac if not record: await callback_query.message.answer("Ключ не найден.") return + if not key_owned_by_user(record, callback_query.from_user.id): + await callback_query.answer("Доступ запрещён.", show_alert=True) + return email = record["email"] client_id = record["client_id"] @@ -137,6 +145,10 @@ async def process_callback_freeze_subscription(callback_query: CallbackQuery, se Показывает пользователю диалог подтверждения заморозки (отключения) подписки. """ key_name = callback_query.data.split("|")[1] + record = await get_key_details(session, key_name) + if not key_owned_by_user(record, callback_query.from_user.id): + await callback_query.answer("Доступ запрещён.", show_alert=True) + return confirm_text = FREEZE_SUBSCRIPTION_CONFIRM_MSG @@ -172,6 +184,9 @@ async def process_callback_freeze_subscription_confirm(callback_query: CallbackQ if not record: await callback_query.message.answer("Ключ не найден.") return + if not key_owned_by_user(record, callback_query.from_user.id): + await callback_query.answer("Доступ запрещён.", show_alert=True) + return email = record["email"] client_id = record["client_id"] diff --git a/handlers/keys/key_mode/key_country_mode.py b/handlers/keys/key_mode/key_country_mode.py index dca7a80b..81d3c839 100644 --- a/handlers/keys/key_mode/key_country_mode.py +++ b/handlers/keys/key_mode/key_country_mode.py @@ -264,6 +264,9 @@ async def change_location_callback(callback_query: CallbackQuery, session: Any): if not record: await callback_query.answer("❌ Ключ не найден", show_alert=True) return + if record.get("tg_id") != callback_query.from_user.id: + await callback_query.answer("Доступ запрещён.", show_alert=True) + return expiry_timestamp = record["expiry_time"] ts = int(expiry_timestamp / 1000) diff --git a/handlers/keys/key_renew.py b/handlers/keys/key_renew.py index fd1a5191..f6ce7ddc 100644 --- a/handlers/keys/key_renew.py +++ b/handlers/keys/key_renew.py @@ -56,7 +56,7 @@ from hooks.processors import ( ) from logger import logger -from .utils import add_tariff_button_generic +from .utils import add_tariff_button_generic, key_owned_by_user router = Router() @@ -86,6 +86,9 @@ async def process_callback_renew_key(callback_query: CallbackQuery, state: FSMCo if not record: await callback_query.message.answer("Ключ не найден.") return + if not key_owned_by_user(record, callback_query.from_user.id): + await callback_query.answer("Доступ запрещён.", show_alert=True) + return client_id = record["client_id"] expiry_time_raw = record["expiry_time"] diff --git a/handlers/keys/key_view.py b/handlers/keys/key_view.py index 99631528..01afe332 100644 --- a/handlers/keys/key_view.py +++ b/handlers/keys/key_view.py @@ -62,6 +62,7 @@ from handlers.texts import ( RENAME_KEY_PROMPT, key_message, ) +from handlers.keys.utils import key_owned_by_user from handlers.utils import ( edit_or_send_message, format_days, @@ -227,8 +228,12 @@ async def build_keys_response(records: list[Key] | None, session: AsyncSession, @router.callback_query(F.data.startswith("rename_key|")) -async def handle_rename_key(callback: CallbackQuery, state: FSMContext): +async def handle_rename_key(callback: CallbackQuery, state: FSMContext, session: AsyncSession): client_id = callback.data.split("|")[1] + key_row = (await session.execute(select(Key).where(Key.client_id == client_id))).scalar_one_or_none() + if not key_row or key_row.tg_id != callback.from_user.id: + await callback.answer("Доступ запрещён.", show_alert=True) + return await state.set_state(RenameKeyState.waiting_for_new_alias) await state.update_data(client_id=client_id) @@ -282,6 +287,10 @@ async def handle_new_alias_input(message: Message, state: FSMContext, session: A @router.callback_query(F.data.startswith("view_key|")) async def process_callback_view_key(callback_query: CallbackQuery, session: AsyncSession): key_name = callback_query.data.split("|")[1] + record = await get_key_details(session, key_name) + if not key_owned_by_user(record, callback_query.from_user.id): + await callback_query.answer("Доступ запрещён.", show_alert=True) + return image_path = os.path.join("img", "pic_view.jpg") await render_key_info(callback_query.message, session, key_name, image_path) @@ -482,6 +491,9 @@ async def handle_reset_hwid(callback_query: CallbackQuery, session: AsyncSession if not record: await callback_query.answer("❌ Ключ не найден.", show_alert=True) return + if not key_owned_by_user(record, callback_query.from_user.id): + await callback_query.answer("Доступ запрещён.", show_alert=True) + return client_id = record.get("client_id") if not client_id: diff --git a/handlers/keys/keys.py b/handlers/keys/keys.py index 8bf69301..5f108ec9 100644 --- a/handlers/keys/keys.py +++ b/handlers/keys/keys.py @@ -7,6 +7,7 @@ from database import delete_key, get_key_details from handlers.buttons import APPLY, BACK, CANCEL from handlers.keys.key_view import process_callback_view_key from handlers.keys.operations import delete_key_from_cluster, update_subscription +from handlers.keys.utils import key_owned_by_user from handlers.texts import DELETE_KEY_CONFIRM_MSG, KEY_DELETED_MSG_SIMPLE from handlers.utils import edit_or_send_message, handle_error from middlewares.session import release_session_early @@ -22,6 +23,10 @@ async def process_callback_update_subscription(callback_query: CallbackQuery, se email = callback_query.data.split("|")[1] try: + record = await get_key_details(session, email) + if not key_owned_by_user(record, callback_query.from_user.id): + await callback_query.answer("Доступ запрещён.", show_alert=True) + return try: await callback_query.message.delete() except TelegramBadRequest as e: @@ -36,12 +41,16 @@ async def process_callback_update_subscription(callback_query: CallbackQuery, se @router.callback_query(F.data.startswith("delete_key|")) -async def process_callback_delete_key(callback_query: CallbackQuery): - client_id = callback_query.data.split("|")[1] +async def process_callback_delete_key(callback_query: CallbackQuery, session: AsyncSession): + key_identifier = callback_query.data.split("|")[1] try: + record = await get_key_details(session, key_identifier) + if not key_owned_by_user(record, callback_query.from_user.id): + await callback_query.answer("Доступ запрещён.", show_alert=True) + return confirmation_keyboard = types.InlineKeyboardMarkup( inline_keyboard=[ - [types.InlineKeyboardButton(text=APPLY, callback_data=f"confirm_delete|{client_id}")], + [types.InlineKeyboardButton(text=APPLY, callback_data=f"confirm_delete|{key_identifier}")], [types.InlineKeyboardButton(text=CANCEL, callback_data="view_keys")], ] ) @@ -54,7 +63,7 @@ async def process_callback_delete_key(callback_query: CallbackQuery): await callback_query.message.edit_text(text=DELETE_KEY_CONFIRM_MSG, reply_markup=confirmation_keyboard) except Exception as e: - logger.error(f"Ошибка при подготовке удаления ключа {client_id}: {e}") + logger.error(f"Ошибка при подготовке удаления ключа {key_identifier}: {e}") @router.callback_query(F.data.startswith("confirm_delete|")) @@ -62,6 +71,9 @@ async def process_callback_confirm_delete(callback_query: CallbackQuery, session email = callback_query.data.split("|")[1] try: record = await get_key_details(session, email) + if not key_owned_by_user(record, callback_query.from_user.id): + await callback_query.answer("Доступ запрещён.", show_alert=True) + return if record: client_id = record["client_id"] server_id = record["server_id"] diff --git a/handlers/keys/operations/traffic.py b/handlers/keys/operations/traffic.py index cd7c5f47..55a8635e 100644 --- a/handlers/keys/operations/traffic.py +++ b/handlers/keys/operations/traffic.py @@ -2,7 +2,7 @@ import asyncio from typing import Any -from sqlalchemy import select +from sqlalchemy import or_, select from sqlalchemy.ext.asyncio import AsyncSession from config import SUPERNODE @@ -21,37 +21,37 @@ async def get_user_traffic(session: AsyncSession, tg_id: int, email: str) -> dic """ Получает трафик пользователя на всех серверах, где у него есть ключ (3x-ui и Remnawave). Для Remnawave трафик считается один раз и отображается как "Remnawave (общий):". + Один запрос: Key + Server через join. """ - result = await session.execute(select(Key.client_id, Key.server_id).where(Key.tg_id == tg_id, Key.email == email)) + join_cond = or_( + Key.server_id == Server.server_name, + Key.server_id == Server.cluster_name, + ) + result = await session.execute( + select(Key.client_id, Key.server_id, Server) + .select_from(Key) + .join(Server, join_cond) + .where(Server.enabled.is_(True), Key.tg_id == tg_id, Key.email == email) + ) rows = result.all() if not rows: return {"status": "error", "message": "У пользователя нет активных ключей."} - server_ids = {row.server_id for row in rows} - server_id = list(server_ids)[0] - result = await session.execute( - select(Server) - .where(Server.enabled.is_(True)) - .where(Server.server_name.in_(server_ids) | Server.cluster_name.in_(server_ids)) - ) - server_rows = result.scalars().all() - if not server_rows: - logger.error(f"Не найдено серверов для: {server_ids}") - return { - "status": "error", - "message": f"Серверы не найдены: {', '.join(server_ids)}", - } - - servers_map = { - s.server_name: { - "server_name": s.server_name, - "cluster_name": s.cluster_name, - "api_url": s.api_url, - "panel_type": s.panel_type, - } - for s in server_rows - } + seen_pairs = set() + unique_rows = [] + servers_map = {} + for client_id, server_id, server in rows: + if (client_id, server_id) not in seen_pairs: + seen_pairs.add((client_id, server_id)) + unique_rows.append((client_id, server_id)) + if server.server_name not in servers_map: + servers_map[server.server_name] = { + "server_name": server.server_name, + "cluster_name": server.cluster_name, + "api_url": server.api_url, + "panel_type": server.panel_type, + } user_traffic_data = {} tasks = [] @@ -80,10 +80,7 @@ async def get_user_traffic(session: AsyncSession, tg_id: int, email: str) -> dic except Exception as e: return server_name, f"Ошибка: {e}" - for row in rows: - client_id = row.client_id - server_id = row.server_id - + for client_id, server_id in unique_rows: matched_servers = [ s for s in servers_map.values() if s["server_name"] == server_id or s["cluster_name"] == server_id ] diff --git a/handlers/keys/utils.py b/handlers/keys/utils.py index 041ee1b8..198df76b 100644 --- a/handlers/keys/utils.py +++ b/handlers/keys/utils.py @@ -7,6 +7,11 @@ from sqlalchemy.ext.asyncio import AsyncSession from handlers.payments.currency_rates import format_for_user +def key_owned_by_user(record: dict | None, user_id: int) -> bool: + """Проверка, что ключ принадлежит пользователю (защита от пересылки callback).""" + return record is not None and record.get("tg_id") == user_id + + async def add_tariff_button_generic( builder: InlineKeyboardBuilder, tariff: dict[str, Any], diff --git a/handlers/notifications/general_notifications.py b/handlers/notifications/general_notifications.py index d2524288..79228ba9 100644 --- a/handlers/notifications/general_notifications.py +++ b/handlers/notifications/general_notifications.py @@ -34,6 +34,7 @@ from database import ( get_all_keys, get_balance, get_last_notification_time, + get_last_notification_times_bulk, update_balance, update_key_expiry, update_key_tariff, @@ -583,6 +584,9 @@ async def handle_expired_keys(ctx: NotificationContext, keys: list): notify_delete_key_enabled = bool(NOTIFICATIONS_CONFIG.get("DELETE_KEY_ENABLED", NOTIFY_DELETE_KEY)) delete_key_delay_minutes = int(NOTIFICATIONS_CONFIG.get("DELETE_KEY_DELAY_MINUTES", NOTIFY_DELETE_DELAY)) + notification_pairs = [(key.tg_id, f"{key.email or ''}_key_expired") for key in expired_keys] + last_times = await get_last_notification_times_bulk(ctx.session, notification_pairs) + for key in expired_keys: tg_id = key.tg_id email = key.email or "" @@ -590,7 +594,7 @@ async def handle_expired_keys(ctx: NotificationContext, keys: list): server_id = key.server_id notification_id = f"{email}_key_expired" - last_notification_time = await get_last_notification_time(ctx.session, tg_id, notification_id) + last_notification_time = last_times.get((tg_id, notification_id)) if notify_renew_expired_enabled: try: @@ -700,7 +704,7 @@ async def periodic_notifications(bot: Bot, *, sessionmaker: async_sessionmaker): trial_time_disable = bool(MODES_CONFIG.get("TRIAL_TIME_DISABLED", TRIAL_TIME_DISABLE)) if not trial_time_disable: try: - await notify_inactive_trial_users(bot, session) + await notify_inactive_trial_users(bot, session, sessionmaker=sessionmaker) except Exception as error: logger.error(f"Ошибка в notify_inactive_trial_users: {error}") diff --git a/handlers/notifications/special_notifications.py b/handlers/notifications/special_notifications.py index 2dfa71ee..e476b0c4 100644 --- a/handlers/notifications/special_notifications.py +++ b/handlers/notifications/special_notifications.py @@ -6,7 +6,7 @@ from aiogram import Bot, Router, types from aiogram.types import InlineKeyboardButton, WebAppInfo from aiogram.utils.keyboard import InlineKeyboardBuilder from sqlalchemy import update -from sqlalchemy.ext.asyncio import AsyncSession +from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker from config import ( NOTIFY_EXTRA_DAYS, @@ -38,7 +38,9 @@ router = Router() moscow_tz = pytz.timezone("Europe/Moscow") -async def notify_inactive_trial_users(bot: Bot, session: AsyncSession): +async def notify_inactive_trial_users( + bot: Bot, session: AsyncSession, *, sessionmaker: async_sessionmaker | None = None +): logger.info("Проверка пользователей, не активировавших пробный период...") inactive_hours = int(NOTIFICATIONS_CONFIG.get("INACTIVE_USER_ENABLED", NOTIFY_INACTIVE)) @@ -119,8 +121,13 @@ async def notify_inactive_trial_users(bot: Bot, session: AsyncSession): sent_tg_ids.append(msg["tg_id"]) if sent_tg_ids: - for tg_id in sent_tg_ids: - await add_notification(session, tg_id, "inactive_trial") + if sessionmaker is not None: + async with sessionmaker() as fresh_session: + for tg_id in sent_tg_ids: + await add_notification(fresh_session, tg_id, "inactive_trial") + else: + for tg_id in sent_tg_ids: + await add_notification(session, tg_id, "inactive_trial") logger.info(f"Отправлено {len(sent_tg_ids)} уведомлений неактивным пользователям.") logger.info("Проверка пользователей с неактивным пробным периодом завершена.") diff --git a/handlers/payments/freekassa/freekassa_pay.py b/handlers/payments/freekassa/freekassa_pay.py index a0ea0f1a..11ce96ce 100644 --- a/handlers/payments/freekassa/freekassa_pay.py +++ b/handlers/payments/freekassa/freekassa_pay.py @@ -16,6 +16,11 @@ from config import ( FREEKASSA_SECRET2, FREEKASSA_SHOP_ID, ) +from core.webhook_abuse import ( + get_webhook_client_ip, + is_webhook_ip_blocked, + record_webhook_signature_failure, +) from database import ( add_payment, add_user, @@ -25,6 +30,7 @@ from database import ( get_key_count, get_payment_by_payment_id, get_temporary_data, + invalidate_payment_cache, update_balance, ) from handlers.buttons import BACK, PAY_2 @@ -198,6 +204,9 @@ def verify_signature(params: dict) -> bool: async def freekassa_webhook(request: web.Request): try: + ip = get_webhook_client_ip(request) + if await is_webhook_ip_blocked(ip): + return web.Response(status=429) params = dict(request.query) logger.info(f"Received Freekassa webhook: {params}") @@ -213,6 +222,7 @@ async def freekassa_webhook(request: web.Request): if not verify_signature(params): logger.error("Invalid signature in webhook") + await record_webhook_signature_failure(ip) return web.Response(status=400, text="Invalid signature") if str(merchant_id) != str(FREEKASSA_SHOP_ID): @@ -244,6 +254,7 @@ async def freekassa_webhook(request: web.Request): await send_payment_success_notification(tg_id_int, amount_float, session) await add_payment(session, tg_id_int, amount_float, "freekassa", payment_id=merchant_order_id) await clear_temporary_data(session, tg_id_int) + await invalidate_payment_cache(merchant_order_id) logger.info(f"Payment processed successfully. User: {tg_id_int}, Amount: {amount_float}") return web.Response(text="YES") diff --git a/handlers/payments/heleket/service.py b/handlers/payments/heleket/service.py index 7b9e9c6b..4da2c8c5 100644 --- a/handlers/payments/heleket/service.py +++ b/handlers/payments/heleket/service.py @@ -21,7 +21,7 @@ from config import ( HELEKET_SUCCESS_URL, PROVIDERS_ENABLED, ) -from database import add_payment, async_session_maker +from database import async_session_maker, register_pending_payment from database.models import User from handlers.buttons import BACK, HELEKET, PAY_2 from handlers.payments.currency_rates import ( @@ -240,7 +240,8 @@ async def handle_custom_amount_input(message: types.Message, state: FSMContext, if currency == "RUB": amount_rub = user_amount else: - async with aiohttp.ClientSession() as session_http: + timeout = aiohttp.ClientTimeout(total=30, connect=10) + async with aiohttp.ClientSession(timeout=timeout) as session_http: amount_rub = int(await to_rub(user_amount, "USD", session=session_http)) await state.update_data(amount=amount_rub) @@ -336,8 +337,9 @@ async def generate_heleket_payment_link( unique_order_id = f"{int(time.time())}_{tg_id}" db_session = session + timeout = aiohttp.ClientTimeout(total=30, connect=10) try: - async with aiohttp.ClientSession() as http_session: + async with aiohttp.ClientSession(timeout=timeout) as http_session: pay_cur = str(method["currency"]).upper() if pay_cur == "RUB": @@ -376,27 +378,13 @@ async def generate_heleket_payment_link( if resp_json.get("state") == 0: payment_url = resp_json.get("result", {}).get("url") if payment_url: - if db_session is not None: - await add_payment( - session=db_session, - tg_id=tg_id, - amount=float(amount), - payment_system="HELEKET", - status="pending", - currency="RUB", - payment_id=unique_order_id, - ) - else: - async with async_session_maker() as dbs: - await add_payment( - session=dbs, - tg_id=tg_id, - amount=float(amount), - payment_system="HELEKET", - status="pending", - currency="RUB", - payment_id=unique_order_id, - ) + await register_pending_payment( + payment_id=unique_order_id, + tg_id=tg_id, + amount=float(amount), + payment_system="heleket", + currency="RUB", + ) logger.info(f"Heleket payment URL created for user {tg_id}") return payment_url else: diff --git a/handlers/payments/heleket/webhook.py b/handlers/payments/heleket/webhook.py index e1cb5951..dc2657b4 100644 --- a/handlers/payments/heleket/webhook.py +++ b/handlers/payments/heleket/webhook.py @@ -5,10 +5,16 @@ import json from aiohttp import web from config import HELEKET_API_KEY +from core.webhook_abuse import ( + get_webhook_client_ip, + is_webhook_ip_blocked, + record_webhook_signature_failure, +) from database import ( add_payment, async_session_maker, get_payment_by_payment_id, + invalidate_payment_cache, update_balance, update_payment_status, ) @@ -106,18 +112,30 @@ async def process_heleket_webhook(data: dict) -> bool: if payment.get("status") == "success": logger.info(f"Heleket: платёж {order_id} уже обработан") return True - ok = await update_payment_status( - session=session, internal_id=int(payment["id"]), new_status="success" - ) - if not ok: - logger.error(f"Heleket: не удалось обновить статус платежа {order_id}") - return False + if payment.get("id") is not None: + ok = await update_payment_status( + session=session, internal_id=int(payment["id"]), new_status="success" + ) + if not ok: + logger.error(f"Heleket: не удалось обновить статус платежа {order_id}") + return False + else: + await add_payment( + session=session, + tg_id=tg_id, + amount=balance_amount, + payment_system="HELEKET", + status="success", + currency="USD", + payment_id=order_id, + metadata=None, + ) else: await add_payment( session=session, tg_id=tg_id, amount=balance_amount, - payment_system="HELEKET", + payment_system="heleket", status="success", currency="USD", payment_id=order_id, @@ -126,6 +144,7 @@ async def process_heleket_webhook(data: dict) -> bool: await update_balance(session, tg_id, balance_amount) await send_payment_success_notification(tg_id, balance_amount, session) + await invalidate_payment_cache(order_id) logger.info( f"Heleket: платёж {order_id} для пользователя {tg_id} " f"успешно обработан, баланс пополнен на {balance_amount} RUB" @@ -136,13 +155,14 @@ async def process_heleket_webhook(data: dict) -> bool: async with async_session_maker() as session: payment = await get_payment_by_payment_id(session, order_id) - if payment: + if payment and payment.get("id") is not None: await update_payment_status( session=session, internal_id=int(payment["id"]), new_status="failed", ) await session.commit() + await invalidate_payment_cache(order_id) return True else: logger.info(f"Heleket: промежуточный статус {status} для платежа {order_id}") @@ -155,11 +175,15 @@ async def process_heleket_webhook(data: dict) -> bool: async def heleket_webhook(request: web.Request): """Обработчик webhook от Heleket для aiohttp.""" try: + ip = get_webhook_client_ip(request) + if await is_webhook_ip_blocked(ip): + return web.Response(status=429) data = await request.json() logger.info(f"Heleket webhook received from {request.remote}") if not verify_heleket_signature(data): logger.error("Heleket webhook: неверная подпись") + await record_webhook_signature_failure(ip) return web.Response(status=400, text="Invalid signature") success = await process_heleket_webhook(data) diff --git a/handlers/payments/kassai/service.py b/handlers/payments/kassai/service.py index 9a794540..6d2cb0cb 100644 --- a/handlers/payments/kassai/service.py +++ b/handlers/payments/kassai/service.py @@ -20,7 +20,7 @@ from config import ( KASSAI_SUCCESS_URL, PROVIDERS_ENABLED, ) -from database import add_payment, async_session_maker +from database import async_session_maker, register_pending_payment from database.models import User from handlers.buttons import BACK, KASSAI_CARDS, KASSAI_SBP, PAY_2 from handlers.payments.currency_rates import ( @@ -263,7 +263,8 @@ async def handle_custom_amount_input(message: types.Message, state: FSMContext, if currency == "RUB": amount_rub = user_amount else: - async with aiohttp.ClientSession() as session_http: + timeout = aiohttp.ClientTimeout(total=30, connect=10) + async with aiohttp.ClientSession(timeout=timeout) as session_http: amount_rub = int(await to_rub(user_amount, "USD", session=session_http)) await state.update_data(amount=amount_rub) @@ -390,8 +391,9 @@ async def generate_kassai_payment_link( db_session = session + timeout = aiohttp.ClientTimeout(total=60, connect=10) try: - async with aiohttp.ClientSession() as http_session: + async with aiohttp.ClientSession(timeout=timeout) as http_session: async with http_session.post(url, headers=headers, json=data, timeout=60) as resp: if resp.status == 200: try: @@ -399,27 +401,13 @@ async def generate_kassai_payment_link( if resp_json.get("type") == "success": payment_url = resp_json.get("location") if payment_url: - if db_session is not None: - await add_payment( - session=db_session, - tg_id=tg_id, - amount=float(amount), - payment_system="KASSAI", - status="pending", - currency="RUB", - payment_id=unique_payment_id, - ) - else: - async with async_session_maker() as dbs: - await add_payment( - session=dbs, - tg_id=tg_id, - amount=float(amount), - payment_system="KASSAI", - status="pending", - currency="RUB", - payment_id=unique_payment_id, - ) + await register_pending_payment( + payment_id=unique_payment_id, + tg_id=tg_id, + amount=float(amount), + payment_system="kassai", + currency="RUB", + ) logger.info(f"KassaAI payment URL created for user {tg_id}") return payment_url logger.error(f"KassaAI: No location in response: {resp_json}") diff --git a/handlers/payments/kassai/webhook.py b/handlers/payments/kassai/webhook.py index 6c1f6dcb..a330e7f1 100644 --- a/handlers/payments/kassai/webhook.py +++ b/handlers/payments/kassai/webhook.py @@ -3,10 +3,16 @@ import hashlib from aiohttp import web from config import KASSAI_SECRET_KEY, KASSAI_SHOP_ID, KASSAI_WEBHOOK_RESPONSE +from core.webhook_abuse import ( + get_webhook_client_ip, + is_webhook_ip_blocked, + record_webhook_signature_failure, +) from database import ( add_payment, async_session_maker, get_payment_by_payment_id, + invalidate_payment_cache, update_balance, update_payment_status, ) @@ -44,14 +50,19 @@ def verify_kassai_signature(data: dict, signature: str) -> bool: async def kassai_webhook(request: web.Request): """Обработчик webhook от KassaAI для aiohttp.""" try: + ip = get_webhook_client_ip(request) + if await is_webhook_ip_blocked(ip): + return web.Response(status=429) data = await request.post() logger.info(f"KassaAI webhook received: {dict(data)}") signature = data.get("SIGN", "") if not signature: logger.error("KassaAI webhook: отсутствует подпись") + await record_webhook_signature_failure(ip) return web.Response(status=400) if not verify_kassai_signature(data, signature): logger.error("KassaAI webhook: неверная подпись") + await record_webhook_signature_failure(ip) return web.Response(status=400) amount_raw = data.get("AMOUNT") @@ -77,10 +88,24 @@ async def kassai_webhook(request: web.Request): if payment.get("status") == "success": logger.info(f"KassaAI: платёж {order_id} уже обработан") return web.Response(text=KASSAI_WEBHOOK_RESPONSE) - ok = await update_payment_status(session=session, internal_id=int(payment["id"]), new_status="success") - if not ok: - logger.error(f"KassaAI: не удалось обновить статус платежа {order_id}") - return web.Response(status=500) + if payment.get("id") is not None: + ok = await update_payment_status( + session=session, internal_id=int(payment["id"]), new_status="success" + ) + if not ok: + logger.error(f"KassaAI: не удалось обновить статус платежа {order_id}") + return web.Response(status=500) + else: + await add_payment( + session=session, + tg_id=tg_id, + amount=amount, + payment_system="kassai", + status="success", + currency="RUB", + payment_id=order_id, + metadata=None, + ) else: await add_payment( session=session, @@ -95,6 +120,7 @@ async def kassai_webhook(request: web.Request): await update_balance(session, tg_id, amount) await send_payment_success_notification(tg_id, amount, session) + await invalidate_payment_cache(order_id) logger.info( f"KassaAI: платёж {order_id} успешно обработан, баланс пользователя {tg_id} пополнен на {amount} RUB" ) diff --git a/handlers/payments/robokassa/service.py b/handlers/payments/robokassa/service.py index 3c1525a6..821f7bfc 100644 --- a/handlers/payments/robokassa/service.py +++ b/handlers/payments/robokassa/service.py @@ -8,7 +8,7 @@ from urllib.parse import quote_plus, urlencode from sqlalchemy.ext.asyncio import AsyncSession from config import ROBOKASSA_LOGIN, ROBOKASSA_PASSWORD1, ROBOKASSA_PASSWORD2, ROBOKASSA_TEST_MODE -from database import add_payment +from database import register_pending_payment from handlers.payments.payment_links import register_payment_creator @@ -61,15 +61,12 @@ async def create_and_store_robokassa_payment( session: AsyncSession, tg_id: int, amount: int | float, description: str, inv_id: int = 0 ) -> tuple[str, str]: url, pid = generate_payment_link(amount, inv_id, description, tg_id) - await add_payment( - session=session, + await register_pending_payment( + payment_id=pid, tg_id=tg_id, amount=float(amount), payment_system="robokassa", - status="pending", currency="RUB", - payment_id=pid, - metadata=None, ) return url, pid diff --git a/handlers/payments/robokassa/webhook.py b/handlers/payments/robokassa/webhook.py index 856c9d49..89c69cee 100644 --- a/handlers/payments/robokassa/webhook.py +++ b/handlers/payments/robokassa/webhook.py @@ -1,6 +1,18 @@ from aiohttp import web -from database import add_payment, async_session_maker, get_payment_by_payment_id, update_balance, update_payment_status +from core.webhook_abuse import ( + get_webhook_client_ip, + is_webhook_ip_blocked, + record_webhook_signature_failure, +) +from database import ( + add_payment, + async_session_maker, + get_payment_by_payment_id, + invalidate_payment_cache, + update_balance, + update_payment_status, +) from handlers.payments.utils import send_payment_success_notification from logger import logger @@ -9,8 +21,12 @@ from .service import check_payment_signature async def robokassa_webhook(request: web.Request): try: + ip = get_webhook_client_ip(request) + if await is_webhook_ip_blocked(ip): + return web.Response(status=429) params = await request.post() if not check_payment_signature(params): + await record_webhook_signature_failure(ip) return web.Response(status=400) amount_raw = params.get("OutSum") @@ -29,15 +45,29 @@ async def robokassa_webhook(request: web.Request): if payment: if payment.get("status") == "success": return web.Response(text=f"OK{inv_id}") - ok = await update_payment_status(session=session, internal_id=int(payment["id"]), new_status="success") - if not ok: - return web.Response(status=500) + if payment.get("id") is not None: + ok = await update_payment_status( + session=session, internal_id=int(payment["id"]), new_status="success" + ) + if not ok: + return web.Response(status=500) + else: + await add_payment( + session=session, + tg_id=tg_id, + amount=amount, + payment_system="robokassa", + status="success", + currency="RUB", + payment_id=shp_pid, + metadata=None, + ) else: await add_payment( session=session, tg_id=tg_id, amount=amount, - payment_system="ROBOKASSA", + payment_system="robokassa", status="success", currency="RUB", payment_id=shp_pid, @@ -46,6 +76,7 @@ async def robokassa_webhook(request: web.Request): await update_balance(session, tg_id, amount) await send_payment_success_notification(tg_id, amount, session) + await invalidate_payment_cache(shp_pid) return web.Response(text=f"OK{inv_id}") except Exception as e: diff --git a/handlers/payments/yookassa/service.cpython-312-x86_64-linux-gnu.so b/handlers/payments/yookassa/service.cpython-312-x86_64-linux-gnu.so index 8383cc89..09901a6d 100644 Binary files a/handlers/payments/yookassa/service.cpython-312-x86_64-linux-gnu.so and b/handlers/payments/yookassa/service.cpython-312-x86_64-linux-gnu.so differ diff --git a/handlers/payments/yookassa/webhook.cpython-312-x86_64-linux-gnu.so b/handlers/payments/yookassa/webhook.cpython-312-x86_64-linux-gnu.so index b8c7431f..d6e88780 100644 Binary files a/handlers/payments/yookassa/webhook.cpython-312-x86_64-linux-gnu.so and b/handlers/payments/yookassa/webhook.cpython-312-x86_64-linux-gnu.so differ diff --git a/handlers/profile.py b/handlers/profile.py index c5827b56..528f3d70 100644 --- a/handlers/profile.py +++ b/handlers/profile.py @@ -15,7 +15,9 @@ from config import ( TRIAL_TIME_DISABLE, ) from core.bootstrap import BUTTONS_CONFIG, MODES_CONFIG -from database import get_balance, get_key_count, get_trial +from core.cache_config import BALANCE_CACHE_TTL_SEC, KEY_COUNT_CACHE_TTL_SEC, PROFILE_DATA_CACHE_TTL_SEC +from core.redis_cache import cache_get, cache_key, cache_set +from database import get_balance_trial_key_count from handlers.buttons import ( ABOUT_VPN, ADD_SUB, @@ -33,6 +35,7 @@ from handlers.payments.currency_rates import format_for_user from handlers.texts import ADD_SUBSCRIPTION_HINT from hooks.hook_buttons import insert_hook_buttons from hooks.hooks import run_hooks +from middlewares.session import release_session_early from .admin.panel.keyboard import AdminPanelCallback from .texts import profile_message_send @@ -62,10 +65,21 @@ async def process_callback_view_profile( chat_id = chat.id username = get_username(user or chat) - key_count = await get_key_count(session, chat_id) - balance_rub = await get_balance(session, chat_id) - trial_status = await get_trial(session, chat_id) - balance_rub = balance_rub or 0 + cached = await cache_get(cache_key("profile_data", chat_id)) + if isinstance(cached, dict) and "key_count" in cached and "balance_rub" in cached and "trial_status" in cached: + key_count = int(cached["key_count"]) + balance_rub = float(cached.get("balance_rub") or 0) + trial_status = int(cached.get("trial_status") or 0) + else: + balance_rub, trial_status, key_count = await get_balance_trial_key_count(session, chat_id) + balance_rub = balance_rub or 0 + await cache_set(cache_key("balance", chat_id), balance_rub, BALANCE_CACHE_TTL_SEC) + await cache_set(cache_key("key_count", chat_id), key_count, KEY_COUNT_CACHE_TTL_SEC) + await cache_set( + cache_key("profile_data", chat_id), + {"key_count": key_count, "balance_rub": balance_rub, "trial_status": trial_status}, + PROFILE_DATA_CACHE_TTL_SEC, + ) balance_text = await format_for_user( session, @@ -134,6 +148,7 @@ async def process_callback_view_profile( else: builder.row(InlineKeyboardButton(text=BACK, callback_data="start")) + await release_session_early(session) await edit_or_send_message( target_message=message, text=profile_message, diff --git a/handlers/start.py b/handlers/start.py index ee253d3e..6aebb2b5 100644 --- a/handlers/start.py +++ b/handlers/start.py @@ -155,9 +155,20 @@ async def process_start_logic( await state.update_data(original_text=text, user_data=user_data) + _MAX_START_PAYLOAD_LEN = 256 + _MAX_START_PARTS = 20 + if text and len(text) > _MAX_START_PAYLOAD_LEN: + text = text[:_MAX_START_PAYLOAD_LEN] + parts = text.split("-") if text else [] + if len(parts) > _MAX_START_PARTS: + parts = parts[:_MAX_START_PARTS] + gift_detected = False - if text: - for part in text.split("-"): + if parts: + for part in parts: + part = part.strip() + if not part: + continue 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) @@ -171,6 +182,8 @@ async def process_start_logic( if "utm" in part: await handle_utm_link(part, message, state, session, user_data) + text = "-".join(parts) if parts else (text or "") + await state.clear() if gift_detected: return diff --git a/handlers/tariffs/addons/key_addons_main.py b/handlers/tariffs/addons/key_addons_main.py index 10347c03..b926daf1 100644 --- a/handlers/tariffs/addons/key_addons_main.py +++ b/handlers/tariffs/addons/key_addons_main.py @@ -318,6 +318,9 @@ async def start_key_addons(callback: CallbackQuery, state: FSMContext, session: logger.warning(f"[ADDONS] Подписка {email} не найдена") await callback.message.answer("❌ Подписка не найдена.") return + if record.get("tg_id") != callback.from_user.id: + await callback.answer("Доступ запрещён.", show_alert=True) + return tariff_id = record.get("tariff_id") if not tariff_id: diff --git a/handlers/tariffs/addons/key_addons_pack.py b/handlers/tariffs/addons/key_addons_pack.py index 77879424..85d10cd1 100644 --- a/handlers/tariffs/addons/key_addons_pack.py +++ b/handlers/tariffs/addons/key_addons_pack.py @@ -396,6 +396,9 @@ async def start_key_addons(callback: CallbackQuery, state: FSMContext, session: logger.warning(f"[ADDONS] PACK_MODE: подписка {email} не найдена") await callback.message.answer("❌ Подписка не найдена.") return + if record.get("tg_id") != callback.from_user.id: + await callback.answer("Доступ запрещён.", show_alert=True) + return tariff_id = record.get("tariff_id") if not tariff_id: diff --git a/hooks/constants.py b/hooks/constants.py index 08d99803..796746cb 100644 --- a/hooks/constants.py +++ b/hooks/constants.py @@ -1 +1 @@ -DEFAULT_HOOK_TIMEOUT = 4.0 +DEFAULT_HOOK_TIMEOUT = 8.0 diff --git a/middlewares/ban_checker.py b/middlewares/ban_checker.py index 88e727a6..32aabaec 100644 --- a/middlewares/ban_checker.py +++ b/middlewares/ban_checker.py @@ -10,7 +10,7 @@ from sqlalchemy.ext.asyncio import AsyncSession from config import ADMIN_ID, SUPPORT_CHAT_URL from core.cache_config import BAN_CACHE_TTL_SEC -from core.redis_cache import cache_get, cache_key, cache_set +from core.redis_cache import cache_delete, cache_get, cache_key, cache_set from database import async_session_maker from database.models import ManualBan from logger import logger @@ -20,6 +20,11 @@ TZ = timezone("Europe/Moscow") _BAN_CACHE_TTL = BAN_CACHE_TTL_SEC +async def invalidate_ban_cache(tg_id: int) -> None: + """Сбросить кэш статуса бана после добавления/снятия бана.""" + await cache_delete(cache_key("ban_status", tg_id)) + + class BanCheckerMiddleware(BaseMiddleware): """Проверка банов.""" @@ -81,10 +86,14 @@ class BanCheckerMiddleware(BaseMiddleware): until_parsed = datetime.fromisoformat(until_raw) except ValueError: until_parsed = None - ban_info = { - "reason": cached.get("reason") or "не указана", - "until": until_parsed, - } + if until_parsed is not None and until_parsed < datetime.utcnow(): + ban_info = None + await cache_delete(cache_key("ban_status", tg_id)) + else: + ban_info = { + "reason": cached.get("reason") or "не указана", + "until": until_parsed, + } else: session = data.get("session") if session is not None and getattr(session, "execute", None) is not None: diff --git a/middlewares/loggings.py b/middlewares/loggings.py index fbb41692..cfd60b99 100644 --- a/middlewares/loggings.py +++ b/middlewares/loggings.py @@ -1,3 +1,4 @@ +import asyncio from collections.abc import Awaitable, Callable from typing import Any, TypedDict @@ -13,8 +14,18 @@ class UserInfo(TypedDict): action: str | None +def _log_activity_sync(user_info: UserInfo) -> None: + """Синхронный вывод в лог, чтобы не блокировать event loop в create_task.""" + logger.info( + f"Активность пользователя │ " + f"ID: {str(user_info['user_id']).ljust(10)} │ " + f"Имя: {user_info['username'] or '—':<15} │ " + f"Действие: {user_info['action'] or '—'}" + ) + + class LoggingMiddleware(BaseMiddleware): - """Middleware для логирования действий пользователя.""" + """Middleware для логирования действий пользователя. Лог пишется в фоне, не задерживая обработчик.""" async def __call__( self, @@ -25,12 +36,7 @@ class LoggingMiddleware(BaseMiddleware): user_info = self._extract_user_info(event) if user_info["user_id"]: - logger.info( - f"Активность пользователя │ " - f"ID: {str(user_info['user_id']).ljust(10)} │ " - f"Имя: {user_info['username'] or '—':<15} │ " - f"Действие: {user_info['action'] or '—'}" - ) + asyncio.create_task(asyncio.to_thread(_log_activity_sync, user_info)) return await handler(event, data) diff --git a/middlewares/session.py b/middlewares/session.py index d421c04c..54427476 100644 --- a/middlewares/session.py +++ b/middlewares/session.py @@ -115,6 +115,10 @@ class SessionMiddleware(BaseMiddleware): t0 = time.perf_counter() if LOG_SESSION_DURATION else None async with self.sessionmaker() as session: + try: + await session.rollback() + except Exception: + pass proxy = _SessionProxy(session, self.sessionmaker, data) data["session"] = proxy committed = False @@ -160,7 +164,7 @@ class SessionMiddleware(BaseMiddleware): rolled_back = True raise finally: - if not committed and not rolled_back and not data.get("_session_released_early"): + if not committed and not data.get("_session_released_early"): try: await session.rollback() except Exception: diff --git a/panels/remnawave_runtime.py b/panels/remnawave_runtime.py index 3b9302e4..2d61e5ac 100644 --- a/panels/remnawave_runtime.py +++ b/panels/remnawave_runtime.py @@ -9,9 +9,11 @@ from core.cache_config import ( REMNAWAVE_MAX_CONCURRENCY, REMNAWAVE_ACTION_TIMEOUT_SEC, REMNAWAVE_PROFILE_CACHE_TTL_SEC, + REMNAWAVE_PROFILE_ERROR_CACHE_TTL_SEC, REMNAWAVE_PROFILE_TIMEOUT_SEC, REMNAWAVE_SERVER_CACHE_TTL_SEC, ) +from core.executor import run_io from core.redis_cache import cache_delete_pattern, cache_get, cache_key, cache_set from database import get_servers from logger import logger @@ -20,6 +22,93 @@ from panels.remnawave import RemnawaveAPI _remnawave_semaphore = asyncio.Semaphore(REMNAWAVE_MAX_CONCURRENCY) +async def _fetch_profile_http_only(api_url: str, client_id: str) -> dict[str, Any] | None: + """Только HTTP к панели: логин + устройства + юзер. Без кэша и без resolve. Вызывается из потока.""" + api = RemnawaveAPI(api_url) + try: + logged_in = True + if not REMNAWAVE_TOKEN_LOGIN_ENABLED: + logged_in = await asyncio.wait_for( + api.login(REMNAWAVE_LOGIN, REMNAWAVE_PASSWORD), + timeout=REMNAWAVE_PROFILE_TIMEOUT_SEC, + ) + if not logged_in: + return None + devices = await asyncio.wait_for( + api.get_user_hwid_devices(client_id), + timeout=REMNAWAVE_PROFILE_TIMEOUT_SEC, + ) + user_data = await asyncio.wait_for( + api.get_user_by_uuid(client_id), + timeout=REMNAWAVE_PROFILE_TIMEOUT_SEC, + ) + hwid_count = len(devices or []) + used_gb = None + traffic_limit_bytes = None + hwid_device_limit = None + if user_data: + user_traffic = user_data.get("userTraffic", {}) + used_bytes = user_traffic.get("usedTrafficBytes", 0) + used_gb = round(used_bytes / 1073741824, 1) + traffic_limit_bytes = user_data.get("trafficLimitBytes") + hwid_device_limit = user_data.get("hwidDeviceLimit") + return { + "api_url": api_url, + "hwid_count": hwid_count, + "used_gb": used_gb, + "traffic_limit_bytes": traffic_limit_bytes, + "hwid_device_limit": hwid_device_limit, + } + except (asyncio.TimeoutError, Exception): + return None + finally: + if hasattr(api, "aclose"): + try: + await api.aclose() + except Exception: + pass + + +def _run_profile_http_in_thread(api_url: str, client_id: str) -> dict[str, Any] | None: + """Синхронная обёртка: свой event loop в потоке, чтобы не блокировать основной цикл бота.""" + loop = asyncio.new_event_loop() + asyncio.set_event_loop(loop) + try: + return loop.run_until_complete(_fetch_profile_http_only(api_url, client_id)) + finally: + loop.close() + + +def _run_with_api_in_thread( + api_url: str, + operation: Callable[[RemnawaveAPI], Awaitable[Any]], + timeout_sec: float, +) -> Any: + """Синхронная обёртка: логин + operation(api) в отдельном event loop в потоке.""" + loop = asyncio.new_event_loop() + asyncio.set_event_loop(loop) + api = RemnawaveAPI(api_url) + try: + logged_in = True + if not REMNAWAVE_TOKEN_LOGIN_ENABLED: + logged_in = loop.run_until_complete( + asyncio.wait_for(api.login(REMNAWAVE_LOGIN, REMNAWAVE_PASSWORD), timeout=timeout_sec) + ) + if not logged_in: + return None + coro = operation(api) + return loop.run_until_complete(asyncio.wait_for(coro, timeout=timeout_sec)) + except (asyncio.TimeoutError, Exception): + return None + finally: + if hasattr(api, "aclose"): + try: + loop.run_until_complete(api.aclose()) + except Exception: + pass + loop.close() + + def invalidate_remnawave_profile_cache(*, api_url: str | None = None, client_id: str | None = None) -> None: """Invalidate cached Remnawave profiles by api_url/client_id (or both).""" import asyncio @@ -96,63 +185,13 @@ async def get_remnawave_profile( if cached_profile is not None: return cached_profile - profile: dict[str, Any] | None = None - async with _remnawave_semaphore: - api = RemnawaveAPI(api_url) - try: - logged_in = True - if not REMNAWAVE_TOKEN_LOGIN_ENABLED: - logged_in = await asyncio.wait_for( - api.login(REMNAWAVE_LOGIN, REMNAWAVE_PASSWORD), - timeout=REMNAWAVE_PROFILE_TIMEOUT_SEC, - ) - if not logged_in: - await cache_set(pkey, None, REMNAWAVE_PROFILE_CACHE_TTL_SEC) - return None + profile = await run_io(_run_profile_http_in_thread, api_url, client_id) + if profile is None: + logger.warning(f"[Remnawave] Таймаут или ошибка профиля для client_id={client_id}") - devices = await asyncio.wait_for( - api.get_user_hwid_devices(client_id), - timeout=REMNAWAVE_PROFILE_TIMEOUT_SEC, - ) - user_data = await asyncio.wait_for( - api.get_user_by_uuid(client_id), - timeout=REMNAWAVE_PROFILE_TIMEOUT_SEC, - ) - - hwid_count = len(devices or []) - used_gb = None - traffic_limit_bytes = None - hwid_device_limit = None - - if user_data: - user_traffic = user_data.get("userTraffic", {}) - used_bytes = user_traffic.get("usedTrafficBytes", 0) - used_gb = round(used_bytes / 1073741824, 1) - traffic_limit_bytes = user_data.get("trafficLimitBytes") - hwid_device_limit = user_data.get("hwidDeviceLimit") - - profile = { - "api_url": api_url, - "hwid_count": hwid_count, - "used_gb": used_gb, - "traffic_limit_bytes": traffic_limit_bytes, - "hwid_device_limit": hwid_device_limit, - } - except asyncio.TimeoutError: - logger.warning(f"[Remnawave] Таймаут профиля для client_id={client_id}") - profile = None - except Exception as e: - logger.warning(f"[Remnawave] Ошибка профиля для client_id={client_id}: {e}") - profile = None - finally: - if hasattr(api, "aclose"): - try: - await api.aclose() - except Exception: - pass - - await cache_set(pkey, profile, REMNAWAVE_PROFILE_CACHE_TTL_SEC) + ttl = REMNAWAVE_PROFILE_CACHE_TTL_SEC if profile else REMNAWAVE_PROFILE_ERROR_CACHE_TTL_SEC + await cache_set(pkey, profile, ttl) return profile @@ -183,26 +222,7 @@ async def with_remnawave_api( return None async with _remnawave_semaphore: - api = RemnawaveAPI(api_url) - try: - logged_in = True - if not REMNAWAVE_TOKEN_LOGIN_ENABLED: - logged_in = await asyncio.wait_for( - api.login(REMNAWAVE_LOGIN, REMNAWAVE_PASSWORD), - timeout=timeout_sec, - ) - if not logged_in: - return None - return await asyncio.wait_for(operation(api), timeout=timeout_sec) - except asyncio.TimeoutError: - logger.warning(f"[Remnawave] Таймаут операции для server_ref={server_ref}") - return None - except Exception as e: - logger.warning(f"[Remnawave] Ошибка операции для server_ref={server_ref}: {e}") - return None - finally: - if hasattr(api, "aclose"): - try: - await api.aclose() - except Exception: - pass + result = await run_io(_run_with_api_in_thread, api_url, operation, timeout_sec) + if result is None: + logger.warning(f"[Remnawave] Таймаут или ошибка операции для server_ref={server_ref}") + return result