ORM update and lots of improvements
This commit is contained in:
+88
-65
@@ -1,20 +1,23 @@
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
import secrets
|
||||
import string
|
||||
|
||||
from datetime import datetime
|
||||
|
||||
import aiofiles
|
||||
import aiohttp
|
||||
import asyncpg
|
||||
|
||||
from aiogram.types import BufferedInputFile, InlineKeyboardMarkup, InputMediaPhoto, Message
|
||||
from aiogram.types import (
|
||||
BufferedInputFile,
|
||||
InlineKeyboardMarkup,
|
||||
InputMediaPhoto,
|
||||
Message,
|
||||
)
|
||||
from sqlalchemy import func, select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from bot import bot
|
||||
from config import ADMIN_ID, DATABASE_URL
|
||||
from database import get_all_keys, get_servers
|
||||
from config import ADMIN_ID
|
||||
from database import get_servers
|
||||
from database.models import Key, Server
|
||||
from logger import logger
|
||||
|
||||
|
||||
@@ -22,14 +25,18 @@ def generate_random_email(length: int = 8) -> str:
|
||||
"""
|
||||
Генерирует случайный email с заданной длиной.
|
||||
"""
|
||||
return "".join(secrets.choice(string.ascii_lowercase + string.digits) for _ in range(length)) if length > 0 else ""
|
||||
return (
|
||||
"".join(
|
||||
secrets.choice(string.ascii_lowercase + string.digits)
|
||||
for _ in range(length)
|
||||
)
|
||||
if length > 0
|
||||
else ""
|
||||
)
|
||||
|
||||
|
||||
async def get_least_loaded_cluster() -> str:
|
||||
"""
|
||||
Возвращает кластер с наименьшей загрузкой, где есть хотя бы один сервер с доступным лимитом.
|
||||
"""
|
||||
servers = await get_servers()
|
||||
async def get_least_loaded_cluster(session: AsyncSession) -> str:
|
||||
servers = await get_servers(session)
|
||||
server_to_cluster = {}
|
||||
cluster_loads = {}
|
||||
|
||||
@@ -38,36 +45,40 @@ async def get_least_loaded_cluster() -> str:
|
||||
for server in cluster_servers:
|
||||
server_to_cluster[server["server_name"]] = cluster_name
|
||||
|
||||
async with asyncpg.create_pool(DATABASE_URL) as pool:
|
||||
async with pool.acquire() as conn:
|
||||
keys = await get_all_keys(conn)
|
||||
for key in keys:
|
||||
server_id = key["server_id"]
|
||||
cluster_id = server_to_cluster.get(server_id, server_id)
|
||||
if cluster_id in cluster_loads:
|
||||
cluster_loads[cluster_id] += 1
|
||||
result = await session.execute(select(Key))
|
||||
keys = result.scalars().all()
|
||||
|
||||
available_clusters = {}
|
||||
for cluster_name, cluster_servers in servers.items():
|
||||
for server in cluster_servers:
|
||||
if server.get("enabled", True) and await check_server_key_limit(server, conn):
|
||||
available_clusters[cluster_name] = cluster_loads[cluster_name]
|
||||
break
|
||||
for key in keys:
|
||||
server_id = key.server_id
|
||||
cluster_id = server_to_cluster.get(server_id, server_id)
|
||||
if cluster_id in cluster_loads:
|
||||
cluster_loads[cluster_id] += 1
|
||||
|
||||
available_clusters = {}
|
||||
for cluster_name, cluster_servers in servers.items():
|
||||
for server in cluster_servers:
|
||||
if server.get("enabled", True) and await check_server_key_limit(
|
||||
server, session
|
||||
):
|
||||
available_clusters[cluster_name] = cluster_loads[cluster_name]
|
||||
break
|
||||
|
||||
if not available_clusters:
|
||||
logger.warning("❌ Нет доступных кластеров с лимитом ключей!")
|
||||
return "cluster1"
|
||||
|
||||
least_loaded_cluster = min(available_clusters, key=lambda k: (available_clusters[k], k))
|
||||
logger.info(f"✅ Выбран наименее загруженный кластер с лимитом: {least_loaded_cluster}")
|
||||
|
||||
least_loaded_cluster = min(
|
||||
available_clusters, key=lambda k: (available_clusters[k], k)
|
||||
)
|
||||
logger.info(
|
||||
f"✅ Выбран наименее загруженный кластер с лимитом: {least_loaded_cluster}"
|
||||
)
|
||||
return least_loaded_cluster
|
||||
|
||||
|
||||
async def check_server_key_limit(server_info: dict, conn) -> bool:
|
||||
"""
|
||||
Универсальная проверка лимита ключей для сервера в режимах кластеров и стран.
|
||||
"""
|
||||
async def check_server_key_limit(server_info: dict, session: AsyncSession) -> bool:
|
||||
from database.models import Key, Notification
|
||||
|
||||
server_name = server_info.get("server_name")
|
||||
cluster_name = server_info.get("cluster_name")
|
||||
max_keys = server_info.get("max_keys")
|
||||
@@ -76,19 +87,30 @@ async def check_server_key_limit(server_info: dict, conn) -> bool:
|
||||
return True
|
||||
|
||||
identifier = cluster_name if cluster_name else server_name
|
||||
total_keys = await conn.fetchval("SELECT COUNT(*) FROM keys WHERE server_id = $1", identifier)
|
||||
|
||||
result = await session.execute(
|
||||
select(func.count()).select_from(Key).where(Key.server_id == identifier)
|
||||
)
|
||||
total_keys = result.scalar() or 0
|
||||
|
||||
if total_keys >= max_keys:
|
||||
logger.warning(f"[Key Limit] Сервер {server_name} достиг лимита: {total_keys}/{max_keys}")
|
||||
logger.warning(
|
||||
f"[Key Limit] Сервер {server_name} достиг лимита: {total_keys}/{max_keys}"
|
||||
)
|
||||
return False
|
||||
|
||||
usage_percent = total_keys / max_keys
|
||||
|
||||
if usage_percent >= 0.9:
|
||||
notif_key = f"server_warn_{server_name}"
|
||||
already_sent = await conn.fetchval(
|
||||
"SELECT EXISTS (SELECT 1 FROM notifications WHERE tg_id = 0 AND notification_type = $1)", notif_key
|
||||
|
||||
result = await session.execute(
|
||||
select(Notification).where(
|
||||
Notification.tg_id == 0, Notification.notification_type == notif_key
|
||||
)
|
||||
)
|
||||
already_sent = result.scalar_one_or_none()
|
||||
|
||||
if not already_sent:
|
||||
for admin_id in ADMIN_ID:
|
||||
try:
|
||||
@@ -99,26 +121,29 @@ async def check_server_key_limit(server_info: dict, conn) -> bool:
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
await conn.execute(
|
||||
"INSERT INTO notifications (tg_id, notification_type) VALUES (0, $1) ON CONFLICT DO NOTHING",
|
||||
notif_key,
|
||||
)
|
||||
|
||||
session.add(Notification(tg_id=0, notification_type=notif_key))
|
||||
await session.commit()
|
||||
|
||||
return True
|
||||
|
||||
|
||||
async def handle_error(tg_id: int, callback_query: object | None = None, message: str = "") -> None:
|
||||
async def handle_error(
|
||||
tg_id: int, callback_query: object | None = None, message: str = ""
|
||||
) -> None:
|
||||
"""
|
||||
Обрабатывает ошибку, отправляя сообщение пользователю.
|
||||
"""
|
||||
try:
|
||||
if callback_query and hasattr(callback_query, "message"):
|
||||
try:
|
||||
await bot.delete_message(chat_id=tg_id, message_id=callback_query.message.message_id)
|
||||
await bot.delete_message(
|
||||
chat_id=tg_id, message_id=callback_query.message.message_id
|
||||
)
|
||||
except Exception as delete_error:
|
||||
logger.warning(f"Не удалось удалить сообщение: {delete_error}")
|
||||
|
||||
await bot.send_message(tg_id, message)
|
||||
await bot.send_message(tg_id, message, parse_mode=None)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Ошибка при обработке ошибки: {e}")
|
||||
@@ -179,7 +204,7 @@ def format_minutes(minutes: int) -> str:
|
||||
async def edit_or_send_message(
|
||||
target_message: Message,
|
||||
text: str,
|
||||
reply_markup: InlineKeyboardMarkup,
|
||||
reply_markup: InlineKeyboardMarkup | None = None,
|
||||
media_path: str = None,
|
||||
disable_web_page_preview: bool = False,
|
||||
force_text: bool = False,
|
||||
@@ -199,7 +224,9 @@ async def edit_or_send_message(
|
||||
return
|
||||
except Exception:
|
||||
await target_message.answer_photo(
|
||||
photo=BufferedInputFile(image_data, filename=os.path.basename(media_path)),
|
||||
photo=BufferedInputFile(
|
||||
image_data, filename=os.path.basename(media_path)
|
||||
),
|
||||
caption=text,
|
||||
reply_markup=reply_markup,
|
||||
disable_web_page_preview=disable_web_page_preview,
|
||||
@@ -208,7 +235,9 @@ async def edit_or_send_message(
|
||||
else:
|
||||
if not force_text and target_message.caption is not None:
|
||||
try:
|
||||
await target_message.edit_caption(caption=text, reply_markup=reply_markup)
|
||||
await target_message.edit_caption(
|
||||
caption=text, reply_markup=reply_markup
|
||||
)
|
||||
return
|
||||
except Exception as e:
|
||||
logger.error(f"Ошибка редактирования подписи: {e}")
|
||||
@@ -239,26 +268,20 @@ def convert_to_bytes(value: float, unit: str) -> int:
|
||||
return int(value * units.get(unit.upper(), 1))
|
||||
|
||||
|
||||
async def is_full_remnawave_cluster(cluster_id: str, session) -> bool:
|
||||
"""
|
||||
Универсальная проверка:
|
||||
- Если cluster_id — это имя кластера, проверяет, что все его сервера используют Remnawave.
|
||||
- Если cluster_id — это имя одиночного сервера, проверяет, что он Remnawave.
|
||||
"""
|
||||
cluster_servers = await session.fetch(
|
||||
"SELECT panel_type FROM servers WHERE cluster_name = $1",
|
||||
cluster_id,
|
||||
async def is_full_remnawave_cluster(cluster_id: str, session: AsyncSession) -> bool:
|
||||
result = await session.execute(
|
||||
select(Server.panel_type).where(Server.cluster_name == cluster_id)
|
||||
)
|
||||
panel_types = result.scalars().all()
|
||||
|
||||
if cluster_servers:
|
||||
panel_types = [s["panel_type"].lower() for s in cluster_servers if s.get("panel_type")]
|
||||
return all(pt == "remnawave" for pt in panel_types)
|
||||
if panel_types:
|
||||
return all(pt.lower() == "remnawave" for pt in panel_types)
|
||||
|
||||
server = await session.fetchrow(
|
||||
"SELECT panel_type FROM servers WHERE server_name = $1",
|
||||
cluster_id,
|
||||
result = await session.execute(
|
||||
select(Server.panel_type).where(Server.server_name == cluster_id)
|
||||
)
|
||||
return server and server["panel_type"].lower() == "remnawave"
|
||||
panel_type = result.scalar_one_or_none()
|
||||
return panel_type and panel_type.lower() == "remnawave"
|
||||
|
||||
|
||||
def sanitize_key_name(key_name: str) -> str:
|
||||
|
||||
Reference in New Issue
Block a user