diff --git a/api/main.py b/api/main.py index 82684327..5a0f5633 100644 --- a/api/main.py +++ b/api/main.py @@ -1,14 +1,50 @@ -from fastapi import FastAPI -from api.routes import users, keys, coupons, servers, tariffs, gifts, referrals, misc, partners +from time import perf_counter + +from fastapi import FastAPI, Request +from api.routes import ( + users, + keys, + coupons, + servers, + tariffs, + gifts, + referrals, + misc, + partners, + modules, + management, + settings, +) +from config import API_LOGGING +from logger import logger app = FastAPI( title="SoloBot API (Alpha)", - version="0.5.1", + version="0.5.2", docs_url="/api/docs", redoc_url="/api/redoc", openapi_url="/api/openapi.json", ) + +@app.middleware("http") +async def api_access_log_middleware(request: Request, call_next): + if not API_LOGGING: + return await call_next(request) + + started = perf_counter() + response = await call_next(request) + duration_ms = int((perf_counter() - started) * 1000) + client_ip = request.client.host if request.client else "-" + path_qs = request.url.path + if request.url.query: + path_qs = f"{path_qs}?{request.url.query}" + + logger.info( + f'[API] {client_ip} "{request.method} {path_qs}" {response.status_code} {duration_ms}ms' + ) + return response + app.include_router(users.router, prefix="/api/users", tags=["Users"]) app.include_router(keys.router, prefix="/api/keys", tags=["Keys"]) app.include_router(coupons.router, prefix="/api/coupons", tags=["Coupons"]) @@ -18,6 +54,9 @@ app.include_router(gifts.router, prefix="/api/gifts", tags=["Gifts"]) app.include_router(referrals.router, prefix="/api/referrals", tags=["Referrals"]) app.include_router(partners.router, prefix="/api/partners", tags=["Partners"]) app.include_router(misc.router, prefix="/api") +app.include_router(modules.router, prefix="/api") +app.include_router(management.router, prefix="/api/management", tags=["Management"]) +app.include_router(settings.router, prefix="/api/settings", tags=["Settings"]) @app.get("/api", include_in_schema=False) diff --git a/api/routes/management.py b/api/routes/management.py new file mode 100644 index 00000000..52bc5a94 --- /dev/null +++ b/api/routes/management.py @@ -0,0 +1,219 @@ +import os +import re +import subprocess +import sys +import asyncio +from typing import Literal + +import psutil +from aiogram import Bot +from aiogram.client.default import DefaultBotProperties +from aiogram.enums import ParseMode +from fastapi import APIRouter, BackgroundTasks, Depends, HTTPException +from pydantic import BaseModel +from sqlalchemy import distinct, exists, func, select, update +from sqlalchemy.ext.asyncio import AsyncSession + +from api.depends import get_session, verify_admin_token +from config import API_TOKEN +from core.bootstrap import MANAGEMENT_CONFIG +from core.settings.management_config import update_management_config +from database.models import Key, User +from database.models import Server +from handlers.admin.sender.sender_service import BroadcastService +from handlers.admin.sender.sender_utils import get_recipients, parse_message_buttons +from logger import logger +from utils.backup import backup_database + + +router = APIRouter() + + +class MaintenanceUpdate(BaseModel): + enabled: bool + + +class DomainChange(BaseModel): + domain: str + + +class BroadcastLaunchPayload(BaseModel): + send_to: Literal["all", "subscribed", "unsubscribed", "untrial", "trial", "hotleads", "cluster"] = "all" + text: str + photo: str | None = None + cluster_name: str | None = None + workers: int = 5 + messages_per_second: int = 35 + + +_broadcast_bot: Bot | None = None + + +def _get_broadcast_bot() -> Bot: + global _broadcast_bot + if _broadcast_bot is None: + _broadcast_bot = Bot(token=API_TOKEN, default=DefaultBotProperties(parse_mode=ParseMode.HTML)) + return _broadcast_bot + + +async def _restart_bot() -> None: + await asyncio.sleep(1) + try: + parent = psutil.Process(os.getpid()).parent() + is_systemd = parent and "systemd" in parent.name().lower() + + if is_systemd: + subprocess.run( + ["sudo", "systemctl", "restart", "bot.service"], + check=True, + ) + else: + python_exe = sys.executable + script_path = os.path.abspath(sys.argv[0]) + os.execv(python_exe, [python_exe, script_path] + sys.argv[1:]) + except Exception: + os._exit(1) + + +@router.get("/status") +async def get_status(admin=Depends(verify_admin_token)): + return { + "maintenance_enabled": bool(MANAGEMENT_CONFIG.get("MAINTENANCE_ENABLED", False)), + "management": dict(MANAGEMENT_CONFIG or {}), + } + + +@router.post("/maintenance") +async def set_maintenance( + payload: MaintenanceUpdate, + admin=Depends(verify_admin_token), + session: AsyncSession = Depends(get_session), +): + current_config = dict(MANAGEMENT_CONFIG or {}) + current_config["MAINTENANCE_ENABLED"] = bool(payload.enabled) + await update_management_config(session, current_config) + return {"maintenance_enabled": bool(MANAGEMENT_CONFIG.get("MAINTENANCE_ENABLED", False))} + + +@router.post("/restart") +async def restart_bot( + background: BackgroundTasks, + admin=Depends(verify_admin_token), +): + background.add_task(_restart_bot) + return {"status": "restarting"} + + +@router.post("/change-domain") +async def change_domain( + payload: DomainChange, + admin=Depends(verify_admin_token), + session: AsyncSession = Depends(get_session), +): + domain = payload.domain.strip() + if not domain or " " in domain or not re.fullmatch(r"[a-zA-Z0-9.-]+", domain): + raise HTTPException(status_code=400, detail="Invalid domain") + + new_domain_url = f"https://{domain}" + + stmt = ( + update(Key) + .values( + key=func.regexp_replace(Key.key, r"^https://[^/]+", new_domain_url), + remnawave_link=func.regexp_replace(Key.remnawave_link, r"^https://[^/]+", new_domain_url), + ) + .where( + (Key.key.startswith("https://") & ~Key.key.startswith(new_domain_url)) + | (Key.remnawave_link.startswith("https://") & ~Key.remnawave_link.startswith(new_domain_url)) + ) + ) + result = await session.execute(stmt) + await session.commit() + + return {"updated": result.rowcount or 0} + + +@router.post("/restore-trials") +async def restore_trials( + admin=Depends(verify_admin_token), + session: AsyncSession = Depends(get_session), +): + stmt = ( + update(User) + .where( + User.trial == 1, + ~exists(select(Key.tg_id).where(Key.tg_id == User.tg_id)), + ) + .values(trial=0) + ) + result = await session.execute(stmt) + await session.commit() + + return {"restored": result.rowcount or 0} + + +@router.post("/backup") +async def trigger_backup(admin=Depends(verify_admin_token)): + async def _run_backup() -> None: + exception = await backup_database() + if exception: + logger.error(f"[Management] Backup finished with error: {exception}") + + asyncio.create_task(_run_backup()) + return {"status": "backup_started"} + + +@router.get("/broadcast/clusters") +async def get_broadcast_clusters( + admin=Depends(verify_admin_token), + session: AsyncSession = Depends(get_session), +): + result = await session.execute(select(distinct(Server.cluster_name)).where(Server.cluster_name.is_not(None))) + clusters = sorted([row[0] for row in result.all() if row and row[0]]) + return {"clusters": clusters} + + +@router.post("/broadcast") +async def launch_broadcast( + payload: BroadcastLaunchPayload, + admin=Depends(verify_admin_token), + session: AsyncSession = Depends(get_session), +): + text_raw = (payload.text or "").strip() + if not text_raw: + raise HTTPException(status_code=400, detail="Broadcast text is required") + + if payload.send_to == "cluster" and not (payload.cluster_name or "").strip(): + raise HTTPException(status_code=400, detail="Cluster name is required for cluster broadcast") + + clean_text, keyboard = parse_message_buttons(text_raw) + + max_len = 1024 if payload.photo else 4096 + if len(clean_text) > max_len: + raise HTTPException(status_code=400, detail=f"Message too long. Max {max_len} symbols") + + tg_ids, total_users = await get_recipients(session, payload.send_to, (payload.cluster_name or None)) + if not tg_ids: + return {"success": False, "message": "No recipients found", "stats": {"total_messages": 0}} + + bot = _get_broadcast_bot() + messages = [ + { + "tg_id": tg_id, + "text": clean_text, + "photo": payload.photo, + "keyboard": keyboard, + } + for tg_id in tg_ids + ] + + workers = max(1, min(int(payload.workers or 5), 30)) + rate = max(1, min(int(payload.messages_per_second or 35), 60)) + broadcast_service = BroadcastService(bot=bot, session=session, messages_per_second=rate) + stats = await broadcast_service.broadcast(messages, workers=workers) + return { + "success": True, + "message": "Broadcast completed", + "recipients": total_users, + "stats": stats, + } diff --git a/api/routes/misc.py b/api/routes/misc.py index 456e9306..a3266840 100644 --- a/api/routes/misc.py +++ b/api/routes/misc.py @@ -9,7 +9,6 @@ from api.schemas import ( ManualBanResponse, NotificationResponse, PaymentResponse, - ReferralResponse, TemporaryDataResponse, TrackingSourceResponse, ) @@ -20,7 +19,6 @@ from database.models import ( ManualBan, Notification, Payment, - Referral, TemporaryData, TrackingSource, ) @@ -56,20 +54,6 @@ async def get_payments_by_tg_id( return payments -router.include_router( - generate_crud_router( - model=Referral, - schema_response=ReferralResponse, - schema_create=None, - schema_update=None, - identifier_field="referred_tg_id", - enabled_methods=["get_all", "get_one", "delete"], - ), - prefix="/referrals", - tags=["Referrals"], - dependencies=[Depends(verify_admin_token)], -) - router.include_router( generate_crud_router( model=Notification, diff --git a/api/routes/modules.py b/api/routes/modules.py new file mode 100644 index 00000000..eba10248 --- /dev/null +++ b/api/routes/modules.py @@ -0,0 +1,109 @@ +import pkgutil +from pathlib import Path +from typing import Literal + +from fastapi import APIRouter, Depends, HTTPException +from pydantic import BaseModel + +from api.depends import verify_admin_token +from utils.modules_loader import _is_safe_module_name +from utils.modules_manager import manager + + +router = APIRouter(prefix="/modules", tags=["Modules"]) + + +MODULES_DIR = Path(__file__).resolve().parents[2] / "modules" + + +class ModuleAction(BaseModel): + action: Literal["start", "stop", "restart"] + + +def _available_module_names() -> list[str]: + candidates: set[str] = set() + + for name in manager.registry.keys(): + if name: + candidates.add(name) + + candidates.update(manager.disabled) + + if MODULES_DIR.is_dir(): + for _finder, name, _ispkg in pkgutil.iter_modules([str(MODULES_DIR)]): + name = (name or "").strip() + if name and _is_safe_module_name(name): + candidates.add(name) + + return sorted(n for n in candidates if _is_safe_module_name(n)) + + +def _module_state(name: str) -> dict: + normalized = name.strip() + record = manager.registry.get(normalized) + is_enabled = manager.is_enabled(normalized) + return { + "name": normalized, + "enabled": is_enabled, + "loaded": bool(record and record.enabled), + "autostart": manager.should_autostart(normalized), + } + + +def _read_local_module_version(name: str) -> str | None: + version_file = MODULES_DIR / name / "VERSION" + if not version_file.exists() or not version_file.is_file(): + return None + + try: + with version_file.open("r", encoding="utf-8") as handle: + for line in handle: + value = line.strip() + if value: + return value + except Exception: + return None + return None + + +@router.get("/") +async def list_modules(admin=Depends(verify_admin_token)): + refresh = getattr(manager, "refresh_state", None) + if callable(refresh): + refresh() + else: + legacy_refresh = getattr(manager, "_load_state", None) + if callable(legacy_refresh): + legacy_refresh() + module_names = _available_module_names() + modules = [_module_state(name) for name in module_names] + + for item in modules: + name = str(item.get("name") or "").strip() + local_version = _read_local_module_version(name) + item["local_version"] = local_version + + return {"items": modules} + + +@router.post("/{module_name}/actions") +async def control_module(module_name: str, payload: ModuleAction, admin=Depends(verify_admin_token)): + name = (module_name or "").strip() + if not _is_safe_module_name(name): + raise HTTPException(status_code=404, detail="Module not found") + + try: + if payload.action == "start": + await manager.start(name) + elif payload.action == "stop": + await manager.stop(name) + elif payload.action == "restart": + await manager.restart(name) + else: + raise HTTPException(status_code=400, detail="Unsupported action") + except ValueError as exc: + raise HTTPException(status_code=400, detail=str(exc)) from exc + except RuntimeError as exc: + raise HTTPException(status_code=500, detail=str(exc)) from exc + + return {"item": _module_state(name)} diff --git a/api/routes/settings.py b/api/routes/settings.py index 94616134..9acafcfc 100644 --- a/api/routes/settings.py +++ b/api/routes/settings.py @@ -1,4 +1,6 @@ -from fastapi import APIRouter, Depends, HTTPException, Query +from typing import Any + +from fastapi import APIRouter, Depends, HTTPException from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession @@ -6,11 +8,23 @@ from api.depends import get_session, verify_admin_token from api.schemas.settings import SettingResponse, SettingUpsert from database.models import Setting from database.settings import set_setting +from core.settings.buttons_config import BUTTONS_CONFIG, update_buttons_config +from core.settings.modes_config import MODES_CONFIG, update_modes_config +from core.settings.money_config import MONEY_CONFIG, update_money_config +from core.settings.notifications_config import NOTIFICATIONS_CONFIG, update_notifications_config +from core.settings.payments_config import PAYMENTS_CONFIG, update_payments_config +from core.settings.providers_order_config import PROVIDERS_ORDER, update_providers_order +from core.settings.tariffs_config import TARIFFS_CONFIG, update_tariffs_config +from pydantic import BaseModel router = APIRouter() +class ConfigUpdatePayload(BaseModel): + value: dict[str, Any] | None = None + + @router.get("/", response_model=list[SettingResponse]) async def get_all_settings( admin=Depends(verify_admin_token), @@ -20,6 +34,77 @@ async def get_all_settings( return result.scalars().all() +@router.get("/configs") +async def get_configs(admin=Depends(verify_admin_token)): + return { + "payments": dict(PAYMENTS_CONFIG), + "buttons": dict(BUTTONS_CONFIG), + "notifications": dict(NOTIFICATIONS_CONFIG), + "modes": dict(MODES_CONFIG), + "money": dict(MONEY_CONFIG), + "providers_order": dict(PROVIDERS_ORDER), + "tariffs": dict(TARIFFS_CONFIG), + } + + +@router.post("/configs/{scope}") +async def update_config_scope( + scope: str, + payload: ConfigUpdatePayload, + admin=Depends(verify_admin_token), + session: AsyncSession = Depends(get_session), +): + data = dict(payload.value or {}) + normalized = scope.strip().lower().replace("-", "_") + + if normalized == "payments": + cleaned = {key: bool(value) for key, value in data.items()} + await update_payments_config(session, cleaned) + return {"payments": dict(PAYMENTS_CONFIG)} + + if normalized == "buttons": + cleaned = {key: bool(value) for key, value in data.items()} + await update_buttons_config(session, cleaned) + return {"buttons": dict(BUTTONS_CONFIG)} + + if normalized == "notifications": + await update_notifications_config(session, data) + return {"notifications": dict(NOTIFICATIONS_CONFIG)} + + if normalized == "modes": + cleaned = {key: bool(value) for key, value in data.items()} + await update_modes_config(session, cleaned) + return {"modes": dict(MODES_CONFIG)} + + if normalized == "money": + await update_money_config(session, data) + return {"money": dict(MONEY_CONFIG)} + + if normalized == "providers_order": + cleaned: dict[str, int] = {} + for key, value in data.items(): + try: + cleaned[key] = int(value) + except (TypeError, ValueError): + continue + await update_providers_order(session, cleaned) + return {"providers_order": dict(PROVIDERS_ORDER)} + + if normalized == "tariffs": + cleaned = dict(data) + if "ALLOW_DOWNGRADE" in cleaned: + cleaned["ALLOW_DOWNGRADE"] = bool(cleaned.get("ALLOW_DOWNGRADE")) + if "KEY_ADDONS_RECALC_PRICE" in cleaned: + cleaned["KEY_ADDONS_RECALC_PRICE"] = bool(cleaned.get("KEY_ADDONS_RECALC_PRICE")) + if "KEY_ADDONS_PACK_MODE" in cleaned: + mode = str(cleaned.get("KEY_ADDONS_PACK_MODE") or "").strip().lower() + cleaned["KEY_ADDONS_PACK_MODE"] = mode if mode in {"", "traffic", "devices", "all"} else "" + await update_tariffs_config(session, cleaned) + return {"tariffs": dict(TARIFFS_CONFIG)} + + raise HTTPException(status_code=404, detail="Unsupported config scope") + + @router.get("/{key}", response_model=SettingResponse) async def get_setting_by_key( key: str, diff --git a/api/schemas/management.py b/api/schemas/management.py new file mode 100644 index 00000000..6ddb3bda --- /dev/null +++ b/api/schemas/management.py @@ -0,0 +1,9 @@ +from pydantic import BaseModel + + +class MaintenanceUpdate(BaseModel): + enabled: bool + + +class DomainChange(BaseModel): + domain: str