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.

This commit is contained in:
Zakhar Izmaylov
2025-01-23 11:36:15 +03:00
parent 7dafb47c38
commit 73a365eeb6
11 changed files with 81 additions and 76 deletions
+34 -34
View File
@@ -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
+13 -13
View File
@@ -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:
+4 -4
View File
@@ -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)
+7 -2
View File
@@ -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 "<b>❌ Вы уже активировали этот купон.</b> 🚫 Купоны могут быть активированы только один раз. 🔒"
+5 -5
View File
@@ -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:
+2 -2
View File
@@ -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:
+3 -3
View File
@@ -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]
+2 -2
View File
@@ -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} не найден в базе данных.")
+6 -6
View File
@@ -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):
+2 -2
View File
@@ -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()}
+3 -3
View File
@@ -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}")