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:
+41
-89
@@ -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):
|
||||
|
||||
@@ -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())
|
||||
|
||||
@@ -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
@@ -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> 💰."
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
Reference in New Issue
Block a user