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:
+34
-34
@@ -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
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
@@ -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> 🚫 Купоны могут быть активированы только один раз. 🔒"
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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} не найден в базе данных.")
|
||||
|
||||
@@ -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
@@ -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
@@ -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}")
|
||||
|
||||
Reference in New Issue
Block a user