fix: key management in admin menu / fixed synchronization

This commit is contained in:
Capybara-z
2025-05-28 00:44:26 +03:00
parent daa9710579
commit 32fda4e8b5
16 changed files with 234 additions and 191 deletions
+11 -1
View File
@@ -8,6 +8,7 @@ from aiogram.filters import ExceptionTypeFilter
from aiogram.fsm.storage.memory import MemoryStorage from aiogram.fsm.storage.memory import MemoryStorage
from aiogram.types import BufferedInputFile, ErrorEvent from aiogram.types import BufferedInputFile, ErrorEvent
from aiogram.utils.markdown import hbold from aiogram.utils.markdown import hbold
import subprocess
from config import ADMIN_ID, API_TOKEN from config import ADMIN_ID, API_TOKEN
from filters.private import IsPrivateFilter from filters.private import IsPrivateFilter
@@ -17,7 +18,16 @@ bot = Bot(token=API_TOKEN, default=DefaultBotProperties(parse_mode=ParseMode.HTM
storage = MemoryStorage() storage = MemoryStorage()
dp = Dispatcher(bot=bot, storage=storage) dp = Dispatcher(bot=bot, storage=storage)
version = "4.3-b240504 (ORM update)"
def get_git_commit_number() -> str:
try:
count = subprocess.check_output(["git", "rev-list", "--count", "HEAD"])
return count.decode("utf-8").strip()
except Exception:
return "unknown"
version = f"4.3-b240504 (commit #{get_git_commit_number()})"
dp.message.filter(IsPrivateFilter()) dp.message.filter(IsPrivateFilter())
dp.callback_query.filter(IsPrivateFilter()) dp.callback_query.filter(IsPrivateFilter())
+10
View File
@@ -185,3 +185,13 @@ async def mark_key_as_unfrozen(
), ),
{"expiry": new_expiry_time, "tg_id": tg_id, "client_id": client_id}, {"expiry": new_expiry_time, "tg_id": tg_id, "client_id": client_id},
) )
async def update_key_tariff(session: AsyncSession, client_id: str, tariff_id: int):
await session.execute(
update(Key)
.where(Key.client_id == client_id)
.values(tariff_id=tariff_id)
)
await session.commit()
logger.info(f"Тариф ключа {client_id} обновлён на {tariff_id}")
+9
View File
@@ -45,6 +45,15 @@ async def get_tariffs_for_cluster(session: AsyncSession, cluster_name: str):
.limit(1) .limit(1)
) )
row = server_row.first() row = server_row.first()
if not row:
server_row = await session.execute(
select(Server.tariff_group)
.where(Server.server_name == cluster_name)
.limit(1)
)
row = server_row.first()
if not row or not row[0]: if not row or not row[0]:
return [] return []
@@ -495,6 +495,7 @@ async def handle_sync_server(
Key.client_id, Key.client_id,
Key.email, Key.email,
Key.expiry_time, Key.expiry_time,
Key.tariff_id,
) )
.join(Key, Server.cluster_name == Key.server_id) .join(Key, Server.cluster_name == Key.server_id)
.where(Server.server_name == server_name) .where(Server.server_name == server_name)
@@ -530,6 +531,8 @@ async def handle_sync_server(
key["email"], key["email"],
key["expiry_time"], key["expiry_time"],
semaphore, semaphore,
plan=key["tariff_id"],
session=session,
) )
await asyncio.sleep(0.6) await asyncio.sleep(0.6)
except Exception as e: except Exception as e:
+3 -3
View File
@@ -156,7 +156,7 @@ async def process_tariff_traffic(message: Message, state: FSMContext):
) )
return return
await state.update_data(traffic_limit=traffic * 1024**3 if traffic > 0 else None) await state.update_data(traffic_limit=traffic if traffic > 0 else None)
await state.set_state(TariffCreateState.device_limit) await state.set_state(TariffCreateState.device_limit)
await message.answer( await message.answer(
"📱 Введите <b>лимит устройств (HWID)</b> для тарифа (например: <i>3</i>, 0 — безлимит):", "📱 Введите <b>лимит устройств (HWID)</b> для тарифа (например: <i>3</i>, 0 — безлимит):",
@@ -315,7 +315,7 @@ async def view_tariff(
return return
traffic_text = ( traffic_text = (
f"{tariff.traffic_limit // 1024**3} ГБ" if tariff.traffic_limit else "Безлимит" f"{tariff.traffic_limit} ГБ" if tariff.traffic_limit else "Безлимит"
) )
device_text = ( device_text = (
f"{tariff.device_limit}" if tariff.device_limit is not None else "Безлимит" f"{tariff.device_limit}" if tariff.device_limit is not None else "Безлимит"
@@ -453,7 +453,7 @@ async def apply_edit(message: Message, state: FSMContext, session: AsyncSession)
if num < 0: if num < 0:
raise ValueError raise ValueError
if field == "traffic_limit": if field == "traffic_limit":
value = num * 1024**3 if num > 0 else None value = num if num > 0 else None
elif field == "device_limit": elif field == "device_limit":
value = num if num > 0 else None value = num if num > 0 else None
else: else:
+46 -48
View File
@@ -917,7 +917,12 @@ async def confirm_admin_key_reissue(
) )
return return
await update_subscription(tg_id, email, session, cluster_override=cluster_id) result = await session.execute(
select(Key.remnawave_link).where(Key.email == email)
)
remnawave_link = result.scalar_one_or_none()
await update_subscription(tg_id, email, session, cluster_override=cluster_id, remnawave_link=remnawave_link)
await handle_key_edit( await handle_key_edit(
callback_query, callback_query,
@@ -1149,58 +1154,51 @@ async def process_user_search(
async def change_expiry_time( async def change_expiry_time(
expiry_time: int, email: str, session: AsyncSession expiry_time: int, email: str, session: AsyncSession
) -> Exception | None: ) -> Exception | None:
result = await session.execute(select(Key.client_id).where(Key.email == email)) result = await session.execute(select(Key.client_id, Key.tariff_id, Key.server_id).where(Key.email == email))
client_id = result.scalar_one_or_none() row = result.first()
if client_id is None: if not row:
return ValueError(f"User with email {email} was not found") return ValueError(f"User with email {email} was not found")
result = await session.execute( client_id, tariff_id, server_id = row
select(Key.server_id).where(Key.client_id == client_id)
)
server_id = result.scalar_one_or_none()
if server_id is None: if server_id is None:
return ValueError(f"Key with client_id {client_id} was not found") return ValueError(f"Key with client_id {client_id} was not found")
result = await session.execute( traffic_limit = 0
select(Server.tariff_group) device_limit = None
.where(or_(Server.server_name == server_id, Server.cluster_name == server_id)) if tariff_id:
.limit(1) result = await session.execute(
select(Tariff.traffic_limit, Tariff.device_limit)
.where(Tariff.id == tariff_id, Tariff.is_active.is_(True))
)
tariff = result.first()
if tariff:
traffic_limit = int(tariff[0]) if tariff[0] is not None else 0
device_limit = int(tariff[1]) if tariff[1] is not None else None
servers = await get_servers(session=session)
if server_id in servers:
target_cluster = server_id
else:
target_cluster = None
for cluster_name, cluster_servers in servers.items():
if any(s.get("server_name") == server_id for s in cluster_servers):
target_cluster = cluster_name
break
if not target_cluster:
return ValueError(f"No suitable cluster found for server {server_id}")
await renew_key_in_cluster(
cluster_id=target_cluster,
email=email,
client_id=client_id,
new_expiry_time=expiry_time,
total_gb=traffic_limit,
session=session,
hwid_device_limit=device_limit
) )
tariff_group = result.scalar_one_or_none()
if not tariff_group:
return ValueError(f"Tariff group not found for server_id={server_id}")
result = await session.execute(
select(Tariff.duration_days, Tariff.traffic_limit)
.where(Tariff.group_code == tariff_group, Tariff.is_active.is_(True))
.order_by(Tariff.duration_days)
)
tariffs = result.all()
if not tariffs:
return ValueError(f"No tariffs found for group {tariff_group}")
added_days = max((expiry_time - int(time.time() * 1000)) / (1000 * 86400), 1)
closest_tariff = min(tariffs, key=lambda t: abs(t[0] - added_days))
traffic_limit = closest_tariff[1] or 0
clusters = await get_servers(session=session)
async def update_key_on_all_servers():
tasks = [
renew_key_in_cluster(
cluster_id=cluster_name,
email=email,
client_id=client_id,
new_expiry_time=expiry_time,
total_gb=traffic_limit,
session=session,
)
for cluster_name in clusters
]
await asyncio.gather(*tasks, return_exceptions=True)
await update_key_on_all_servers()
await update_key_expiry(session, client_id, expiry_time) await update_key_expiry(session, client_id, expiry_time)
return None return None
+9 -5
View File
@@ -77,15 +77,18 @@ async def key_cluster_mode(
is_trial = data.get("is_trial", False) is_trial = data.get("is_trial", False)
device_limit = 0 device_limit = 0
traffic_limit_bytes = None traffic_limit_gb = None
if is_trial: if is_trial:
device_limit = TRIAL_CONFIG.get("hwid_limit", 1) device_limit = TRIAL_CONFIG.get("hwid_limit", 1)
traffic_limit_bytes = int(TRIAL_CONFIG.get("traffic_limit_gb", 100) * 1024**3) traffic_limit_gb = TRIAL_CONFIG.get("traffic_limit_gb", 100)
elif plan: elif plan:
tariff = await get_tariff_by_id(session, plan) tariff = await get_tariff_by_id(session, plan)
if tariff and tariff.get("device_limit") is not None: if tariff:
device_limit = int(tariff["device_limit"]) if tariff.get("device_limit") is not None:
device_limit = int(tariff["device_limit"])
if tariff.get("traffic_limit") is not None:
traffic_limit_gb = int(tariff["traffic_limit"])
least_loaded_cluster = await get_least_loaded_cluster(session) least_loaded_cluster = await get_least_loaded_cluster(session)
await create_key_on_cluster( await create_key_on_cluster(
@@ -97,7 +100,8 @@ async def key_cluster_mode(
plan=plan, plan=plan,
session=session, session=session,
hwid_limit=device_limit, hwid_limit=device_limit,
traffic_limit_bytes=traffic_limit_bytes, traffic_limit_bytes=traffic_limit_gb,
is_trial=is_trial,
) )
logger.info( logger.info(
+50 -69
View File
@@ -41,6 +41,7 @@ from handlers.buttons import (
PC_BUTTON, PC_BUTTON,
SUPPORT, SUPPORT,
TV_BUTTON, TV_BUTTON,
MY_SUB
) )
from handlers.keys.key_utils import create_client_on_server from handlers.keys.key_utils import create_client_on_server
from handlers.texts import SELECT_COUNTRY_MSG, key_message_success from handlers.texts import SELECT_COUNTRY_MSG, key_message_success
@@ -349,7 +350,6 @@ async def finalize_key_creation(
client_id = old_key_details["client_id"] client_id = old_key_details["client_id"]
email = old_key_details["email"] email = old_key_details["email"]
expiry_timestamp = old_key_details["expiry_time"] expiry_timestamp = old_key_details["expiry_time"]
else: else:
while True: while True:
key_name = generate_random_email() key_name = generate_random_email()
@@ -360,6 +360,25 @@ async def finalize_key_creation(
email = key_name.lower() email = key_name.lower()
expiry_timestamp = int(expiry_time.timestamp() * 1000) expiry_timestamp = int(expiry_time.timestamp() * 1000)
traffic_limit_bytes = None
device_limit = 0
data = await state.get_data() if state else {}
is_trial = data.get("is_trial", False)
if is_trial:
from config import TRIAL_CONFIG
traffic_limit_bytes = int(TRIAL_CONFIG.get("traffic_limit_gb", 100)) * 1024**3
device_limit = TRIAL_CONFIG.get("hwid_limit", 1)
elif data.get("tariff_id") or tariff_id:
tariff_id = data.get("tariff_id") or tariff_id
result = await session.execute(select(Tariff).where(Tariff.id == tariff_id))
tariff = result.scalar_one_or_none()
if tariff:
if tariff.traffic_limit is not None:
traffic_limit_bytes = int(tariff.traffic_limit)
if tariff.device_limit is not None:
device_limit = int(tariff.device_limit)
public_link = None public_link = None
remnawave_link = None remnawave_link = None
created_at = int(datetime.now(moscow_tz).timestamp() * 1000) created_at = int(datetime.now(moscow_tz).timestamp() * 1000)
@@ -373,31 +392,22 @@ async def finalize_key_creation(
raise ValueError(f"Сервер {selected_country} не найден") raise ValueError(f"Сервер {selected_country} не найден")
panel_type = server_info.panel_type.lower() panel_type = server_info.panel_type.lower()
cluster_info = await check_server_name_by_cluster( cluster_info = await check_server_name_by_cluster(session, server_info.server_name)
session, server_info.server_name
)
if not cluster_info: if not cluster_info:
raise ValueError(f"Кластер для сервера {server_info.server_name} не найден") raise ValueError(f"Кластер для сервера {server_info.server_name} не найден")
is_full_remnawave = await is_full_remnawave_cluster( is_full_remnawave = await is_full_remnawave_cluster(cluster_info["cluster_name"], session)
cluster_info["cluster_name"], session
)
if old_key_name: if old_key_name:
old_server_id = old_key_details["server_id"] old_server_id = old_key_details["server_id"]
if old_server_id: if old_server_id:
result = await session.execute( result = await session.execute(select(Server).where(Server.server_name == old_server_id))
select(Server).where(Server.server_name == old_server_id)
)
old_server_info = result.scalar_one_or_none() old_server_info = result.scalar_one_or_none()
if old_server_info: if old_server_info:
try: try:
if old_server_info.panel_type.lower() == "3x-ui": if old_server_info.panel_type.lower() == "3x-ui":
xui = await get_xui_instance(old_server_info.api_url) xui = await get_xui_instance(old_server_info.api_url)
await delete_client( await delete_client(xui, old_server_info.inbound_id, email, client_id)
xui, old_server_info.inbound_id, email, client_id
)
await session.execute( await session.execute(
update(Key) update(Key)
.where(Key.tg_id == tg_id, Key.email == email) .where(Key.tg_id == tg_id, Key.email == email)
@@ -418,21 +428,21 @@ async def finalize_key_creation(
if panel_type == "remnawave" or is_full_remnawave: if panel_type == "remnawave" or is_full_remnawave:
remna = RemnawaveAPI(server_info.api_url) remna = RemnawaveAPI(server_info.api_url)
if not await remna.login(REMNAWAVE_LOGIN, REMNAWAVE_PASSWORD): if not await remna.login(REMNAWAVE_LOGIN, REMNAWAVE_PASSWORD):
raise ValueError( raise ValueError(f"❌ Не удалось авторизоваться в Remnawave ({server_info.server_name})")
f"❌ Не удалось авторизоваться в Remnawave ({server_info.server_name})"
)
expire_at = ( expire_at = datetime.utcfromtimestamp(expiry_timestamp / 1000).isoformat() + "Z"
datetime.utcfromtimestamp(expiry_timestamp / 1000).isoformat() + "Z"
)
user_data = { user_data = {
"username": email, "username": email,
"trafficLimitStrategy": "NO_RESET", "trafficLimitStrategy": "NO_RESET",
"expireAt": expire_at, "expireAt": expire_at,
"telegramId": tg_id, "telegramId": tg_id,
"activeUserInbounds": [server_info.inbound_id], "activeUserInbounds": [server_info.inbound_id],
"hwidDeviceLimit": 0,
} }
if traffic_limit_bytes:
user_data["trafficLimitBytes"] = traffic_limit_bytes
if device_limit:
user_data["hwidDeviceLimit"] = device_limit
result = await remna.create_user(user_data) result = await remna.create_user(user_data)
if not result: if not result:
raise ValueError("❌ Ошибка при создании пользователя в Remnawave") raise ValueError("❌ Ошибка при создании пользователя в Remnawave")
@@ -442,9 +452,7 @@ async def finalize_key_creation(
if old_key_name: if old_key_name:
await session.execute( await session.execute(
update(Key) update(Key).where(Key.tg_id == tg_id, Key.email == email).values(client_id=client_id)
.where(Key.tg_id == tg_id, Key.email == email)
.values(client_id=client_id)
) )
if panel_type == "3x-ui": if panel_type == "3x-ui":
@@ -461,12 +469,13 @@ async def finalize_key_creation(
email=email, email=email,
expiry_timestamp=expiry_timestamp, expiry_timestamp=expiry_timestamp,
semaphore=semaphore, semaphore=semaphore,
session=session,
plan=tariff_id,
is_trial=is_trial,
) )
public_link = f"{PUBLIC_LINK}{email}/{tg_id}" public_link = f"{PUBLIC_LINK}{email}/{tg_id}"
logger.info( logger.info(f"[Key Creation] Подписка создана для пользователя {tg_id} на сервере {selected_country}")
f"[Key Creation] Подписка создана для пользователя {tg_id} на сервере {selected_country}"
)
if old_key_name: if old_key_name:
update_data = {"server_id": selected_country} update_data = {"server_id": selected_country}
@@ -474,19 +483,8 @@ async def finalize_key_creation(
update_data["key"] = public_link update_data["key"] = public_link
elif panel_type == "remnawave": elif panel_type == "remnawave":
update_data["remnawave_link"] = remnawave_link update_data["remnawave_link"] = remnawave_link
await session.execute(update(Key).where(Key.tg_id == tg_id, Key.email == email).values(**update_data))
await session.execute(
update(Key)
.where(Key.tg_id == tg_id, Key.email == email)
.values(**update_data)
)
else: else:
data = {}
if state:
data = await state.get_data()
tariff_id = data.get("tariff_id") or tariff_id
new_key = Key( new_key = Key(
tg_id=tg_id, tg_id=tg_id,
client_id=client_id, client_id=client_id,
@@ -496,19 +494,17 @@ async def finalize_key_creation(
key=public_link, key=public_link,
remnawave_link=remnawave_link, remnawave_link=remnawave_link,
server_id=selected_country, server_id=selected_country,
tariff_id=data.get("tariff_id"), tariff_id=tariff_id,
) )
session.add(new_key) session.add(new_key)
data = await state.get_data() if is_trial:
if data.get("is_trial"):
trial_status = await get_trial(session, tg_id) trial_status = await get_trial(session, tg_id)
if trial_status in [0, -1]: if trial_status in [0, -1]:
await update_trial(session, tg_id, 1) await update_trial(session, tg_id, 1)
if data.get("tariff_id"):
result = await session.execute( if tariff_id:
select(Tariff.price_rub).where(Tariff.id == data["tariff_id"]) result = await session.execute(select(Tariff.price_rub).where(Tariff.id == tariff_id))
)
row = result.scalar_one_or_none() row = result.scalar_one_or_none()
if row: if row:
await update_balance(session, tg_id, -row) await update_balance(session, tg_id, -row)
@@ -516,21 +512,13 @@ async def finalize_key_creation(
await session.commit() await session.commit()
except Exception as e: except Exception as e:
logger.error( logger.error(f"[Key Finalize] Ошибка при создании ключа для пользователя {tg_id}: {e}")
f"[Key Finalize] Ошибка при создании ключа для пользователя {tg_id}: {e}" await callback_query.message.answer("❌ Произошла ошибка при создании подписки. Попробуйте снова.")
)
await callback_query.message.answer(
"❌ Произошла ошибка при создании подписки. Попробуйте снова."
)
return return
builder = InlineKeyboardBuilder() builder = InlineKeyboardBuilder()
is_full_remnawave = await is_full_remnawave_cluster( is_full_remnawave = await is_full_remnawave_cluster(cluster_info["cluster_name"], session)
cluster_info["cluster_name"], session if (panel_type == "remnawave" or is_full_remnawave) and (public_link or remnawave_link):
)
if (panel_type == "remnawave" or is_full_remnawave) and (
public_link or remnawave_link
):
builder.row( builder.row(
InlineKeyboardButton( InlineKeyboardButton(
text=CONNECT_DEVICE, text=CONNECT_DEVICE,
@@ -539,9 +527,7 @@ async def finalize_key_creation(
) )
elif CONNECT_PHONE_BUTTON: elif CONNECT_PHONE_BUTTON:
builder.row( builder.row(
InlineKeyboardButton( InlineKeyboardButton(text=CONNECT_PHONE, callback_data=f"connect_phone|{key_name}")
text=CONNECT_PHONE, callback_data=f"connect_phone|{key_name}"
)
) )
builder.row( builder.row(
InlineKeyboardButton(text=PC_BUTTON, callback_data=f"connect_pc|{email}"), InlineKeyboardButton(text=PC_BUTTON, callback_data=f"connect_pc|{email}"),
@@ -549,15 +535,10 @@ async def finalize_key_creation(
) )
else: else:
builder.row( builder.row(
InlineKeyboardButton( InlineKeyboardButton(text=CONNECT_DEVICE, callback_data=f"connect_device|{key_name}")
text=CONNECT_DEVICE, callback_data=f"connect_device|{key_name}"
)
) )
builder.row(
InlineKeyboardButton( builder.row(InlineKeyboardButton(text=MY_SUB, callback_data=f"view_key|{key_name}"))
text="🔐 Моя подписка", callback_data=f"view_key|{key_name}"
)
)
builder.row(InlineKeyboardButton(text=SUPPORT, url=SUPPORT_CHAT_URL)) builder.row(InlineKeyboardButton(text=SUPPORT, url=SUPPORT_CHAT_URL))
builder.row(InlineKeyboardButton(text=MAIN_MENU, callback_data="profile")) builder.row(InlineKeyboardButton(text=MAIN_MENU, callback_data="profile"))
+3
View File
@@ -27,6 +27,7 @@ from handlers.buttons import MAIN_MENU, PAYMENT
from handlers.payments.robokassa_pay import handle_custom_amount_input from handlers.payments.robokassa_pay import handle_custom_amount_input
from handlers.payments.stars_pay import process_custom_amount_input_stars from handlers.payments.stars_pay import process_custom_amount_input_stars
from handlers.payments.yookassa_pay import process_custom_amount_input from handlers.payments.yookassa_pay import process_custom_amount_input
from handlers.payments.yoomoney_pay import process_custom_amount_input_yoomoney
from handlers.texts import ( from handlers.texts import (
CREATING_CONNECTION_MSG, CREATING_CONNECTION_MSG,
INSUFFICIENT_FUNDS_MSG, INSUFFICIENT_FUNDS_MSG,
@@ -165,6 +166,8 @@ async def select_tariff_plan(
await handle_custom_amount_input(callback_query, session) await handle_custom_amount_input(callback_query, session)
elif USE_NEW_PAYMENT_FLOW == "STARS": elif USE_NEW_PAYMENT_FLOW == "STARS":
await process_custom_amount_input_stars(callback_query, session) await process_custom_amount_input_stars(callback_query, session)
elif USE_NEW_PAYMENT_FLOW == "YOOMONEY":
await process_custom_amount_input_yoomoney(callback_query, session)
else: else:
builder = InlineKeyboardBuilder() builder = InlineKeyboardBuilder()
builder.row(InlineKeyboardButton(text=PAYMENT, callback_data="pay")) builder.row(InlineKeyboardButton(text=PAYMENT, callback_data="pay"))
+10 -2
View File
@@ -4,7 +4,7 @@ from typing import Any
from aiogram import F, Router from aiogram import F, Router
from aiogram.types import CallbackQuery, InlineKeyboardButton from aiogram.types import CallbackQuery, InlineKeyboardButton
from aiogram.utils.keyboard import InlineKeyboardBuilder from aiogram.utils.keyboard import InlineKeyboardBuilder
from sqlalchemy import or_, select from sqlalchemy import or_, select, update
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from bot import bot from bot import bot
@@ -19,12 +19,13 @@ from database import (
update_balance, update_balance,
update_key_expiry, update_key_expiry,
) )
from database.models import Server from database.models import Server, Key
from handlers.buttons import BACK, MAIN_MENU, PAYMENT from handlers.buttons import BACK, MAIN_MENU, PAYMENT
from handlers.keys.key_utils import renew_key_in_cluster from handlers.keys.key_utils import renew_key_in_cluster
from handlers.payments.robokassa_pay import handle_custom_amount_input from handlers.payments.robokassa_pay import handle_custom_amount_input
from handlers.payments.stars_pay import process_custom_amount_input_stars from handlers.payments.stars_pay import process_custom_amount_input_stars
from handlers.payments.yookassa_pay import process_custom_amount_input from handlers.payments.yookassa_pay import process_custom_amount_input
from handlers.payments.yoomoney_pay import process_custom_amount_input_yoomoney
from handlers.texts import ( from handlers.texts import (
INSUFFICIENT_FUNDS_RENEWAL_MSG, INSUFFICIENT_FUNDS_RENEWAL_MSG,
KEY_NOT_FOUND_MSG, KEY_NOT_FOUND_MSG,
@@ -191,6 +192,8 @@ async def process_callback_renew_plan(callback_query: CallbackQuery, session: An
await handle_custom_amount_input(callback_query, session) await handle_custom_amount_input(callback_query, session)
elif USE_NEW_PAYMENT_FLOW == "STARS": elif USE_NEW_PAYMENT_FLOW == "STARS":
await process_custom_amount_input_stars(callback_query, session) await process_custom_amount_input_stars(callback_query, session)
elif USE_NEW_PAYMENT_FLOW == "YOOMONEY":
await process_custom_amount_input_yoomoney(callback_query, session)
else: else:
builder = InlineKeyboardBuilder() builder = InlineKeyboardBuilder()
builder.row(InlineKeyboardButton(text=PAYMENT, callback_data="pay")) builder.row(InlineKeyboardButton(text=PAYMENT, callback_data="pay"))
@@ -316,6 +319,11 @@ async def complete_key_renewal(
) )
await update_key_expiry(session, client_id, new_expiry_time) await update_key_expiry(session, client_id, new_expiry_time)
await session.execute(
update(Key)
.where(Key.client_id == client_id)
.values(tariff_id=tariff_id)
)
await update_balance(session, tg_id, -cost) await update_balance(session, tg_id, -cost)
logger.info( logger.info(
+72 -47
View File
@@ -5,7 +5,7 @@ from typing import Any
from sqlalchemy import delete, select from sqlalchemy import delete, select
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from config import PUBLIC_LINK, REMNAWAVE_LOGIN, REMNAWAVE_PASSWORD, SUPERNODE from config import PUBLIC_LINK, REMNAWAVE_LOGIN, REMNAWAVE_PASSWORD, SUPERNODE, TRIAL_CONFIG
from database import delete_notification, get_servers, get_tariff_by_id, store_key from database import delete_notification, get_servers, get_tariff_by_id, store_key
from database.models import Key, Server, Tariff from database.models import Key, Server, Tariff
from handlers.utils import check_server_key_limit, get_least_loaded_cluster from handlers.utils import check_server_key_limit, get_least_loaded_cluster
@@ -32,7 +32,8 @@ async def create_key_on_cluster(
session: AsyncSession = None, session: AsyncSession = None,
remnawave_link: str = None, remnawave_link: str = None,
hwid_limit: int = None, hwid_limit: int = None,
traffic_limit_bytes: int = None, traffic_limit_bytes: int = None,
is_trial: bool = False,
): ):
try: try:
servers = await get_servers(session, include_enabled=True) servers = await get_servers(session, include_enabled=True)
@@ -125,7 +126,7 @@ async def create_key_on_cluster(
} }
if traffic_limit_bytes and traffic_limit_bytes > 0: if traffic_limit_bytes and traffic_limit_bytes > 0:
user_data["trafficLimitBytes"] = traffic_limit_bytes user_data["trafficLimitBytes"] = traffic_limit_bytes * 1024 * 1024 * 1024
if short_uuid: if short_uuid:
user_data["shortUuid"] = short_uuid user_data["shortUuid"] = short_uuid
@@ -161,6 +162,7 @@ async def create_key_on_cluster(
semaphore, semaphore,
plan=plan, plan=plan,
session=session, session=session,
is_trial=is_trial,
) )
else: else:
await asyncio.gather( await asyncio.gather(
@@ -174,6 +176,7 @@ async def create_key_on_cluster(
semaphore, semaphore,
plan=plan, plan=plan,
session=session, session=session,
is_trial=is_trial,
) )
for server in xui_servers for server in xui_servers
], ],
@@ -207,12 +210,13 @@ async def create_client_on_server(
semaphore: asyncio.Semaphore, semaphore: asyncio.Semaphore,
plan: int = None, plan: int = None,
session=None, session=None,
is_trial: bool = False,
): ):
""" """
Создает клиента на указанном 3x-ui сервере с лимитом по тарифу (через ORM). Создает клиента на указанном 3x-ui сервере с лимитом по тарифу или триалу.
""" """
logger.info( logger.info(
f"[Client] Вход в create_client_on_server: сервер={server_info.get('server_name')}, план={plan}" f"[Client] Вход в create_client_on_server: сервер={server_info.get('server_name')}, план={plan}, is_trial={is_trial}"
) )
async with semaphore: async with semaphore:
@@ -221,9 +225,7 @@ async def create_client_on_server(
server_name = server_info.get("server_name", "unknown") server_name = server_info.get("server_name", "unknown")
if not inbound_id: if not inbound_id:
logger.warning( logger.warning(f"[Client] INBOUND_ID отсутствует для сервера {server_name}. Пропуск.")
f"INBOUND_ID отсутствует для сервера {server_name}. Пропуск."
)
return return
if SUPERNODE: if SUPERNODE:
@@ -236,26 +238,24 @@ async def create_client_on_server(
total_gb_value = 0 total_gb_value = 0
device_limit_value = None device_limit_value = None
if plan is not None: if is_trial:
total_gb_value = TRIAL_CONFIG.get("traffic_limit_gb", 0)
device_limit_value = TRIAL_CONFIG.get("hwid_limit")
logger.info(f"[Trial] Используются параметры триала: {total_gb_value} GB, {device_limit_value} устройств")
elif plan is not None:
tariff = await get_tariff_by_id(session, plan) tariff = await get_tariff_by_id(session, plan)
logger.info(f"[Tariff Debug] Получен тариф: {tariff}") logger.info(f"[Tariff Debug] Получен тариф: {tariff}")
if not tariff: if not tariff:
raise ValueError(f"Тариф с id={plan} не найден.") raise ValueError(f"Тариф с id={plan} не найден.")
total_gb_value = ( total_gb_value = int(tariff["traffic_limit"]) if tariff["traffic_limit"] else 0
int(tariff["traffic_limit"]) if tariff.get("traffic_limit") else 0 device_limit_value = int(tariff["device_limit"]) if tariff.get("device_limit") is not None else None
)
device_limit_value = (
int(tariff["device_limit"])
if tariff.get("device_limit") is not None
else None
)
try: try:
logger.info( logger.info(
f"[Client] Вызов add_client: email={unique_email}, client_id={client_id}, GB={total_gb_value}, Devices={device_limit_value}" f"[Client] Вызов add_client: email={email}, client_id={client_id}, GB={total_gb_value}, Devices={device_limit_value}"
) )
traffic_limit_bytes = total_gb_value * 1024 * 1024 * 1024
await add_client( await add_client(
xui, xui,
ClientConfig( ClientConfig(
@@ -263,7 +263,7 @@ async def create_client_on_server(
email=unique_email, email=unique_email,
tg_id=tg_id, tg_id=tg_id,
limit_ip=device_limit_value, limit_ip=device_limit_value,
total_gb=total_gb_value, total_gb=traffic_limit_bytes,
expiry_time=expiry_timestamp, expiry_time=expiry_timestamp,
enable=True, enable=True,
flow="xtls-rprx-vision", flow="xtls-rprx-vision",
@@ -273,9 +273,7 @@ async def create_client_on_server(
) )
logger.info(f"[Client] Клиент успешно добавлен на сервер {server_name}") logger.info(f"[Client] Клиент успешно добавлен на сервер {server_name}")
except Exception as e: except Exception as e:
logger.error( logger.error(f"[Client Error] Не удалось создать клиента на {server_name}: {e}")
f"[Client Error] Не удалось создать клиента на {server_name}: {e}"
)
if SUPERNODE: if SUPERNODE:
await asyncio.sleep(0.7) await asyncio.sleep(0.7)
@@ -335,12 +333,6 @@ async def renew_key_in_cluster(
if tariff and tariff.device_limit is not None: if tariff and tariff.device_limit is not None:
hwid_device_limit = int(tariff.device_limit) hwid_device_limit = int(tariff.device_limit)
notification_prefixes = ["key_24h", "key_10h", "key_expired", "renew"]
for notif in notification_prefixes:
notification_id = f"{email}_{notif}"
await delete_notification(session, tg_id, notification_id)
logger.info(f"🧹 Уведомления для ключа {email} очищены при продлении.")
remnawave_inbound_ids = [] remnawave_inbound_ids = []
tasks = [] tasks = []
for server_info in cluster: for server_info in cluster:
@@ -366,11 +358,12 @@ async def renew_key_in_cluster(
datetime.utcfromtimestamp(new_expiry_time // 1000).isoformat() datetime.utcfromtimestamp(new_expiry_time // 1000).isoformat()
+ "Z" + "Z"
) )
traffic_limit_bytes = total_gb * 1024 * 1024 * 1024 if total_gb else 0
updated = await remna.update_user( updated = await remna.update_user(
uuid=client_id, uuid=client_id,
expire_at=expire_iso, expire_at=expire_iso,
active_user_inbounds=remnawave_inbound_ids, active_user_inbounds=remnawave_inbound_ids,
traffic_limit_bytes=total_gb, traffic_limit_bytes=traffic_limit_bytes,
hwid_device_limit=hwid_device_limit, hwid_device_limit=hwid_device_limit,
) )
if updated: if updated:
@@ -404,6 +397,7 @@ async def renew_key_in_cluster(
unique_email = email unique_email = email
sub_id = unique_email sub_id = unique_email
traffic_bytes = total_gb * 1024 * 1024 * 1024 if total_gb else 0
tasks.append( tasks.append(
extend_client_key( extend_client_key(
xui=xui, xui=xui,
@@ -411,7 +405,7 @@ async def renew_key_in_cluster(
email=unique_email, email=unique_email,
new_expiry_time=new_expiry_time, new_expiry_time=new_expiry_time,
client_id=client_id, client_id=client_id,
total_gb=total_gb, total_gb=traffic_bytes,
sub_id=sub_id, sub_id=sub_id,
tg_id=tg_id, tg_id=tg_id,
limit_ip=hwid_device_limit, limit_ip=hwid_device_limit,
@@ -420,6 +414,12 @@ async def renew_key_in_cluster(
await asyncio.gather(*tasks, return_exceptions=True) await asyncio.gather(*tasks, return_exceptions=True)
notification_prefixes = ["key_24h", "key_10h", "key_expired", "renew"]
for notif in notification_prefixes:
notification_id = f"{email}_{notif}"
await delete_notification(session, tg_id, notification_id)
logger.info(f"🧹 Уведомления для ключа {email} очищены при продлении.")
except Exception as e: except Exception as e:
logger.error( logger.error(
f"Не удалось продлить ключ {client_id} в кластере/на сервере {cluster_id}: {e}" f"Не удалось продлить ключ {client_id} в кластере/на сервере {cluster_id}: {e}"
@@ -510,6 +510,9 @@ async def update_key_on_cluster(
expiry_time: int, expiry_time: int,
cluster_id: str, cluster_id: str,
session: AsyncSession, session: AsyncSession,
traffic_limit: int = None,
device_limit: int = None,
remnawave_link: str = None,
): ):
""" """
Пересоздаёт ключ на всех серверах указанного кластера (или сервера, если передано имя). Пересоздаёт ключ на всех серверах указанного кластера (или сервера, если передано имя).
@@ -568,6 +571,11 @@ async def update_key_on_cluster(
) )
tariff = result.scalar_one_or_none() tariff = result.scalar_one_or_none()
short_uuid = None
if remnawave_link and "/" in remnawave_link:
short_uuid = remnawave_link.rstrip("/").split("/")[-1]
logger.info(f"[Update] Извлечен short_uuid из ссылки: {short_uuid}")
user_data = { user_data = {
"username": email, "username": email,
"trafficLimitStrategy": "NO_RESET", "trafficLimitStrategy": "NO_RESET",
@@ -575,12 +583,13 @@ async def update_key_on_cluster(
"telegramId": tg_id, "telegramId": tg_id,
"activeUserInbounds": inbound_ids, "activeUserInbounds": inbound_ids,
} }
if traffic_limit is not None:
if tariff: user_data["trafficLimitBytes"] = traffic_limit
if tariff.traffic_limit is not None: if device_limit is not None:
user_data["trafficLimitBytes"] = int(tariff.traffic_limit) user_data["hwidDeviceLimit"] = device_limit
if tariff.device_limit is not None: if short_uuid:
user_data["hwidDeviceLimit"] = int(tariff.device_limit) user_data["shortUuid"] = short_uuid
logger.info(f"[Update] Добавлен short_uuid в user_data: {short_uuid}")
result = await remna.create_user(user_data) result = await remna.create_user(user_data)
if result: if result:
@@ -628,20 +637,14 @@ async def update_key_on_cluster(
) )
tariff = result.scalar_one_or_none() tariff = result.scalar_one_or_none()
total_gb_bytes = ( total_gb_bytes = int(traffic_limit * 1024 ** 3) if traffic_limit else 0
int(tariff.traffic_limit) if tariff and tariff.traffic_limit else 0 device_limit_value = device_limit if device_limit is not None else None
)
device_limit = (
int(tariff.device_limit)
if tariff and tariff.device_limit is not None
else None
)
config = ClientConfig( config = ClientConfig(
client_id=remnawave_client_id, client_id=remnawave_client_id,
email=unique_email, email=unique_email,
tg_id=tg_id, tg_id=tg_id,
limit_ip=device_limit, limit_ip=device_limit_value,
total_gb=total_gb_bytes, total_gb=total_gb_bytes,
expiry_time=expiry_time, expiry_time=expiry_time,
enable=True, enable=True,
@@ -673,6 +676,7 @@ async def update_subscription(
session: AsyncSession, session: AsyncSession,
cluster_override: str = None, cluster_override: str = None,
country_override: str = None, country_override: str = None,
remnawave_link: str = None,
) -> None: ) -> None:
result = await session.execute( result = await session.execute(
select(Key).where(Key.tg_id == tg_id, Key.email == email) select(Key).where(Key.tg_id == tg_id, Key.email == email)
@@ -685,8 +689,25 @@ async def update_subscription(
expiry_time = record.expiry_time expiry_time = record.expiry_time
client_id = record.client_id client_id = record.client_id
old_cluster_id = record.server_id old_cluster_id = record.server_id
tariff_id = record.tariff_id
remnawave_link = remnawave_link or record.remnawave_link
public_link = f"{PUBLIC_LINK}{email}/{tg_id}" public_link = f"{PUBLIC_LINK}{email}/{tg_id}"
traffic_limit = None
device_limit = None
if tariff_id:
result = await session.execute(
select(Tariff).where(Tariff.id == tariff_id, Tariff.is_active.is_(True))
)
tariff = result.scalar_one_or_none()
if tariff:
traffic_limit = int(tariff.traffic_limit) if tariff.traffic_limit is not None else None
device_limit = int(tariff.device_limit) if tariff.device_limit is not None else None
else:
logger.warning(f"[LOG] update_subscription: тариф с id={tariff_id} не найден!")
else:
logger.warning(f"[LOG] update_subscription: tariff_id отсутствует!")
await delete_key_from_cluster(old_cluster_id, email, client_id, session=session) await delete_key_from_cluster(old_cluster_id, email, client_id, session=session)
await session.execute(delete(Key).where(Key.tg_id == tg_id, Key.email == email)) await session.execute(delete(Key).where(Key.tg_id == tg_id, Key.email == email))
@@ -703,6 +724,9 @@ async def update_subscription(
expiry_time=expiry_time, expiry_time=expiry_time,
cluster_id=new_cluster_id, cluster_id=new_cluster_id,
session=session, session=session,
traffic_limit=traffic_limit,
device_limit=device_limit,
remnawave_link=remnawave_link
) )
servers = await get_servers(session) servers = await get_servers(session)
@@ -720,6 +744,7 @@ async def update_subscription(
key=final_key_link, key=final_key_link,
remnawave_link=remnawave_key, remnawave_link=remnawave_key,
server_id=new_cluster_id, server_id=new_cluster_id,
tariff_id=tariff_id,
) )
@@ -1012,4 +1037,4 @@ async def reset_traffic_in_cluster(
logger.error( logger.error(
f"[Reset Traffic] Ошибка при сбросе трафика клиента {email} в кластере {cluster_id}: {e}" f"[Reset Traffic] Ошибка при сбросе трафика клиента {email} в кластере {cluster_id}: {e}"
) )
raise raise
+7 -15
View File
@@ -76,6 +76,7 @@ async def process_callback_confirm_delete(
record = await get_key_details(session, email) record = await get_key_details(session, email)
if record: if record:
client_id = record["client_id"] client_id = record["client_id"]
server_id = record["server_id"]
response_message = KEY_DELETED_MSG_SIMPLE response_message = KEY_DELETED_MSG_SIMPLE
back_button = types.InlineKeyboardButton( back_button = types.InlineKeyboardButton(
text=BACK, callback_data="view_keys" text=BACK, callback_data="view_keys"
@@ -90,21 +91,12 @@ async def process_callback_confirm_delete(
reply_markup=keyboard, reply_markup=keyboard,
) )
servers = await get_servers(session) try:
await delete_key_from_cluster(server_id, email, client_id, session)
async def delete_key_from_servers(): except Exception as e:
try: logger.error(
tasks = [ f"Ошибка при удалении ключа {client_id} с сервера {server_id}: {e}"
delete_key_from_cluster(cluster_id, email, client_id, session) )
for cluster_id in servers
]
await asyncio.gather(*tasks, return_exceptions=True)
except Exception as e:
logger.error(
f"Ошибка при удалении ключа {client_id} с серверов: {e}"
)
asyncio.create_task(delete_key_from_servers())
else: else:
response_message = "Ключ не найден или уже удален." response_message = "Ключ не найден или уже удален."
@@ -597,4 +597,4 @@ async def process_auto_renew_or_notify(
) )
except Exception as e: except Exception as e:
logger.error(f"❌ Ошибка в process_auto_renew_or_notify: {e}") logger.error(f"❌ Ошибка в process_auto_renew_or_notify: {e}")