From 906601f5cdc504cc1a7cf01d7af123e5b7f0d5fb Mon Sep 17 00:00:00 2001 From: Zakhar Izmaylov Date: Thu, 23 Jan 2025 11:06:25 +0300 Subject: [PATCH] Refactor database interaction in key management functions to utilize session objects for improved performance and consistency. Updated `get_keys`, `get_keys_by_server`, `update_balance`, and `update_key_expiry` functions to accept a session parameter, allowing for better session management across handlers. Adjusted related handler functions to pass the session, enhancing code clarity and maintainability. --- database.py | 130 +++++++++------------------- handlers/admin/admin_servers.py | 15 +--- handlers/admin/admin_user_editor.py | 2 +- handlers/coupons.py | 2 +- handlers/keys/key_management.py | 2 +- handlers/keys/keys.py | 10 +-- handlers/notifications.py | 12 +-- 7 files changed, 57 insertions(+), 116 deletions(-) diff --git a/database.py b/database.py index 83e528c3..688f0ebf 100644 --- a/database.py +++ b/database.py @@ -349,7 +349,7 @@ async def store_key( raise -async def get_keys(tg_id: int): +async def get_keys(tg_id: int, session: Any): """ Получает список ключей для указанного пользователя. @@ -362,10 +362,8 @@ async def get_keys(tg_id: int): Raises: Exception: В случае ошибки при подключении к базе данных или выполнении запроса """ - conn = None try: - conn = await asyncpg.connect(DATABASE_URL) - records = await conn.fetch( + records = await session.fetch( """ SELECT client_id, email, created_at, key FROM keys @@ -378,17 +376,14 @@ async def get_keys(tg_id: int): except Exception as e: logger.error(f"Ошибка при получении ключей для пользователя {tg_id}: {e}") raise - finally: - if conn: - await conn.close() -async def get_keys_by_server(tg_id: int, server_id: str): +async def get_keys_by_server(tg_id: int | None, server_id: str, session: Any): """ - Получает список ключей для указанного пользователя на определенном сервере. + Получает список ключей на определенном сервере. Если tg_id=None, возвращает все ключи на сервере. Args: - tg_id (int): Telegram ID пользователя + tg_id (int | None): Telegram ID пользователя или None для всех пользователей server_id (str): Идентификатор сервера Returns: @@ -397,53 +392,36 @@ async def get_keys_by_server(tg_id: int, server_id: str): Raises: Exception: В случае ошибки при подключении к базе данных или выполнении запроса """ - conn = None try: - conn = await asyncpg.connect(DATABASE_URL) - records = await conn.fetch( - """ - SELECT client_id, email, created_at, key - FROM keys - WHERE tg_id = $1 AND server_id = $2 - """, - tg_id, - server_id, - ) - logger.info(f"Успешно получено {len(records)} ключей для пользователя {tg_id} на сервере {server_id}") + if tg_id is not None: + records = await session.fetch( + """ + SELECT * + FROM keys + WHERE tg_id = $1 AND server_id = $2 + """, + tg_id, + server_id, + ) + logger.info(f"Успешно получено {len(records)} ключей для пользователя {tg_id} на сервере {server_id}") + else: + records = await session.fetch( + """ + SELECT * + FROM keys + WHERE server_id = $1 + """, + server_id, + ) + logger.info(f"Успешно получено {len(records)} ключей на сервере {server_id}") + return records except Exception as e: - logger.error(f"Ошибка при получении ключей для пользователя {tg_id} на сервере {server_id}: {e}") + error_msg = f"Ошибка при получении ключей на сервере {server_id}" + if tg_id is not None: + error_msg += f" для пользователя {tg_id}" + logger.error(f"{error_msg}: {e}") raise - finally: - if conn: - await conn.close() - - -async def has_active_key(tg_id: int) -> bool: - """ - Проверяет наличие активных ключей для указанного пользователя. - - Args: - tg_id (int): Telegram ID пользователя - - Returns: - bool: True, если у пользователя есть активные ключи, иначе False - - Raises: - Exception: В случае ошибки при подключении к базе данных или выполнении запроса - """ - conn = None - try: - conn = await asyncpg.connect(DATABASE_URL) - count = await conn.fetchval("SELECT COUNT(*) FROM keys WHERE tg_id = $1", tg_id) - logger.info(f"Проверка наличия ключей для пользователя {tg_id}. Найдено ключей: {count}") - return count > 0 - except Exception as e: - logger.error(f"Ошибка при проверке наличия ключей для пользователя {tg_id}: {e}") - raise - finally: - if conn: - await conn.close() async def get_balance(tg_id: int) -> float: @@ -473,21 +451,25 @@ async def get_balance(tg_id: int) -> float: await conn.close() -async def update_balance(tg_id: int, amount: float): +async def update_balance(tg_id: int, amount: float, session: Any = None): """ Обновляет баланс пользователя в базе данных. Args: tg_id (int): Telegram ID пользователя amount (float): Сумма для обновления баланса + session (Any, optional): Сессия базы данных. Если не передана, создается новая. Raises: Exception: В случае ошибки при подключении к базе данных или обновлении баланса """ conn = None try: - conn = await asyncpg.connect(DATABASE_URL) - await conn.execute( + if session is None: + conn = await asyncpg.connect(DATABASE_URL) + session = conn + + await session.execute( """ UPDATE connections SET balance = balance + $1 @@ -504,7 +486,7 @@ async def update_balance(tg_id: int, amount: float): logger.error(f"Ошибка при обновлении баланса для пользователя {tg_id}: {e}") raise finally: - if conn: + if conn is not None: await conn.close() @@ -555,28 +537,6 @@ async def get_key_count(tg_id: int) -> int: await conn.close() -async def get_all_users(conn): - """ - Получает список всех пользователей из базы данных. - - Args: - conn: Подключение к базе данных - - Returns: - list: Список Telegram ID всех пользователей - - Raises: - Exception: В случае ошибки при получении данных - """ - try: - users = await conn.fetch("SELECT tg_id FROM connections") - logger.info(f"Получен список всех пользователей. Количество: {len(users)}") - return users - except Exception as e: - logger.error(f"Ошибка при получении списка пользователей: {e}") - raise - - async def add_referral(referred_tg_id: int, referrer_tg_id: int, session: Any): try: if referred_tg_id == referrer_tg_id: @@ -784,7 +744,7 @@ async def get_referral_stats(referrer_tg_id: int): logger.info("Закрытие подключения к базе данных") -async def update_key_expiry(client_id: str, new_expiry_time: int): +async def update_key_expiry(client_id: str, new_expiry_time: int, session: Any): """ Обновление времени истечения ключа для указанного клиента. @@ -795,12 +755,8 @@ async def update_key_expiry(client_id: str, new_expiry_time: int): Raises: Exception: В случае ошибки при подключении к базе данных или обновлении ключа """ - conn = None try: - conn = await asyncpg.connect(DATABASE_URL) - logger.info(f"Установлено подключение к базе данных для обновления времени истечения ключа клиента {client_id}") - - await conn.execute( + await session.execute( """ UPDATE keys SET expiry_time = $1, notified = FALSE, notified_24h = FALSE @@ -814,10 +770,6 @@ async def update_key_expiry(client_id: str, new_expiry_time: int): except Exception as e: logger.error(f"Ошибка при обновлении времени истечения ключа для клиента {client_id}: {e}") raise - finally: - if conn: - await conn.close() - logger.info("Закрытие подключения к базе данных") async def add_balance_to_client(client_id: str, amount: float): diff --git a/handlers/admin/admin_servers.py b/handlers/admin/admin_servers.py index 151475a8..c110fda7 100644 --- a/handlers/admin/admin_servers.py +++ b/handlers/admin/admin_servers.py @@ -11,7 +11,7 @@ from py3xui import AsyncApi from backup import create_backup_and_send_to_admins from config import ADMIN_PASSWORD, ADMIN_USERNAME, DATABASE_URL -from database import add_server_to_db, check_unique_server_name, get_servers_from_db +from database import add_server_to_db, check_unique_server_name, get_keys_by_server, get_servers_from_db from filters.admin import IsAdminFilter from handlers.keys.key_utils import create_key_on_cluster from logger import logger @@ -294,19 +294,12 @@ async def handle_manage_cluster(callback_query: types.CallbackQuery, state: FSMC @router.callback_query(F.data.startswith("sync_cluster|"), IsAdminFilter()) -async def sync_cluster_handler(callback_query: types.CallbackQuery): +async def sync_cluster_handler(callback_query: types.CallbackQuery, session: Any): """Обработчик для синхронизации ключей на всех серверах выбранного кластера.""" cluster_name = callback_query.data.split("|")[1] - conn = await asyncpg.connect(DATABASE_URL) try: - # TODO - query_keys = """ - SELECT tg_id, client_id, email, expiry_time - FROM keys - WHERE server_id = $1 - """ - keys_to_sync = await conn.fetch(query_keys, cluster_name) + keys_to_sync = await get_keys_by_server(None, cluster_name, session) if not keys_to_sync: await callback_query.message.answer( @@ -348,8 +341,6 @@ async def sync_cluster_handler(callback_query: types.CallbackQuery): .row(InlineKeyboardButton(text="🔙 Назад", callback_data="servers_editor")) .as_markup(), ) - finally: - await conn.close() @router.callback_query(F.data.startswith("server_availability|"), IsAdminFilter()) diff --git a/handlers/admin/admin_user_editor.py b/handlers/admin/admin_user_editor.py index 9a577846..31dac3bc 100644 --- a/handlers/admin/admin_user_editor.py +++ b/handlers/admin/admin_user_editor.py @@ -431,7 +431,7 @@ async def handle_expiry_time_input(message: types.Message, state: FSMContext, se await update_key_on_all_servers() - await update_key_expiry(client_id, expiry_time) + await update_key_expiry(client_id, expiry_time, session) response_message = ( f"✅ Время истечения ключа для клиента {client_id} ({email}) успешно обновлено на всех серверах." diff --git a/handlers/coupons.py b/handlers/coupons.py index d010a028..6c340090 100644 --- a/handlers/coupons.py +++ b/handlers/coupons.py @@ -89,5 +89,5 @@ async def activate_coupon(user_id: int, coupon_code: str, session: Any): datetime.utcnow(), ) - await update_balance(user_id, coupon_amount) + await update_balance(user_id, coupon_amount, session) return f"✅ Купон успешно активирован! 🎉\n\nНа ваш баланс добавлено {coupon_amount} рублей 💰." diff --git a/handlers/keys/key_management.py b/handlers/keys/key_management.py index 458ac4c9..946689de 100644 --- a/handlers/keys/key_management.py +++ b/handlers/keys/key_management.py @@ -158,7 +158,7 @@ async def select_tariff_plan(callback_query: CallbackQuery, session: Any): expiry_time = datetime.utcnow() + timedelta(days=duration_days) await create_key(tg_id, expiry_time, None, session, callback_query) - await update_balance(tg_id, -plan_price) + await update_balance(tg_id, -plan_price, session) async def create_key( diff --git a/handlers/keys/keys.py b/handlers/keys/keys.py index ecbb1171..64a094d2 100644 --- a/handlers/keys/keys.py +++ b/handlers/keys/keys.py @@ -30,6 +30,7 @@ from database import ( delete_key, get_balance, get_key_details, + get_keys_by_server, get_servers_from_db, create_temporary_data, store_key, @@ -424,10 +425,7 @@ async def process_callback_renew_plan(callback_query: types.CallbackQuery, sessi total_gb = TOTAL_GB * gb_multiplier.get(plan, 1) if TOTAL_GB > 0 else 0 try: - record = await session.fetchrow( - "SELECT email, expiry_time FROM keys WHERE client_id = $1", - client_id, - ) + record = await get_keys_by_server(tg_id, client_id, session) if record: email = record["email"] @@ -545,8 +543,8 @@ async def complete_key_renewal(tg_id, client_id, email, new_expiry_time, total_g total_gb, ) - await update_key_expiry(client_id, new_expiry_time) - await update_balance(tg_id, -cost) + await update_key_expiry(client_id, new_expiry_time, conn) + await update_balance(tg_id, -cost, conn) logger.info(f"[RENEW] Ключ {client_id} успешно продлён на {plan} мес. для пользователя {tg_id}.") await renew_key_on_cluster() diff --git a/handlers/notifications.py b/handlers/notifications.py index 67cd9890..e038e218 100644 --- a/handlers/notifications.py +++ b/handlers/notifications.py @@ -152,9 +152,9 @@ async def process_10h_record(record, bot, conn): if AUTO_RENEW_KEYS and balance >= RENEWAL_PLANS["1"]["price"]: try: - await update_balance(tg_id, -RENEWAL_PLANS["1"]["price"]) + await update_balance(tg_id, -RENEWAL_PLANS["1"]["price"], conn) new_expiry_time = int((datetime.utcnow() + timedelta(days=30)).timestamp() * 1000) - await update_key_expiry(record["client_id"], new_expiry_time) + await update_key_expiry(record["client_id"], new_expiry_time, conn) servers = await get_servers_from_db() for cluster_id in servers: @@ -241,9 +241,9 @@ async def process_24h_record(record, bot, conn): if AUTO_RENEW_KEYS and balance >= RENEWAL_PLANS["1"]["price"]: try: - await update_balance(tg_id, -RENEWAL_PLANS["1"]["price"]) + await update_balance(tg_id, -RENEWAL_PLANS["1"]["price"], conn) new_expiry_time = int((datetime.utcnow() + timedelta(days=30)).timestamp() * 1000) - await update_key_expiry(record["client_id"], new_expiry_time) + await update_key_expiry(record["client_id"], new_expiry_time, conn) servers = await get_servers_from_db() for cluster_id in servers: @@ -433,10 +433,10 @@ async def process_key(record, bot, conn): try: if AUTO_RENEW_KEYS and balance >= RENEWAL_PLANS["1"]["price"]: - await update_balance(tg_id, -RENEWAL_PLANS["1"]["price"]) + await update_balance(tg_id, -RENEWAL_PLANS["1"]["price"], conn) new_expiry_time = int((datetime.now(moscow_tz) + timedelta(days=30)).timestamp() * 1000) - await update_key_expiry(client_id, new_expiry_time) + await update_key_expiry(client_id, new_expiry_time, conn) servers = await get_servers_from_db()