From 73a365eeb6240b6442784cc463411e1f53293d83 Mon Sep 17 00:00:00 2001 From: Zakhar Izmaylov Date: Thu, 23 Jan 2025 11:36:15 +0300 Subject: [PATCH] Refactor server management and database interactions. Renamed `get_servers_from_db` to `get_servers` for clarity and updated related handlers to utilize session management. Refactored `add_server_to_db` to `create_server` for consistency. Improved code readability and maintainability across coupon and notification handlers by streamlining database calls and enhancing error handling. --- database.py | 68 ++++++++++++++--------------- handlers/admin/admin_servers.py | 26 +++++------ handlers/admin/admin_user_editor.py | 8 ++-- handlers/coupons.py | 9 +++- handlers/keys/key_utils.py | 10 ++--- handlers/keys/keys.py | 4 +- handlers/keys/subscriptions.py | 6 +-- handlers/keys/trial_key.py | 4 +- handlers/notifications.py | 12 ++--- handlers/utils.py | 4 +- servers.py | 6 +-- 11 files changed, 81 insertions(+), 76 deletions(-) diff --git a/database.py b/database.py index a833e664..fee9d8da 100644 --- a/database.py +++ b/database.py @@ -86,15 +86,10 @@ async def check_unique_server_name(server_name: str, session: Any, cluster_name: """ if cluster_name: result = await session.fetchrow( - "SELECT 1 FROM servers WHERE server_name = $1 AND cluster_name = $2 LIMIT 1", - server_name, - cluster_name + "SELECT 1 FROM servers WHERE server_name = $1 AND cluster_name = $2 LIMIT 1", server_name, cluster_name ) else: - result = await session.fetchrow( - "SELECT 1 FROM servers WHERE server_name = $1 LIMIT 1", - server_name - ) + result = await session.fetchrow("SELECT 1 FROM servers WHERE server_name = $1 LIMIT 1", server_name) return result is None @@ -130,6 +125,7 @@ async def create_coupon(coupon_code: str, amount: float, usage_limit: int, sessi logger.error(f"Ошибка при создании купона {coupon_code}: {e}") raise + async def get_coupon_by_code(coupon_code: str, session: Any) -> dict | None: """ Получает информацию о купоне по его коду. @@ -1101,34 +1097,37 @@ async def check_notification_time(tg_id: int, notification_type: str, hours: int await conn.close() -async def get_servers_from_db(): - conn = await asyncpg.connect(DATABASE_URL) +async def get_servers(session: Any = None): + conn = None + try: + conn = session if session is not None else await asyncpg.connect(DATABASE_URL) - result = await conn.fetch( - """ - SELECT cluster_name, server_name, api_url, subscription_url, inbound_id - FROM servers - """ - ) - - await conn.close() - - servers = {} - for row in result: - cluster_name = row["cluster_name"] - if cluster_name not in servers: - servers[cluster_name] = [] - - servers[cluster_name].append( - { - "server_name": row["server_name"], - "api_url": row["api_url"], - "subscription_url": row["subscription_url"], - "inbound_id": row["inbound_id"], - } + result = await conn.fetch( + """ + SELECT cluster_name, server_name, api_url, subscription_url, inbound_id + FROM servers + """ ) + servers = {} + for row in result: + cluster_name = row["cluster_name"] + if cluster_name not in servers: + servers[cluster_name] = [] - return servers + servers[cluster_name].append( + { + "server_name": row["server_name"], + "api_url": row["api_url"], + "subscription_url": row["subscription_url"], + "inbound_id": row["inbound_id"], + } + ) + + return servers + + finally: + if conn is not None and session is None: + await conn.close() async def delete_user_data(session: Any, tg_id: int): @@ -1267,7 +1266,7 @@ async def delete_key(identifier, session): logger.error(f"Ошибка при удалении ключа с идентификатором {identifier} из базы данных: {e}") -async def add_server_to_db( +async def create_server( cluster_name: str, server_name: str, api_url: str, subscription_url: str, inbound_id: int, session: Any ): """ @@ -1301,6 +1300,7 @@ async def add_server_to_db( logger.error(f"Ошибка при добавлении сервера {server_name} в кластер {cluster_name}: {e}") raise + async def delete_server(server_name: str, session: Any): """ Удаляет сервер из базы данных по его названию. @@ -1352,6 +1352,7 @@ async def create_coupon_usage(coupon_id: int, user_id: int, session: Any): logger.error(f"Ошибка при создании записи об использовании купона {coupon_id} пользователем {user_id}: {e}") raise + async def check_coupon_usage(coupon_id: int, user_id: int, session: Any) -> bool: """ Проверяет, использовал ли пользователь данный купон. @@ -1406,4 +1407,3 @@ async def update_coupon_usage_count(coupon_id: int, session: Any): except Exception as e: logger.error(f"Ошибка при обновлении счетчика использования купона {coupon_id}: {e}") raise - diff --git a/handlers/admin/admin_servers.py b/handlers/admin/admin_servers.py index fd429596..e0d27e0d 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, delete_server, get_keys_by_server, get_servers_from_db +from database import create_server, check_unique_server_name, delete_server, get_keys_by_server, get_servers from filters.admin import IsAdminFilter from handlers.keys.key_utils import create_key_on_cluster from logger import logger @@ -28,8 +28,8 @@ class UserEditorState(StatesGroup): @router.callback_query(F.data == "servers_editor", IsAdminFilter()) -async def handle_servers_editor(callback_query: types.CallbackQuery): - servers = await get_servers_from_db() +async def handle_servers_editor(callback_query: types.CallbackQuery, session: Any): + servers = await get_servers(session) builder = InlineKeyboardBuilder() @@ -225,7 +225,7 @@ async def handle_inbound_id_input(message: types.Message, state: FSMContext, ses api_url = user_data.get("api_url") subscription_url = user_data.get("subscription_url") - await add_server_to_db( + await create_server( cluster_name=cluster_name, server_name=server_name, api_url=api_url, @@ -246,10 +246,10 @@ async def handle_inbound_id_input(message: types.Message, state: FSMContext, ses @router.callback_query(F.data.startswith("manage_cluster|"), IsAdminFilter()) -async def handle_manage_cluster(callback_query: types.CallbackQuery, state: FSMContext): +async def handle_manage_cluster(callback_query: types.CallbackQuery, state: FSMContext, session: Any): cluster_name = callback_query.data.split("|")[1] - servers = await get_servers_from_db() + servers = await get_servers(session) cluster_servers = servers.get(cluster_name, []) builder = InlineKeyboardBuilder() @@ -310,7 +310,7 @@ async def sync_cluster_handler(callback_query: types.CallbackQuery, session: Any ) return - servers = await get_servers_from_db() + servers = await get_servers(session) cluster_servers = servers.get(cluster_name, []) for key in keys_to_sync: @@ -344,10 +344,10 @@ async def sync_cluster_handler(callback_query: types.CallbackQuery, session: Any @router.callback_query(F.data.startswith("server_availability|"), IsAdminFilter()) -async def handle_check_server_availability(callback_query: types.CallbackQuery): +async def handle_check_server_availability(callback_query: types.CallbackQuery, session: Any): cluster_name = callback_query.data.split("|")[1] - servers = await get_servers_from_db() + servers = await get_servers(session) cluster_servers = servers.get(cluster_name, []) if not cluster_servers: @@ -382,10 +382,10 @@ async def handle_check_server_availability(callback_query: types.CallbackQuery): @router.callback_query(F.data.startswith("manage_server|"), IsAdminFilter()) -async def handle_manage_server(callback_query: types.CallbackQuery, state: FSMContext): +async def handle_manage_server(callback_query: types.CallbackQuery, state: FSMContext, session: Any): server_name = callback_query.data.split("|")[1] - servers = await get_servers_from_db() + servers = await get_servers(session) server = None cluster_name = None @@ -463,10 +463,10 @@ async def handle_add_server(callback_query: types.CallbackQuery, state: FSMConte @router.callback_query(F.data.startswith("backup_cluster|"), IsAdminFilter()) -async def handle_backup_cluster(callback_query: types.CallbackQuery): +async def handle_backup_cluster(callback_query: types.CallbackQuery, session: Any): cluster_name = callback_query.data.split("|")[1] - servers = await get_servers_from_db() + servers = await get_servers(session) cluster_servers = servers.get(cluster_name, []) for server in cluster_servers: diff --git a/handlers/admin/admin_user_editor.py b/handlers/admin/admin_user_editor.py index 31dac3bc..87c5193d 100644 --- a/handlers/admin/admin_user_editor.py +++ b/handlers/admin/admin_user_editor.py @@ -15,7 +15,7 @@ from database import ( get_client_id_by_email, get_key_details, get_keys, - get_servers_from_db, + get_servers, set_trial, update_key_expiry, ) @@ -410,7 +410,7 @@ async def handle_expiry_time_input(message: types.Message, state: FSMContext, se await state.clear() return - clusters = await get_servers_from_db() + clusters = await get_servers(session) async def update_key_on_all_servers(): tasks = [] @@ -488,7 +488,7 @@ async def process_callback_confirm_delete(callback_query: types.CallbackQuery, s builder = InlineKeyboardBuilder() builder.row(InlineKeyboardButton(text="⬅️ Назад", callback_data="view_keys")) - clusters = await get_servers_from_db() + clusters = await get_servers(session) async def delete_key_from_servers(email, client_id): tasks = [] @@ -569,7 +569,7 @@ async def delete_user(callback_query: types.CallbackQuery, session: Any): try: tasks = [] for email, client_id in key_records: - servers = await get_servers_from_db() + servers = await get_servers(session) for cluster_id, cluster in servers.items(): tasks.append(delete_key_from_cluster(cluster_id, email, client_id)) await asyncio.gather(*tasks) diff --git a/handlers/coupons.py b/handlers/coupons.py index 6c74359e..b652cf1c 100644 --- a/handlers/coupons.py +++ b/handlers/coupons.py @@ -7,7 +7,13 @@ from aiogram.fsm.state import State, StatesGroup from aiogram.types import InlineKeyboardButton from aiogram.utils.keyboard import InlineKeyboardBuilder -from database import check_coupon_usage, create_coupon_usage, get_coupon_by_code, update_balance, update_coupon_usage_count +from database import ( + check_coupon_usage, + create_coupon_usage, + get_coupon_by_code, + update_balance, + update_coupon_usage_count, +) class CouponActivationState(StatesGroup): @@ -50,7 +56,6 @@ async def activate_coupon(user_id: int, coupon_code: str, session: Any): usage_exists = await check_coupon_usage(coupon_record["id"], user_id, session) - if usage_exists: return "❌ Вы уже активировали этот купон. 🚫 Купоны могут быть активированы только один раз. 🔒" diff --git a/handlers/keys/key_utils.py b/handlers/keys/key_utils.py index 98877801..00e7098c 100644 --- a/handlers/keys/key_utils.py +++ b/handlers/keys/key_utils.py @@ -4,7 +4,7 @@ from py3xui import AsyncApi from client import ClientConfig, add_client, delete_client, extend_client_key from config import ADMIN_PASSWORD, ADMIN_USERNAME, LIMIT_IP, SUPERNODE, TOTAL_GB -from database import get_servers_from_db +from database import get_servers from logger import logger @@ -13,7 +13,7 @@ async def create_key_on_cluster(cluster_id: str, tg_id: int, client_id: str, ema Создает ключ на всех серверах указанного кластера. """ try: - servers = await get_servers_from_db() + servers = await get_servers() cluster = servers.get(cluster_id) if not cluster: @@ -105,7 +105,7 @@ async def create_client_on_server( async def renew_key_in_cluster(cluster_id, email, client_id, new_expiry_time, total_gb): try: - servers = await get_servers_from_db() + servers = await get_servers() cluster = servers.get(cluster_id) if not cluster: @@ -147,7 +147,7 @@ async def renew_key_in_cluster(cluster_id, email, client_id, new_expiry_time, to async def delete_key_from_cluster(cluster_id, email, client_id): """Удаление ключа с серверов в кластере""" try: - servers = await get_servers_from_db() + servers = await get_servers() cluster = servers.get(cluster_id) if not cluster: @@ -186,7 +186,7 @@ async def delete_key_from_cluster(cluster_id, email, client_id): async def update_key_on_cluster(tg_id, client_id, email, expiry_time, cluster_id): try: - servers = await get_servers_from_db() + servers = await get_servers() cluster = servers.get(cluster_id) if not cluster: diff --git a/handlers/keys/keys.py b/handlers/keys/keys.py index 64a094d2..8213f7eb 100644 --- a/handlers/keys/keys.py +++ b/handlers/keys/keys.py @@ -31,7 +31,7 @@ from database import ( get_balance, get_key_details, get_keys_by_server, - get_servers_from_db, + get_servers, create_temporary_data, store_key, update_balance, @@ -385,7 +385,7 @@ async def process_callback_confirm_delete(callback_query: types.CallbackQuery, s reply_markup=keyboard, ) - servers = await get_servers_from_db() + servers = await get_servers(session) async def delete_key_from_servers(): try: diff --git a/handlers/keys/subscriptions.py b/handlers/keys/subscriptions.py index d3f6d36e..6203252b 100644 --- a/handlers/keys/subscriptions.py +++ b/handlers/keys/subscriptions.py @@ -7,7 +7,7 @@ import asyncpg from aiohttp import web from config import DATABASE_URL, PROJECT_NAME, SUB_MESSAGE, SUPERNODE, TRANSITION_DATE_STR -from database import get_key_details, get_servers_from_db +from database import get_key_details, get_servers from logger import logger # Глобальная переменная для пула соединений @@ -122,7 +122,7 @@ async def handle_old_subscription(request): status=400, ) - servers = await get_servers_from_db() + servers = await get_servers() cluster_servers = servers.get(cluster_name, []) logger.info(f"Сервера в кластере: {cluster_servers}") @@ -180,7 +180,7 @@ async def handle_new_subscription(request): status=403, ) - servers = await get_servers_from_db() + servers = await get_servers() cluster_servers = servers.get(cluster_name, []) urls = [f"{server['subscription_url']}/{email}" for server in cluster_servers] diff --git a/handlers/keys/trial_key.py b/handlers/keys/trial_key.py index 12725c02..1c433e5c 100644 --- a/handlers/keys/trial_key.py +++ b/handlers/keys/trial_key.py @@ -8,7 +8,7 @@ from py3xui import AsyncApi from client import ClientConfig, add_client from config import ADMIN_PASSWORD, ADMIN_USERNAME, LIMIT_IP, PUBLIC_LINK, SUPERNODE, TOTAL_GB, TRIAL_TIME -from database import get_servers_from_db, get_trial, store_key, set_trial +from database import get_servers, get_trial, store_key, set_trial from handlers.texts import INSTRUCTIONS from handlers.utils import generate_random_email, get_least_loaded_cluster from logger import logger @@ -33,7 +33,7 @@ async def create_trial_key(tg_id: int, session: Any): expiry_time = current_time + timedelta(days=TRIAL_TIME) expiry_timestamp = int(expiry_time.timestamp() * 1000) - clusters = await get_servers_from_db() + clusters = await get_servers(session) least_loaded_cluster = await get_least_loaded_cluster() if least_loaded_cluster not in clusters: raise ValueError(f"Кластер {least_loaded_cluster} не найден в базе данных.") diff --git a/handlers/notifications.py b/handlers/notifications.py index e038e218..35c7c2c6 100644 --- a/handlers/notifications.py +++ b/handlers/notifications.py @@ -28,7 +28,7 @@ from database import ( check_notification_time, delete_key, get_balance, - get_servers_from_db, + get_servers, update_balance, update_key_expiry, ) @@ -156,7 +156,7 @@ async def process_10h_record(record, bot, conn): new_expiry_time = int((datetime.utcnow() + timedelta(days=30)).timestamp() * 1000) await update_key_expiry(record["client_id"], new_expiry_time, conn) - servers = await get_servers_from_db() + servers = await get_servers(conn) for cluster_id in servers: await renew_key_in_cluster(cluster_id, email, record["client_id"], new_expiry_time, TOTAL_GB) logger.info(f"Ключ для пользователя {tg_id} успешно продлен в кластере {cluster_id}.") @@ -245,7 +245,7 @@ async def process_24h_record(record, bot, conn): new_expiry_time = int((datetime.utcnow() + timedelta(days=30)).timestamp() * 1000) await update_key_expiry(record["client_id"], new_expiry_time, conn) - servers = await get_servers_from_db() + servers = await get_servers(conn) for cluster_id in servers: await renew_key_in_cluster(cluster_id, email, record["client_id"], new_expiry_time, TOTAL_GB) logger.info(f"Ключ для пользователя {tg_id} успешно продлен в кластере {cluster_id}.") @@ -438,7 +438,7 @@ async def process_key(record, bot, conn): new_expiry_time = int((datetime.now(moscow_tz) + timedelta(days=30)).timestamp() * 1000) await update_key_expiry(client_id, new_expiry_time, conn) - servers = await get_servers_from_db() + servers = await get_servers(conn) for cluster_id in servers: await renew_key_in_cluster(cluster_id, email, client_id, new_expiry_time, TOTAL_GB) @@ -487,7 +487,7 @@ async def process_key(record, bot, conn): logger.error(f"Не удалось отправить уведомление об истечении клиенту {tg_id}: {e}") if AUTO_DELETE_EXPIRED_KEYS: - servers = await get_servers_from_db() + servers = await get_servers(conn) for cluster_id in servers: try: @@ -509,7 +509,7 @@ async def process_key(record, bot, conn): async def check_online_users(): - servers = await get_servers_from_db() + servers = await get_servers() for cluster_id, cluster in servers.items(): for server_id, server in enumerate(cluster): diff --git a/handlers/utils.py b/handlers/utils.py index 608a0f4c..dba6fb7b 100644 --- a/handlers/utils.py +++ b/handlers/utils.py @@ -7,7 +7,7 @@ import asyncpg from bot import bot from config import DATABASE_URL -from database import get_servers_from_db +from database import get_servers from logger import logger @@ -59,7 +59,7 @@ async def get_least_loaded_cluster() -> str: Returns: str: Идентификатор наименее загруженного кластера. """ - servers = await get_servers_from_db() + servers = await get_servers() cluster_loads: dict[str, int] = {cluster_id: 0 for cluster_id in servers.keys()} diff --git a/servers.py b/servers.py index f97b7bbb..521b55f1 100644 --- a/servers.py +++ b/servers.py @@ -9,7 +9,7 @@ from ping3 import ping from bot import bot from config import ADMIN_ID, DATABASE_URL, PING_TIME -from database import add_server_to_db, check_unique_server_name, get_servers_from_db +from database import create_server, check_unique_server_name, get_servers from logger import logger try: @@ -37,7 +37,7 @@ async def sync_servers_with_db(): exists = await check_unique_server_name(server_info["name"], conn, cluster_name) if not exists: - await add_server_to_db( + await create_server( cluster_name=cluster_name, server_name=server_info["name"], api_url=server_info["API_URL"], @@ -117,7 +117,7 @@ async def check_servers(): Периодическая проверка серверов с учетом извлечения хоста из `api_url`. """ while True: - servers = await get_servers_from_db() + servers = await get_servers() current_time = datetime.now() logger.info(f"Начинаю проверку серверов: {current_time}")