from sqlalchemy import delete, func, insert, select, update from sqlalchemy.exc import SQLAlchemyError from sqlalchemy.ext.asyncio import AsyncSession from core.cache_config import SERVERS_CACHE_TTL_SEC from core.redis_cache import cache_delete_pattern, cache_get, cache_key, cache_set from database.models import Key, Server, ServerSpecialgroup, ServerSubgroup, Tariff from logger import logger async def _invalidate_servers_cache() -> None: await cache_delete_pattern("servers:*") async def create_server( session: AsyncSession, cluster_name: str, server_name: str, api_url: str, subscription_url: str, inbound_id: str, ): try: stmt = insert(Server).values( cluster_name=cluster_name, server_name=server_name, api_url=api_url, subscription_url=subscription_url, inbound_id=inbound_id, ) await session.execute(stmt) await session.commit() await _invalidate_servers_cache() logger.info(f"✅ Сервер {server_name} добавлен в кластер {cluster_name}") except SQLAlchemyError as e: logger.error(f"❌ Ошибка при добавлении сервера {server_name}: {e}") await session.rollback() raise async def delete_server(session: AsyncSession, server_name: str): try: stmt = delete(Server).where(Server.server_name == server_name) await session.execute(stmt) await session.commit() await _invalidate_servers_cache() logger.info(f"🗑 Сервер {server_name} удалён") except SQLAlchemyError as e: logger.error(f"❌ Ошибка при удалении сервера {server_name}: {e}") await session.rollback() raise async def get_servers(session: AsyncSession, include_enabled: bool = False) -> dict: from handlers.utils import ALLOWED_GROUP_CODES cache_key_servers = cache_key("servers", int(include_enabled)) cached = await cache_get(cache_key_servers) if isinstance(cached, dict): return cached try: stmt = select(Server) result = await session.execute(stmt) servers = result.scalars().all() ids = [s.id for s in servers] subs_map = {} tariffs_map = {} if ids: r = await session.execute( select(ServerSubgroup.server_id, ServerSubgroup.subgroup_title).where(ServerSubgroup.server_id.in_(ids)) ) for sid, sg in r.all(): if sg and sg.isdigit(): tariffs_map.setdefault(sid, []).append(int(sg)) else: subs_map.setdefault(sid, []).append(sg) groups_map = {} if ids: r2 = await session.execute( select(ServerSpecialgroup.server_id, ServerSpecialgroup.group_code).where( ServerSpecialgroup.server_id.in_(ids) ) ) for sid, gc in r2.all(): groups_map.setdefault(sid, []).append(gc) allowed = set(ALLOWED_GROUP_CODES) grouped = {} for s in servers: if not include_enabled and not s.enabled: continue cluster = s.cluster_name special = sorted({g for g in groups_map.get(s.id, []) if g in allowed}) grouped.setdefault(cluster, []).append({ "server_name": s.server_name, "api_url": s.api_url, "subscription_url": s.subscription_url, "inbound_id": s.inbound_id, "panel_type": s.panel_type, "enabled": s.enabled, "max_keys": s.max_keys, "tariff_group": s.tariff_group, "tariff_subgroups": subs_map.get(s.id, []), "tariff_ids": tariffs_map.get(s.id, []), "special_groups": special, "cluster_name": cluster, "server_id": s.id, }) await cache_set(cache_key_servers, grouped, SERVERS_CACHE_TTL_SEC) return grouped except SQLAlchemyError as e: logger.error(f"Ошибка при получении серверов: {e}") await session.rollback() return {} async def get_clusters(session: AsyncSession) -> list[str]: stmt = select(Server.cluster_name).distinct().order_by(Server.cluster_name) result = await session.execute(stmt) return [r[0] for r in result.all()] async def check_unique_server_name(session: AsyncSession, server_name: str, cluster_name: str | None = None) -> bool: stmt = select(Server).where(Server.server_name == server_name) if cluster_name: stmt = stmt.where(Server.cluster_name == cluster_name) result = await session.execute(stmt.limit(1)) return result.scalar_one_or_none() is None async def check_server_name_by_cluster(session: AsyncSession, server_name: str) -> dict | None: try: result = await session.execute(select(Server.cluster_name).where(Server.server_name == server_name)) row = result.first() return {"cluster_name": row[0]} if row else None except SQLAlchemyError as e: logger.error(f"Ошибка при поиске кластера для сервера {server_name}: {e}") await session.rollback() return None async def get_cluster_name_by_server(session: AsyncSession, server_id_or_name: str) -> str | None: stmt = ( select(Server.cluster_name) .where((Server.id == server_id_or_name) | (Server.server_name == server_id_or_name)) .limit(1) ) result = await session.execute(stmt) row = result.scalar_one_or_none() return row async def get_server_by_name(session: AsyncSession, server_name: str) -> dict | None: try: stmt = select(Server).where(Server.server_name == server_name) result = await session.execute(stmt) server = result.scalar_one_or_none() if server: return { "id": server.id, "cluster_name": server.cluster_name, "server_name": server.server_name, "api_url": server.api_url, "subscription_url": server.subscription_url, "inbound_id": server.inbound_id, "panel_type": server.panel_type, "enabled": server.enabled, "max_keys": server.max_keys, "tariff_group": server.tariff_group, } return None except SQLAlchemyError as e: logger.error(f"Ошибка при получении сервера {server_name}: {e}") await session.rollback() return None async def update_server_field(session: AsyncSession, server_name: str, field: str, value: any) -> bool: try: stmt = update(Server).where(Server.server_name == server_name).values(**{field: value}) await session.execute(stmt) await session.commit() await _invalidate_servers_cache() logger.info(f"✅ Поле {field} сервера {server_name} обновлено на {value}") return True except SQLAlchemyError as e: logger.error(f"❌ Ошибка при обновлении поля {field} сервера {server_name}: {e}") await session.rollback() return False async def update_server_name_with_keys(session: AsyncSession, old_name: str, new_name: str) -> bool: try: from sqlalchemy import update from database.models import Key if not await check_unique_server_name(session, new_name): logger.error(f"❌ Сервер с именем {new_name} уже существует") return False stmt_server = update(Server).where(Server.server_name == old_name).values(server_name=new_name) await session.execute(stmt_server) stmt_keys = update(Key).where(Key.server_id == old_name).values(server_id=new_name) await session.execute(stmt_keys) await session.commit() await _invalidate_servers_cache() logger.info(f"✅ Сервер переименован с {old_name} на {new_name}") return True except SQLAlchemyError as e: logger.error(f"❌ Ошибка при переименовании сервера {old_name}: {e}") await session.rollback() return False async def get_available_clusters(session: AsyncSession) -> list[str]: try: stmt = select(Server.cluster_name).distinct().order_by(Server.cluster_name) result = await session.execute(stmt) return [row[0] for row in result.all()] except SQLAlchemyError as e: logger.error(f"Ошибка при получении списка кластеров: {e}") await session.rollback() return [] async def update_server_cluster(session: AsyncSession, server_name: str, new_cluster: str) -> bool: try: server_data = await get_server_by_name(session, server_name) if not server_data: return False old_cluster = server_data["cluster_name"] stmt_remaining = select(func.count()).where( (Server.cluster_name == old_cluster) & (Server.server_name != server_name) ) result = await session.execute(stmt_remaining) remaining_servers = result.scalar_one() if remaining_servers == 0: stmt_update_keys = update(Key).where(Key.server_id == old_cluster).values(server_id=new_cluster) await session.execute(stmt_update_keys) stmt_new_cluster = select(Server.tariff_group).where(Server.cluster_name == new_cluster).limit(1) result = await session.execute(stmt_new_cluster) new_tariff_group = result.scalar_one_or_none() await session.execute( update(Server) .where(Server.server_name == server_name) .values(cluster_name=new_cluster, tariff_group=new_tariff_group) ) if server_data.get("id") is None: rid = await session.execute(select(Server.id).where(Server.server_name == server_name).limit(1)) server_id = rid.scalar_one_or_none() else: server_id = server_data["id"] if server_id is not None and new_tariff_group is not None: await session.execute( update(ServerSubgroup).where(ServerSubgroup.server_id == server_id).values(group_code=new_tariff_group) ) await session.commit() await _invalidate_servers_cache() logger.info( f"✅ Сервер {server_name} перемещен в кластер {new_cluster} с обновлением тарифной группы и привязок подгрупп" ) return True except SQLAlchemyError as e: logger.error(f"❌ Ошибка при обновлении кластера сервера {server_name}: {e}") await session.rollback() return False async def resolve_device_limit_from_group(session: AsyncSession, server_id: str) -> int | None: r = await session.execute(select(Server.tariff_group).where(Server.server_name == server_id)) group = r.scalar_one_or_none() if not group: return None q = await session.execute( select(Tariff.device_limit) .where(Tariff.group_code == group, Tariff.is_active.is_(True)) .order_by(Tariff.duration_days.desc()) .limit(1) ) dl = q.scalar_one_or_none() return int(dl) if dl is not None else None async def filter_cluster_by_subgroup( session: AsyncSession, cluster: list, target_subgroup: str, cluster_id: str, tariff_id: int | None = None, ) -> list: names = [s.get("server_name") for s in cluster if s.get("server_name")] if not names: return [] if tariff_id: tariff_id_str = str(tariff_id) q_by_tariff = await session.execute( select(Server.server_name) .join(ServerSubgroup, ServerSubgroup.server_id == Server.id) .where( Server.server_name.in_(names), Server.enabled.is_(True), ServerSubgroup.subgroup_title == tariff_id_str, ) ) allowed_by_tariff = {n for (n,) in q_by_tariff.all()} if allowed_by_tariff: logger.debug(f"Найдены серверы по tariff_id={tariff_id}: {allowed_by_tariff}") return [s for s in cluster if s.get("server_name") in allowed_by_tariff] q_allowed = await session.execute( select(Server.server_name) .join(ServerSubgroup, ServerSubgroup.server_id == Server.id) .where( Server.server_name.in_(names), Server.enabled.is_(True), ServerSubgroup.subgroup_title == target_subgroup, ) ) allowed = {n for (n,) in q_allowed.all()} if allowed: return [s for s in cluster if s.get("server_name") in allowed] check_values = [target_subgroup] if tariff_id: check_values.append(str(tariff_id)) total_bindings = await session.scalar( select(func.count()).select_from(ServerSubgroup).where(ServerSubgroup.subgroup_title.in_(check_values)) ) if not total_bindings: logger.info(f"Для подгруппы/тарифа нет привязок. Используем весь кластер {cluster_id}.") return cluster q_any = await session.execute( select(Server.server_name) .join(ServerSubgroup, ServerSubgroup.server_id == Server.id) .where( Server.server_name.in_(names), Server.enabled.is_(True), ) ) any_bound = {n for (n,) in q_any.all()} if any_bound: logger.warning(f"Нет серверов под подгруппу {target_subgroup} в кластере {cluster_id}.") return [] logger.info(f"В кластере {cluster_id} нет привязок. Используем весь кластер.") return cluster async def filter_cluster_by_tariff(session: AsyncSession, cluster: list, tariff_id: int, cluster_id: str) -> list: names = [s.get("server_name") for s in cluster if s.get("server_name")] if not names: return [] tariff_id_str = str(tariff_id) q_allowed = await session.execute( select(Server.server_name) .join(ServerSubgroup, ServerSubgroup.server_id == Server.id) .where( Server.server_name.in_(names), Server.enabled.is_(True), ServerSubgroup.subgroup_title == tariff_id_str, ) ) allowed = {n for (n,) in q_allowed.all()} if allowed: return [s for s in cluster if s.get("server_name") in allowed] total_for_tariff = await session.scalar( select(func.count()).select_from(ServerSubgroup).where(ServerSubgroup.subgroup_title == tariff_id_str) ) if not total_for_tariff: logger.info(f"Для тарифа {tariff_id} нет привязок серверов. Используем весь кластер {cluster_id}.") return cluster q_any = await session.execute( select(Server.server_name) .join(ServerSubgroup, ServerSubgroup.server_id == Server.id) .where( Server.server_name.in_(names), Server.enabled.is_(True), ) ) any_bound = {n for (n,) in q_any.all()} if any_bound: logger.warning(f"Нет серверов под тариф {tariff_id} в кластере {cluster_id}.") return [] logger.info(f"В кластере {cluster_id} нет привязок тарифов. Используем весь кластер.") return cluster async def has_legacy_subgroup_bindings(session: AsyncSession, server_ids: list[int]) -> bool: if not server_ids: return False result = await session.execute( select(ServerSubgroup.subgroup_title).where(ServerSubgroup.server_id.in_(server_ids)) ) for (title,) in result.all(): if title and not title.isdigit(): return True return False