Files
Amnezia-Web-Panel-main/db/store.py
T

559 lines
20 KiB
Python

"""Panel data store backed by PostgreSQL 17.
Preserves the same dict shape that used to live in data.json so existing
FastAPI handlers keep working without a full rewrite.
"""
from __future__ import annotations
import copy
import json
import logging
import os
import shutil
import threading
from datetime import datetime, timezone
from pathlib import Path
from typing import Any, Optional
from uuid import UUID
from psycopg.types.json import Jsonb
from .connection import get_pool, init_schema
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 = {
'appearance': {
'title': 'Amnezia',
'logo': '❤️',
'subtitle': 'Web Panel',
},
'sync': {
'remnawave_url': '',
'remnawave_api_key': '',
'remnawave_sync': False,
'remnawave_sync_users': False,
'remnawave_create_conns': False,
'remnawave_server_id': 0,
'remnawave_protocol': 'awg',
},
'guest': {
'enabled': False,
'token': '',
'password_hash': None,
'user_id': '',
'allow_create': False,
'create_protocol': 'awg',
'create_server_id': 0,
},
}
def _parse_ts(value: Any) -> Optional[datetime]:
if value is None or value == '':
return None
if isinstance(value, datetime):
return value
text = str(value).strip()
if not text:
return None
if text.endswith('Z'):
text = text[:-1] + '+00:00'
try:
return datetime.fromisoformat(text)
except ValueError:
return None
def _ts_iso(value: Any) -> Optional[str]:
if value is None:
return None
if isinstance(value, datetime):
return value.isoformat()
return str(value)
def _as_uuid(value: Any) -> UUID:
return UUID(str(value))
def _merge_settings(raw: Optional[dict]) -> dict:
settings = json.loads(json.dumps(DEFAULT_SETTINGS))
if not isinstance(raw, dict):
return settings
for section, defaults in DEFAULT_SETTINGS.items():
incoming = raw.get(section)
if isinstance(incoming, dict) and isinstance(defaults, dict):
merged = dict(defaults)
merged.update(incoming)
settings[section] = merged
elif section in raw:
settings[section] = raw[section]
for key, value in raw.items():
if key not in settings:
settings[key] = value
return settings
def _as_dict(value, default=None):
"""Coerce JSONB / legacy values to a plain dict."""
if default is None:
default = {}
if value is None:
return dict(default)
if isinstance(value, dict):
return dict(value)
# Legacy: ssh.test_connection() used to return a plain string
if isinstance(value, str):
text = value.strip()
if not text:
return dict(default)
if text.startswith('{') or text.startswith('['):
try:
parsed = json.loads(text)
if isinstance(parsed, dict):
return parsed
except Exception:
pass
return {'raw': text}
try:
return dict(value)
except Exception:
return {'raw': str(value)}
def _row_to_server(row) -> dict:
return {
'name': row['name'] or '',
'host': row['host'] or '',
'ssh_port': int(row['ssh_port'] or 22),
'username': row['username'] or '',
'password': row['password'],
'private_key': row['private_key'],
'server_info': _as_dict(row['server_info']),
'protocols': _as_dict(row['protocols']),
}
def _row_to_user(row) -> dict:
return {
'id': str(row['id']),
'username': row['username'],
'password_hash': row['password_hash'] or '',
'role': row['role'] or 'user',
'enabled': bool(row['enabled']),
'created_at': _ts_iso(row['created_at']),
'telegramId': row['telegram_id'],
'email': row['email'],
'description': row['description'],
'traffic_limit': int(row['traffic_limit'] or 0),
'traffic_used': int(row['traffic_used'] or 0),
'traffic_total': int(row['traffic_total'] or 0),
'traffic_reset_strategy': row['traffic_reset_strategy'] or 'never',
'last_reset_at': _ts_iso(row['last_reset_at']),
'expiration_date': _ts_iso(row['expiration_date']),
'expire_after_first_use': bool(row.get('expire_after_first_use') if hasattr(row, 'get') else row['expire_after_first_use']),
'expiration_days': int((row.get('expiration_days') if hasattr(row, 'get') else row['expiration_days']) or 0),
'remnawave_uuid': row['remnawave_uuid'],
'xui_email': row.get('xui_email'),
'share_enabled': bool(row['share_enabled']),
'share_token': row['share_token'],
'share_password_hash': row['share_password_hash'],
}
def _row_to_connection(row) -> dict:
return {
'id': str(row['id']),
'user_id': str(row['user_id']),
'server_id': int(row['server_id'] or 0),
'protocol': row['protocol'] or '',
'client_id': row['client_id'] or '',
'name': row['name'] or '',
'xui_panel_id': (row.get('xui_panel_id') or '') if hasattr(row, 'get') else (row['xui_panel_id'] if 'xui_panel_id' in row else ''),
'created_at': _ts_iso(row['created_at']),
'last_bytes': int(row['last_bytes'] or 0),
}
def _row_to_token(row) -> dict:
return {
'id': str(row['id']),
'name': row['name'] or '',
'token_hash': row['token_hash'] or '',
'token_prefix': row['token_prefix'] or '',
'user_id': str(row['user_id']),
'created_at': _ts_iso(row['created_at']),
'last_used_at': _ts_iso(row['last_used_at']),
}
def _row_to_invite(row) -> dict:
return {
'id': str(row['id']),
'name': row['name'] or '',
'token': row['token'] or '',
'enabled': bool(row['enabled']),
'max_uses': int(row['max_uses'] or 0),
'used_count': int(row['used_count'] or 0),
'user_id': str(row['user_id']) if row['user_id'] else '',
'protocol': row['protocol'] or 'awg',
'server_id': int(row['server_id'] or 0),
'xui_inbound_id': int(row['xui_inbound_id'] or 0),
'xui_panel_id': (row.get('xui_panel_id') or '') if hasattr(row, 'get') else '',
'password_hash': row['password_hash'],
'expires_at': _ts_iso(row['expires_at']),
'duration_days': int(row.get('duration_days') or 0),
'note': row['note'] or '',
'created_at': _ts_iso(row['created_at']),
}
def invalidate_data_cache() -> None:
global _DATA_CACHE
with _DATA_CACHE_LOCK:
_DATA_CACHE = None
def _fetch_data_from_db() -> dict:
init_schema()
pool = get_pool()
with pool.connection() as conn:
with conn.cursor() as cur:
cur.execute(
'SELECT position, name, host, ssh_port, username, password, '
'private_key, server_info, protocols '
'FROM servers ORDER BY position ASC'
)
servers = [_row_to_server(r) for r in cur.fetchall()]
cur.execute(
'SELECT id, username, password_hash, role, enabled, created_at, '
'telegram_id, email, description, traffic_limit, traffic_used, '
'traffic_total, traffic_reset_strategy, last_reset_at, '
'expiration_date, expire_after_first_use, expiration_days, '
'remnawave_uuid, xui_email, share_enabled, share_token, '
'share_password_hash FROM users ORDER BY created_at NULLS LAST, username'
)
users = [_row_to_user(r) for r in cur.fetchall()]
cur.execute(
'SELECT id, user_id, server_id, protocol, client_id, name, '
'xui_panel_id, created_at, last_bytes FROM user_connections '
'ORDER BY created_at NULLS LAST, id'
)
user_connections = [_row_to_connection(r) for r in cur.fetchall()]
cur.execute(
'SELECT id, name, token_hash, token_prefix, user_id, '
'created_at, last_used_at FROM api_tokens '
'ORDER BY created_at NULLS LAST, id'
)
api_tokens = [_row_to_token(r) for r in cur.fetchall()]
cur.execute(
'SELECT id, name, token, enabled, max_uses, used_count, user_id, '
'protocol, server_id, xui_inbound_id, xui_panel_id, password_hash, expires_at, '
'duration_days, note, created_at FROM invite_links '
'ORDER BY created_at DESC NULLS LAST, name'
)
invite_links = [_row_to_invite(r) for r in cur.fetchall()]
cur.execute('SELECT data FROM settings WHERE id = 1')
settings_row = cur.fetchone()
settings = _merge_settings(settings_row['data'] if settings_row else None)
return {
'servers': servers,
'users': users,
'user_connections': user_connections,
'api_tokens': api_tokens,
'invite_links': invite_links,
'settings': settings,
}
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:
"""Replace panel state in a single transaction (same semantics as rewriting data.json)."""
global _DATA_CACHE
init_schema()
servers = data.get('servers') or []
users = data.get('users') or []
connections = data.get('user_connections') or []
tokens = data.get('api_tokens') or []
invite_links = data.get('invite_links') or []
settings = _merge_settings(data.get('settings'))
# Keep only connections whose user still exists
user_ids = {str(u.get('id')) for u in users if u.get('id')}
connections = [c for c in connections if str(c.get('user_id')) in user_ids]
tokens = [t for t in tokens if str(t.get('user_id')) in user_ids]
# Clear holder if user was deleted
for link in invite_links:
if link.get('user_id') and str(link['user_id']) not in user_ids:
link['user_id'] = ''
pool = get_pool()
with pool.connection() as conn:
with conn.cursor() as cur:
# Order matters for FKs: children first on delete, parents first on insert
cur.execute('DELETE FROM invite_links')
cur.execute('DELETE FROM api_tokens')
cur.execute('DELETE FROM user_connections')
cur.execute('DELETE FROM users')
cur.execute('DELETE FROM servers')
for pos, server in enumerate(servers):
cur.execute(
'INSERT INTO servers (position, name, host, ssh_port, username, '
'password, private_key, server_info, protocols) '
'VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s)',
(
pos,
server.get('name') or '',
server.get('host') or '',
int(server.get('ssh_port') or 22),
server.get('username') or '',
server.get('password'),
server.get('private_key'),
Jsonb(_as_dict(server.get('server_info'))),
Jsonb(_as_dict(server.get('protocols'))),
),
)
for user in users:
cur.execute(
'INSERT INTO users ('
'id, username, password_hash, role, enabled, created_at, '
'telegram_id, email, description, traffic_limit, traffic_used, '
'traffic_total, traffic_reset_strategy, last_reset_at, '
'expiration_date, expire_after_first_use, expiration_days, remnawave_uuid, xui_email, share_enabled, share_token, '
'share_password_hash'
') VALUES ('
'%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s'
')',
(
_as_uuid(user['id']),
user.get('username') or '',
user.get('password_hash') or '',
user.get('role') or 'user',
bool(user.get('enabled', True)),
_parse_ts(user.get('created_at')),
str(user['telegramId']) if user.get('telegramId') is not None else None,
user.get('email'),
user.get('description'),
int(user.get('traffic_limit') or 0),
int(user.get('traffic_used') or 0),
int(user.get('traffic_total') or 0),
user.get('traffic_reset_strategy') or 'never',
_parse_ts(user.get('last_reset_at')),
_parse_ts(user.get('expiration_date')),
bool(user.get('expire_after_first_use', False)),
int(user.get('expiration_days') or 0),
user.get('remnawave_uuid'),
user.get('xui_email'),
bool(user.get('share_enabled', False)),
user.get('share_token'),
user.get('share_password_hash'),
),
)
for conn_row in connections:
cur.execute(
'INSERT INTO user_connections ('
'id, user_id, server_id, protocol, client_id, name, xui_panel_id, created_at, last_bytes'
') VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s)',
(
_as_uuid(conn_row['id']),
_as_uuid(conn_row['user_id']),
int(conn_row.get('server_id') or 0),
conn_row.get('protocol') or '',
conn_row.get('client_id') or '',
conn_row.get('name') or '',
conn_row.get('xui_panel_id') or '',
_parse_ts(conn_row.get('created_at')),
int(conn_row.get('last_bytes') or 0),
),
)
for token in tokens:
cur.execute(
'INSERT INTO api_tokens ('
'id, name, token_hash, token_prefix, user_id, created_at, last_used_at'
') VALUES (%s, %s, %s, %s, %s, %s, %s)',
(
_as_uuid(token['id']),
token.get('name') or '',
token.get('token_hash') or '',
token.get('token_prefix') or '',
_as_uuid(token['user_id']),
_parse_ts(token.get('created_at')),
_parse_ts(token.get('last_used_at')),
),
)
for link in invite_links:
uid = link.get('user_id') or None
cur.execute(
'INSERT INTO invite_links ('
'id, name, token, enabled, max_uses, used_count, user_id, '
'protocol, server_id, xui_inbound_id, xui_panel_id, password_hash, expires_at, '
'duration_days, note, created_at'
') VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s)',
(
_as_uuid(link['id']),
link.get('name') or '',
link.get('token') or '',
bool(link.get('enabled', True)),
int(link.get('max_uses') or 0),
int(link.get('used_count') or 0),
_as_uuid(uid) if uid else None,
link.get('protocol') or 'awg',
int(link.get('server_id') or 0),
int(link.get('xui_inbound_id') or 0),
link.get('xui_panel_id') or '',
link.get('password_hash'),
_parse_ts(link.get('expires_at')),
int(link.get('duration_days') or 0),
link.get('note') or '',
_parse_ts(link.get('created_at')),
),
)
cur.execute(
'INSERT INTO settings (id, data) VALUES (1, %s) '
'ON CONFLICT (id) DO UPDATE SET data = EXCLUDED.data',
(Jsonb(settings),),
)
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:
"""Export current DB state as the legacy data.json document."""
return load_data()
def is_database_empty() -> bool:
init_schema()
pool = get_pool()
with pool.connection() as conn:
with conn.cursor() as cur:
cur.execute('SELECT COUNT(*) AS n FROM users')
users = cur.fetchone()['n']
cur.execute('SELECT COUNT(*) AS n FROM servers')
servers = cur.fetchone()['n']
return users == 0 and servers == 0
def import_from_json_file(path: str | os.PathLike, *, backup: bool = True) -> bool:
"""Import legacy data.json into Postgres. Returns True if import ran."""
path = Path(path)
if not path.exists():
return False
with path.open('r', encoding='utf-8') as f:
data = json.load(f)
if not isinstance(data, dict):
raise ValueError(f'Invalid data.json: expected object, got {type(data).__name__}')
data.setdefault('servers', [])
data.setdefault('users', [])
data.setdefault('user_connections', [])
data.setdefault('api_tokens', [])
data.setdefault('invite_links', [])
data.setdefault('settings', {})
save_data(data)
if backup:
stamp = datetime.now(timezone.utc).strftime('%Y%m%dT%H%M%SZ')
backup_path = path.with_name(f'data.json.migrated.{stamp}.bak')
try:
shutil.copy2(path, backup_path)
logger.info('Backed up legacy data.json to %s', backup_path)
except OSError as e:
logger.warning('Could not backup data.json: %s', e)
logger.info(
'Imported data.json → PostgreSQL (%s servers, %s users, %s connections)',
len(data.get('servers', [])),
len(data.get('users', [])),
len(data.get('user_connections', [])),
)
return True
def load_tunnel_state() -> dict:
init_schema()
pool = get_pool()
with pool.connection() as conn:
with conn.cursor() as cur:
cur.execute('SELECT provider, data FROM tunnel_state')
rows = cur.fetchall()
return {row['provider']: dict(row['data'] or {}) for row in rows}
def save_tunnel_state(state: dict) -> None:
init_schema()
pool = get_pool()
with pool.connection() as conn:
with conn.cursor() as cur:
cur.execute('DELETE FROM tunnel_state')
for provider, payload in (state or {}).items():
cur.execute(
'INSERT INTO tunnel_state (provider, data) VALUES (%s, %s)',
(str(provider), Jsonb(payload or {})),
)
conn.commit()
def update_tunnel_state(provider: str, **updates) -> None:
state = load_tunnel_state()
provider_state = state.get(provider, {})
provider_state.update(updates)
state[provider] = provider_state
save_tunnel_state(state)
def clear_tunnel_state(provider: str) -> None:
state = load_tunnel_state()
if provider in state:
state.pop(provider)
save_tunnel_state(state)
def ensure_db_ready(legacy_data_file: Optional[str] = None) -> None:
"""Init schema and one-shot import from data.json when DB is empty."""
init_schema()
if legacy_data_file and is_database_empty() and os.path.exists(legacy_data_file):
logger.info('Empty database — importing legacy %s', legacy_data_file)
import_from_json_file(legacy_data_file)