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.

This commit is contained in:
Zakhar Izmaylov
2025-01-23 11:06:25 +03:00
parent e42a66dde7
commit 906601f5cd
7 changed files with 57 additions and 116 deletions
+41 -89
View File
@@ -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):
+3 -12
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, 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())
+1 -1
View File
@@ -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}) успешно обновлено на всех серверах."
+1 -1
View File
@@ -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"<b>✅ Купон успешно активирован! 🎉</b>\n\nНа ваш баланс добавлено <b>{coupon_amount} рублей</b> 💰."
+1 -1
View File
@@ -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(
+4 -6
View File
@@ -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()
+6 -6
View File
@@ -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()