Speed up API reads with data cache and faster server check/stats.

Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
orohi
2026-07-26 06:07:44 +03:00
co-authored by Cursor
parent 2b28892a9b
commit fa3569f81f
3 changed files with 280 additions and 190 deletions
+148 -93
View File
@@ -2,6 +2,7 @@ import os
import sys import sys
import json import json
import logging import logging
from collections import Counter
from dotenv import load_dotenv from dotenv import load_dotenv
load_dotenv() load_dotenv()
@@ -167,6 +168,11 @@ load_translations()
DATA_LOCK = asyncio.Lock() DATA_LOCK = asyncio.Lock()
async def load_data_async():
"""Load panel state off the event loop (uses in-process cache)."""
return await asyncio.to_thread(load_data)
async def save_data_async(data): async def save_data_async(data):
"""Saves panel state to PostgreSQL in a thread-safe way.""" """Saves panel state to PostgreSQL in a thread-safe way."""
async with DATA_LOCK: async with DATA_LOCK:
@@ -183,6 +189,37 @@ def get_ssh(server):
) )
def get_current_user(request: Request, data: Optional[dict] = None):
user_id = request.session.get('user_id')
if not user_id:
return None
snapshot = data if data is not None else load_data()
for u in snapshot.get('users', []):
if u['id'] == user_id:
return u
return None
def tpl(request, template, **kwargs):
data = load_data()
lang = request.cookies.get('lang', 'en')
ctx = {
'request': request,
'current_user': get_current_user(request, data),
'site_settings': data.get('settings', {}).get('appearance', {}),
'captcha_settings': data.get('settings', {}).get('captcha', {}),
'telegram_settings': data.get('settings', {}).get('telegram', {}),
'bot_running': tg_bot.is_running(),
'lang': lang,
'_': lambda text_id: _t(text_id, lang),
'translations_json': json.dumps(TRANSLATIONS.get(lang, TRANSLATIONS.get('en', {}))),
# Keep for legacy JS; prefer translations_json on new pages
'all_translations_json': json.dumps(TRANSLATIONS),
}
ctx.update(kwargs)
return templates.TemplateResponse(template, ctx)
def get_panel_local_url(request: Optional[Request] = None): def get_panel_local_url(request: Optional[Request] = None):
data = load_data() data = load_data()
ssl_conf = data.get('settings', {}).get('ssl', {}) ssl_conf = data.get('settings', {}).get('ssl', {})
@@ -1337,36 +1374,6 @@ async def sync_users_with_remnawave(data: dict):
return 0, f"Error: {str(e)}" return 0, f"Error: {str(e)}"
def get_current_user(request: Request):
user_id = request.session.get('user_id')
if not user_id:
return None
data = load_data()
for u in data.get('users', []):
if u['id'] == user_id:
return u
return None
def tpl(request, template, **kwargs):
data = load_data()
lang = request.cookies.get('lang', 'en')
ctx = {
'request': request,
'current_user': get_current_user(request),
'site_settings': data.get('settings', {}).get('appearance', {}),
'captcha_settings': data.get('settings', {}).get('captcha', {}),
'telegram_settings': data.get('settings', {}).get('telegram', {}),
'bot_running': tg_bot.is_running(),
'lang': lang,
'_': lambda text_id: _t(text_id, lang),
'translations_json': json.dumps(TRANSLATIONS.get(lang, TRANSLATIONS.get('en', {}))),
'all_translations_json': json.dumps(TRANSLATIONS)
}
ctx.update(kwargs)
return templates.TemplateResponse(template, ctx)
# ======================== Pydantic Models ======================== # ======================== Pydantic Models ========================
class LoginRequest(BaseModel): class LoginRequest(BaseModel):
@@ -1946,22 +1953,24 @@ async def server_detail(request: Request, server_id: int):
@app.get('/users', response_class=HTMLResponse, tags=["System Templates"]) @app.get('/users', response_class=HTMLResponse, tags=["System Templates"])
async def users_page(request: Request): async def users_page(request: Request):
user = get_current_user(request) data = load_data()
user = get_current_user(request, data)
if not user: if not user:
return RedirectResponse(url='/login', status_code=302) return RedirectResponse(url='/login', status_code=302)
if user['role'] not in ('admin', 'support'): if user['role'] not in ('admin', 'support'):
return RedirectResponse(url='/my', status_code=302) return RedirectResponse(url='/my', status_code=302)
data = load_data()
users_list = data.get('users', []) users_list = data.get('users', [])
# Count connections per user # Count connections per user
conns = data.get('user_connections', []) conns = data.get('user_connections', [])
conn_counts = Counter(c.get('user_id') for c in conns if c.get('user_id'))
for u in users_list: for u in users_list:
u['connections_count'] = sum(1 for c in conns if c['user_id'] == u['id']) u['connections_count'] = conn_counts.get(u['id'], 0)
servers = data['servers'] servers = data['servers']
return tpl( return tpl(
request, 'users.html', request, 'users.html',
users=users_list, users=users_list,
servers=servers, servers=servers,
current_user=user,
) )
@@ -2309,54 +2318,97 @@ rm -rf /opt/amnezia
return JSONResponse({'error': str(e)}, status_code=500) return JSONResponse({'error': str(e)}, status_code=500)
def _docker_container_inventory(ssh) -> dict:
"""One SSH round-trip: map container name -> running bool."""
out, _, _ = ssh.run_sudo_command(
"docker ps -a --format '{{.Names}}|{{.Status}}' 2>/dev/null"
)
inventory = {}
for line in (out or '').splitlines():
if '|' not in line:
continue
name, status = line.split('|', 1)
name = name.strip()
if not name:
continue
inventory[name] = status.strip().lower().startswith('up')
return inventory
def _docker_is_ready(ssh) -> bool:
ver, _, vcode = ssh.run_command("docker --version 2>/dev/null")
return vcode == 0 and bool((ver or '').strip())
@app.post('/api/servers/{server_id}/stats', tags=["Servers"]) @app.post('/api/servers/{server_id}/stats', tags=["Servers"])
async def api_server_stats(request: Request, server_id: int): async def api_server_stats(request: Request, server_id: int):
if not _check_admin(request): if not _check_admin(request):
return JSONResponse({'error': 'Forbidden'}, status_code=403) return JSONResponse({'error': 'Forbidden'}, status_code=403)
try: try:
data = load_data() data = await load_data_async()
if server_id >= len(data['servers']): if server_id >= len(data['servers']):
return JSONResponse({'error': 'Server not found'}, status_code=404) return JSONResponse({'error': 'Server not found'}, status_code=404)
server = data['servers'][server_id] server = data['servers'][server_id]
def _fetch_stats():
ssh = get_ssh(server) ssh = get_ssh(server)
ssh.connect() ssh.connect()
stats = {}
out, _, _ = ssh.run_command(
"top -bn1 | grep 'Cpu(s)' | awk '{print $2}' | cut -d'%' -f1 2>/dev/null || "
"awk '{u=$2+$4; t=$2+$4+$5; if(NR==1){pu=u;pt=t} else printf \"%.1f\", (u-pu)/(t-pt)*100}' "
"<(grep 'cpu ' /proc/stat) <(sleep 0.5 && grep 'cpu ' /proc/stat) 2>/dev/null"
)
try: try:
stats['cpu'] = round(float(out.strip().split('\n')[0]), 1) script = r"""
except (ValueError, IndexError): cpu=$(top -bn1 2>/dev/null | grep 'Cpu(s)' | awk '{print $2}' | cut -d'%' -f1)
stats['cpu'] = 0 if [ -z "$cpu" ]; then cpu=0; fi
out, _, _ = ssh.run_command("free -b | awk 'NR==2{printf \"%d %d\", $3, $2}'") ram=$(free -b 2>/dev/null | awk 'NR==2{printf "%d %d", $3, $2}')
try: disk=$(df -B1 / 2>/dev/null | awk 'NR==2{printf "%d %d", $3, $2}')
parts = out.strip().split() dev=$(ip route 2>/dev/null | awk '/default/ {print $5; exit}')
used, total = int(parts[0]), int(parts[1]) net=$(awk -v d="${dev}:" '$1==d{printf "%d %d", $2, $10}' /proc/net/dev 2>/dev/null)
stats.update(ram_used=used, ram_total=total, ram_percent=round(used / total * 100, 1) if total > 0 else 0) up=$(uptime -p 2>/dev/null || uptime)
except (ValueError, IndexError): printf 'CPU|%s\nRAM|%s\nDISK|%s\nNET|%s\nUP|%s\n' "${cpu:-0}" "${ram:-0 0}" "${disk:-0 0}" "${net:-0 0}" "$up"
stats.update(ram_used=0, ram_total=0, ram_percent=0) """
out, _, _ = ssh.run_command("df -B1 / | awk 'NR==2{printf \"%d %d\", $3, $2}'") out, _, _ = ssh.run_command(script)
try: finally:
parts = out.strip().split()
used, total = int(parts[0]), int(parts[1])
stats.update(disk_used=used, disk_total=total, disk_percent=round(used / total * 100, 1) if total > 0 else 0)
except (ValueError, IndexError):
stats.update(disk_used=0, disk_total=0, disk_percent=0)
out, _, _ = ssh.run_command(
"DEV=$(ip route | awk '/default/ {print $5}' | head -1); "
"cat /proc/net/dev | awk -v dev=\"$DEV:\" '$1==dev{printf \"%d %d\", $2, $10}'"
)
try:
parts = out.strip().split()
stats['net_rx'], stats['net_tx'] = int(parts[0]), int(parts[1])
except (ValueError, IndexError):
stats['net_rx'] = stats['net_tx'] = 0
out, _, _ = ssh.run_command("uptime -p 2>/dev/null || uptime")
stats['uptime'] = out.strip()
ssh.disconnect() ssh.disconnect()
stats = {
'cpu': 0,
'ram_used': 0, 'ram_total': 0, 'ram_percent': 0,
'disk_used': 0, 'disk_total': 0, 'disk_percent': 0,
'net_rx': 0, 'net_tx': 0,
'uptime': '',
}
for line in (out or '').splitlines():
if '|' not in line:
continue
key, val = line.split('|', 1)
key, val = key.strip(), val.strip()
try:
if key == 'CPU':
stats['cpu'] = round(float(val.split()[0] if val else 0), 1)
elif key == 'RAM':
parts = val.split()
used, total = int(parts[0]), int(parts[1])
stats.update(
ram_used=used,
ram_total=total,
ram_percent=round(used / total * 100, 1) if total > 0 else 0,
)
elif key == 'DISK':
parts = val.split()
used, total = int(parts[0]), int(parts[1])
stats.update(
disk_used=used,
disk_total=total,
disk_percent=round(used / total * 100, 1) if total > 0 else 0,
)
elif key == 'NET':
parts = val.split()
stats['net_rx'], stats['net_tx'] = int(parts[0]), int(parts[1])
elif key == 'UP':
stats['uptime'] = val
except (ValueError, IndexError):
continue
return stats return stats
return await asyncio.to_thread(_fetch_stats)
except Exception as e: except Exception as e:
logger.exception("Error getting server stats") logger.exception("Error getting server stats")
return JSONResponse({'error': str(e)}, status_code=500) return JSONResponse({'error': str(e)}, status_code=500)
@@ -2367,27 +2419,27 @@ async def api_check_server(request: Request, server_id: int):
if not _check_admin(request): if not _check_admin(request):
return JSONResponse({'error': 'Forbidden'}, status_code=403) return JSONResponse({'error': 'Forbidden'}, status_code=403)
try: try:
data = load_data() data = await load_data_async()
if server_id >= len(data['servers']): if server_id >= len(data['servers']):
return JSONResponse({'error': 'Server not found'}, status_code=404) return JSONResponse({'error': 'Server not found'}, status_code=404)
server = data['servers'][server_id] server = data['servers'][server_id]
def _run_check():
ssh = get_ssh(server) ssh = get_ssh(server)
ssh.connect() ssh.connect()
# Just use awg's docker checker since it uses the same command try:
manager = get_protocol_manager(ssh, 'awg') status = {
status = {'connection': 'ok', 'docker_installed': manager.check_docker_installed(), 'protocols': {}} 'connection': 'ok',
'docker_installed': _docker_is_ready(ssh),
'protocols': {},
}
inventory = _docker_container_inventory(ssh) if status['docker_installed'] else {}
changed = False changed = False
if 'protocols' not in server: if 'protocols' not in server:
server['protocols'] = {} server['protocols'] = {}
def merge_saved_protocol_status(proto, result=None, error=None): def merge_saved_protocol_status(proto, result=None, error=None):
"""Merge live status with saved protocol metadata.
Multi-instance protocols are source-of-truth in data.json because
they cannot be discovered from BASE_PROTOCOLS alone. A transient
check failure must not delete awg2__2/awg2__3 from the panel.
"""
db_proto = server.get('protocols', {}).get(proto, {}) or {} db_proto = server.get('protocols', {}).get(proto, {}) or {}
merged = dict(result or {}) merged = dict(result or {})
merged.setdefault('protocol', proto) merged.setdefault('protocol', proto)
@@ -2416,22 +2468,24 @@ async def api_check_server(request: Request, server_id: int):
return merged return merged
def should_preserve_saved_protocol(proto, result=None, err=None): def should_preserve_saved_protocol(proto, result=None, err=None):
"""Return True when check must not remove a saved protocol record."""
db_proto = server.get('protocols', {}).get(proto) db_proto = server.get('protocols', {}).get(proto)
if not db_proto: if not db_proto:
return False return False
# Additional AWG-family instances are only known by their saved
# dynamic keys (awg__2/awg2__2/awg_legacy__2). Keep them unless
# the user explicitly uninstalls them.
if protocol_base(proto) in MULTI_INSTANCE_PROTOCOLS and protocol_instance(proto) > 1: if protocol_base(proto) in MULTI_INSTANCE_PROTOCOLS and protocol_instance(proto) > 1:
return True return True
# Do not delete any saved protocol on command/check errors; only a
# clean live result with container_exists=False may prune base apps.
if err or (result and result.get('error')): if err or (result and result.get('error')):
return True return True
return False return False
def check_proto(proto): def check_proto(proto):
cname = protocol_container_name(proto)
# Fast path: skip deep manager probes when container is absent
if cname and cname not in inventory:
return proto, merge_saved_protocol_status(proto, {
'container_exists': False,
'container_running': False,
'protocol': proto,
}), None
try: try:
p_manager = get_protocol_manager(ssh, proto) p_manager = get_protocol_manager(ssh, proto)
result = _manager_call(p_manager, 'get_server_status', proto) result = _manager_call(p_manager, 'get_server_status', proto)
@@ -2439,10 +2493,9 @@ async def api_check_server(request: Request, server_id: int):
except Exception as e: except Exception as e:
return proto, merge_saved_protocol_status(proto, {}, str(e)), str(e) return proto, merge_saved_protocol_status(proto, {}, str(e)), str(e)
protocols_to_check = list(dict.fromkeys(BASE_PROTOCOLS + list(server.get('protocols', {}).keys()))) protocols_to_check = list(dict.fromkeys(
# Run checks sequentially. Several managers use the same SSH connection; BASE_PROTOCOLS + list(server.get('protocols', {}).keys())
# checking them in parallel through one SSH object can produce false ))
# negatives and previously caused dynamic AWG instances to be removed.
for proto in protocols_to_check: for proto in protocols_to_check:
proto, result, err = check_proto(proto) proto, result, err = check_proto(proto)
status['protocols'][proto] = result status['protocols'][proto] = result
@@ -2476,7 +2529,6 @@ async def api_check_server(request: Request, server_id: int):
else: else:
if proto in server['protocols']: if proto in server['protocols']:
if should_preserve_saved_protocol(proto, result, err): if should_preserve_saved_protocol(proto, result, err):
# Keep saved dynamic instances visible as installed but stopped/unchecked.
status['protocols'][proto]['container_exists'] = True status['protocols'][proto]['container_exists'] = True
status['protocols'][proto].setdefault('container_running', False) status['protocols'][proto].setdefault('container_running', False)
status['protocols'][proto]['status_preserved'] = True status['protocols'][proto]['status_preserved'] = True
@@ -2486,9 +2538,11 @@ async def api_check_server(request: Request, server_id: int):
if changed: if changed:
save_data(data) save_data(data)
ssh.disconnect()
return status return status
finally:
ssh.disconnect()
return await asyncio.to_thread(_run_check)
except Exception as e: except Exception as e:
logger.exception("Error checking server") logger.exception("Error checking server")
return JSONResponse({'error': str(e), 'connection': 'failed'}, status_code=500) return JSONResponse({'error': str(e), 'connection': 'failed'}, status_code=500)
@@ -3470,9 +3524,10 @@ async def api_toggle_connection(request: Request, server_id: int, req: ToggleCon
async def api_list_users(request: Request, search: str = '', page: int = 1, size: int = 10): async def api_list_users(request: Request, search: str = '', page: int = 1, size: int = 10):
if not _check_admin(request): if not _check_admin(request):
return JSONResponse({'error': 'Forbidden'}, status_code=403) return JSONResponse({'error': 'Forbidden'}, status_code=403)
data = load_data() data = await load_data_async()
all_users = data.get('users', []) all_users = data.get('users', [])
conns = data.get('user_connections', []) conns = data.get('user_connections', [])
conn_counts = Counter(c.get('user_id') for c in conns if c.get('user_id'))
# Filter # Filter
filtered = [] filtered = []
@@ -3500,7 +3555,7 @@ async def api_list_users(request: Request, search: str = '', page: int = 1, size
'telegramId': u.get('telegramId'), 'telegramId': u.get('telegramId'),
'email': u.get('email'), 'email': u.get('email'),
'description': u.get('description'), 'description': u.get('description'),
'connections_count': sum(1 for c in conns if c['user_id'] == u['id']), 'connections_count': conn_counts.get(u['id'], 0),
'traffic_used': u.get('traffic_used', 0), 'traffic_used': u.get('traffic_used', 0),
'traffic_total': u.get('traffic_total', 0), 'traffic_total': u.get('traffic_total', 0),
'traffic_limit': u.get('traffic_limit', 0), 'traffic_limit': u.get('traffic_limit', 0),
+36 -1
View File
@@ -6,10 +6,12 @@ FastAPI handlers keep working without a full rewrite.
from __future__ import annotations from __future__ import annotations
import copy
import json import json
import logging import logging
import os import os
import shutil import shutil
import threading
from datetime import datetime, timezone from datetime import datetime, timezone
from pathlib import Path from pathlib import Path
from typing import Any, Optional from typing import Any, Optional
@@ -21,6 +23,11 @@ from .connection import get_pool, init_schema
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
# Process-wide snapshot cache. load_data() returns a deep copy so callers can
# mutate safely; save_data() refreshes the cache after a successful write.
_DATA_CACHE: Optional[dict] = None
_DATA_CACHE_LOCK = threading.RLock()
DEFAULT_SETTINGS = { DEFAULT_SETTINGS = {
'appearance': { 'appearance': {
'title': 'Amnezia', 'title': 'Amnezia',
@@ -208,7 +215,13 @@ def _row_to_invite(row) -> dict:
} }
def load_data() -> dict: def invalidate_data_cache() -> None:
global _DATA_CACHE
with _DATA_CACHE_LOCK:
_DATA_CACHE = None
def _fetch_data_from_db() -> dict:
init_schema() init_schema()
pool = get_pool() pool = get_pool()
with pool.connection() as conn: with pool.connection() as conn:
@@ -266,8 +279,20 @@ def load_data() -> dict:
} }
def load_data() -> dict:
"""Return panel state. Uses an in-process cache; always returns a deep copy."""
global _DATA_CACHE
with _DATA_CACHE_LOCK:
if _DATA_CACHE is not None:
return copy.deepcopy(_DATA_CACHE)
data = _fetch_data_from_db()
_DATA_CACHE = data
return copy.deepcopy(data)
def save_data(data: dict) -> None: def save_data(data: dict) -> None:
"""Replace panel state in a single transaction (same semantics as rewriting data.json).""" """Replace panel state in a single transaction (same semantics as rewriting data.json)."""
global _DATA_CACHE
init_schema() init_schema()
servers = data.get('servers') or [] servers = data.get('servers') or []
users = data.get('users') or [] users = data.get('users') or []
@@ -419,6 +444,16 @@ def save_data(data: dict) -> None:
) )
conn.commit() conn.commit()
# Keep caller dict and cache aligned with what was actually persisted
data['servers'] = servers
data['users'] = users
data['user_connections'] = connections
data['api_tokens'] = tokens
data['invite_links'] = invite_links
data['settings'] = settings
with _DATA_CACHE_LOCK:
_DATA_CACHE = copy.deepcopy(data)
def export_data_dict() -> dict: def export_data_dict() -> dict:
"""Export current DB state as the legacy data.json document.""" """Export current DB state as the legacy data.json document."""
+1 -1
View File
@@ -66,7 +66,7 @@ class SSHManager:
if not self.client: if not self.client:
raise ConnectionError("Not connected to server") raise ConnectionError("Not connected to server")
logger.info(f"Running command: {command[:100]}...") logger.debug(f"Running command: {command[:100]}...")
stdin, stdout, stderr = self.client.exec_command(command, timeout=timeout) stdin, stdout, stderr = self.client.exec_command(command, timeout=timeout)
# Crucial: set timeout on the channel to prevent hanging indefinitely # Crucial: set timeout on the channel to prevent hanging indefinitely