464 lines
15 KiB
Python
464 lines
15 KiB
Python
from __future__ import annotations
|
|
|
|
import asyncio
|
|
from datetime import datetime, timezone
|
|
from pathlib import Path
|
|
|
|
from sqlalchemy import func, select
|
|
from sqlalchemy.ext.asyncio import AsyncSession
|
|
from sqlalchemy.orm import selectinload
|
|
|
|
from app.config import settings
|
|
from app.models import Protocol, VpnClient, VpnServer
|
|
from app.services.crypto import (
|
|
generate_awg_params,
|
|
generate_keypair,
|
|
generate_preshared_key,
|
|
next_client_ip,
|
|
render_client_config,
|
|
render_server_config,
|
|
server_address,
|
|
)
|
|
from app.services.remote import deploy_config, install_vpn
|
|
from app.services.ssh import probe_tcp, verify_ssh
|
|
|
|
|
|
async def ensure_default_servers(session: AsyncSession) -> None:
|
|
existing = (await session.execute(select(VpnServer))).scalars().all()
|
|
by_protocol = {s.protocol: s for s in existing}
|
|
|
|
if Protocol.WIREGUARD.value not in by_protocol:
|
|
keys = generate_keypair()
|
|
private = settings.server_private_key or keys.private_key
|
|
public = settings.server_public_key or keys.public_key
|
|
if settings.server_private_key and settings.server_public_key:
|
|
private, public = settings.server_private_key, settings.server_public_key
|
|
elif settings.server_private_key:
|
|
private = settings.server_private_key
|
|
# derive not available without raw parse; keep generated public if only private set
|
|
session.add(
|
|
VpnServer(
|
|
name="WireGuard",
|
|
protocol=Protocol.WIREGUARD.value,
|
|
public_host=settings.public_host,
|
|
public_port=settings.public_port,
|
|
interface_name=settings.wg_interface,
|
|
subnet=settings.wg_subnet,
|
|
dns=settings.wg_dns,
|
|
server_private_key=private,
|
|
server_public_key=public,
|
|
)
|
|
)
|
|
|
|
if Protocol.AWG2.value not in by_protocol:
|
|
keys = generate_keypair()
|
|
private = settings.awg_server_private_key or keys.private_key
|
|
public = settings.awg_server_public_key or keys.public_key
|
|
session.add(
|
|
VpnServer(
|
|
name="AmneziaWG 2.0",
|
|
protocol=Protocol.AWG2.value,
|
|
public_host=settings.public_host,
|
|
public_port=settings.awg_public_port,
|
|
interface_name=settings.awg_interface,
|
|
subnet=settings.awg_subnet,
|
|
dns=settings.wg_dns,
|
|
server_private_key=private,
|
|
server_public_key=public,
|
|
jc=4,
|
|
jmin=40,
|
|
jmax=70,
|
|
s1=0,
|
|
s2=0,
|
|
h1=1,
|
|
h2=2,
|
|
h3=3,
|
|
h4=4,
|
|
)
|
|
)
|
|
|
|
await session.commit()
|
|
|
|
|
|
def _defaults_for_protocol(protocol: str) -> dict:
|
|
if protocol == Protocol.AWG2.value:
|
|
awg = generate_awg_params()
|
|
return {
|
|
"interface_name": settings.awg_interface,
|
|
"subnet": settings.awg_subnet,
|
|
"public_port": settings.awg_public_port,
|
|
"mtu": 1280,
|
|
"jc": awg.jc,
|
|
"jmin": awg.jmin,
|
|
"jmax": awg.jmax,
|
|
"s1": awg.s1,
|
|
"s2": awg.s2,
|
|
"s3": awg.s3,
|
|
"s4": awg.s4,
|
|
"h1": awg.h1,
|
|
"h2": awg.h2,
|
|
"h3": awg.h3,
|
|
"h4": awg.h4,
|
|
}
|
|
return {
|
|
"interface_name": settings.wg_interface,
|
|
"subnet": settings.wg_subnet,
|
|
"public_port": settings.public_port,
|
|
"mtu": 1420,
|
|
"jc": None,
|
|
"jmin": None,
|
|
"jmax": None,
|
|
"s1": None,
|
|
"s2": None,
|
|
"s3": None,
|
|
"s4": None,
|
|
"h1": None,
|
|
"h2": None,
|
|
"h3": None,
|
|
"h4": None,
|
|
}
|
|
|
|
|
|
async def create_server(
|
|
session: AsyncSession,
|
|
*,
|
|
name: str,
|
|
protocol: str,
|
|
ssh_host: str,
|
|
ssh_username: str,
|
|
ssh_port: int = 22,
|
|
ssh_password: str | None = None,
|
|
ssh_private_key: str | None = None,
|
|
public_host: str | None = None,
|
|
public_port: int | None = None,
|
|
dns: str = "1.1.1.1",
|
|
) -> VpnServer:
|
|
protocol = protocol.strip().lower()
|
|
if protocol not in {Protocol.WIREGUARD.value, Protocol.AWG2.value}:
|
|
raise ValueError("Unsupported protocol")
|
|
|
|
password = (ssh_password or "").strip() or None
|
|
private_key = (ssh_private_key or "").strip() or None
|
|
if not password and not private_key:
|
|
raise ValueError("Укажите пароль SSH или приватный ключ")
|
|
|
|
host = ssh_host.strip()
|
|
if not host:
|
|
raise ValueError("Укажите SSH host")
|
|
username = ssh_username.strip()
|
|
if not username:
|
|
raise ValueError("Укажите SSH логин")
|
|
|
|
# Like Amnezia-Web-Panel: verify SSH before saving
|
|
try:
|
|
server_info = await asyncio.to_thread(
|
|
verify_ssh, host, ssh_port or 22, username, password, private_key
|
|
)
|
|
except Exception as exc: # noqa: BLE001
|
|
raise ValueError(f"Подключение не удалось: {exc}") from exc
|
|
|
|
defaults = _defaults_for_protocol(protocol)
|
|
keys = generate_keypair()
|
|
# Unique interface per server id is assigned after insert; use temp then update
|
|
existing_count = await session.scalar(select(func.count()).select_from(VpnServer)) or 0
|
|
iface_base = "awg" if protocol == Protocol.AWG2.value else "wg"
|
|
interface_name = f"{iface_base}{existing_count}"
|
|
|
|
# Unique subnet per server: 10.{8+n}.0.0/24 or 10.{9+n}.0.0/24
|
|
octet = (8 if protocol == Protocol.WIREGUARD.value else 9) + int(existing_count)
|
|
if octet > 250:
|
|
raise ValueError("Слишком много серверов для авто-подсети")
|
|
subnet = f"10.{octet}.0.0/24"
|
|
|
|
now = datetime.now(timezone.utc)
|
|
server = VpnServer(
|
|
name=name.strip() or host,
|
|
protocol=protocol,
|
|
public_host=(public_host or host).strip(),
|
|
public_port=public_port or int(defaults["public_port"]),
|
|
interface_name=interface_name,
|
|
subnet=subnet,
|
|
dns=(dns or "1.1.1.1").strip(),
|
|
server_private_key=keys.private_key,
|
|
server_public_key=keys.public_key,
|
|
ssh_host=host,
|
|
ssh_port=ssh_port or 22,
|
|
ssh_username=username,
|
|
ssh_password=password,
|
|
ssh_private_key=private_key,
|
|
ssh_server_info=server_info,
|
|
ssh_reachable=True,
|
|
ssh_last_checked=now,
|
|
jc=defaults["jc"],
|
|
jmin=defaults["jmin"],
|
|
jmax=defaults["jmax"],
|
|
s1=defaults["s1"],
|
|
s2=defaults["s2"],
|
|
s3=defaults["s3"],
|
|
s4=defaults["s4"],
|
|
h1=defaults["h1"],
|
|
h2=defaults["h2"],
|
|
h3=defaults["h3"],
|
|
h4=defaults["h4"],
|
|
mtu=int(defaults["mtu"]),
|
|
vpn_installed=False,
|
|
is_enabled=True,
|
|
)
|
|
session.add(server)
|
|
await session.commit()
|
|
await session.refresh(server)
|
|
await sync_server_config(session, server.id)
|
|
return server
|
|
|
|
|
|
async def check_server_ssh(session: AsyncSession, server_id: int) -> VpnServer:
|
|
server = await session.get(VpnServer, server_id)
|
|
if not server:
|
|
raise ValueError("Server not found")
|
|
if not server.ssh_host or not server.ssh_username:
|
|
raise ValueError("SSH данные не заданы")
|
|
if not server.ssh_password and not server.ssh_private_key:
|
|
raise ValueError("Нет пароля и ключа SSH")
|
|
|
|
now = datetime.now(timezone.utc)
|
|
try:
|
|
info = await asyncio.to_thread(
|
|
verify_ssh,
|
|
server.ssh_host,
|
|
server.ssh_port or 22,
|
|
server.ssh_username,
|
|
server.ssh_password,
|
|
server.ssh_private_key,
|
|
)
|
|
server.ssh_server_info = info
|
|
server.ssh_reachable = True
|
|
server.ssh_last_checked = now
|
|
except Exception as exc: # noqa: BLE001
|
|
server.ssh_reachable = False
|
|
server.ssh_last_checked = now
|
|
await session.commit()
|
|
raise ValueError(f"Подключение не удалось: {exc}") from exc
|
|
|
|
await session.commit()
|
|
await session.refresh(server)
|
|
return server
|
|
|
|
|
|
async def probe_servers(servers: list[VpnServer]) -> dict[int, dict]:
|
|
"""TCP ping SSH ports for UI online indicators."""
|
|
|
|
async def one(s: VpnServer) -> tuple[int, dict]:
|
|
if not s.ssh_host:
|
|
return s.id, {"online": False, "latency_ms": None}
|
|
result = await asyncio.to_thread(probe_tcp, s.ssh_host, s.ssh_port or 22)
|
|
return s.id, {"online": result.online, "latency_ms": result.latency_ms}
|
|
|
|
pairs = await asyncio.gather(*(one(s) for s in servers))
|
|
return dict(pairs)
|
|
|
|
|
|
async def update_server_settings(
|
|
session: AsyncSession,
|
|
server_id: int,
|
|
*,
|
|
name: str | None = None,
|
|
public_host: str,
|
|
public_port: int,
|
|
dns: str = "1.1.1.1",
|
|
ssh_host: str | None = None,
|
|
ssh_port: int = 22,
|
|
ssh_username: str | None = None,
|
|
ssh_password: str | None = None,
|
|
ssh_private_key: str | None = None,
|
|
clear_ssh_password: bool = False,
|
|
clear_ssh_private_key: bool = False,
|
|
) -> VpnServer:
|
|
server = await session.get(VpnServer, server_id)
|
|
if not server:
|
|
raise ValueError("Server not found")
|
|
|
|
if name is not None and name.strip():
|
|
server.name = name.strip()
|
|
server.public_host = public_host.strip()
|
|
server.public_port = public_port
|
|
server.dns = dns.strip() or "1.1.1.1"
|
|
|
|
if ssh_host is not None:
|
|
server.ssh_host = ssh_host.strip() or None
|
|
server.ssh_port = ssh_port or 22
|
|
if ssh_username is not None:
|
|
server.ssh_username = ssh_username.strip() or None
|
|
|
|
if clear_ssh_password:
|
|
server.ssh_password = None
|
|
elif ssh_password is not None and ssh_password.strip():
|
|
server.ssh_password = ssh_password.strip()
|
|
|
|
if clear_ssh_private_key:
|
|
server.ssh_private_key = None
|
|
elif ssh_private_key is not None and ssh_private_key.strip():
|
|
server.ssh_private_key = ssh_private_key.strip()
|
|
|
|
await session.commit()
|
|
await session.refresh(server)
|
|
await sync_server_config(session, server_id)
|
|
return server
|
|
|
|
|
|
async def create_client(
|
|
session: AsyncSession,
|
|
*,
|
|
server_id: int,
|
|
name: str,
|
|
notes: str | None = None,
|
|
allowed_ips: str = "0.0.0.0/0, ::/0",
|
|
) -> VpnClient:
|
|
server = await session.get(VpnServer, server_id)
|
|
if not server:
|
|
raise ValueError("Server not found")
|
|
|
|
used = (
|
|
await session.execute(select(VpnClient.address).where(VpnClient.server_id == server_id))
|
|
).scalars().all()
|
|
address = next_client_ip(server.subnet, list(used))
|
|
keys = generate_keypair()
|
|
client = VpnClient(
|
|
server_id=server_id,
|
|
name=name.strip(),
|
|
private_key=keys.private_key,
|
|
public_key=keys.public_key,
|
|
preshared_key=generate_preshared_key(),
|
|
address=address,
|
|
allowed_ips=allowed_ips,
|
|
notes=notes,
|
|
is_enabled=True,
|
|
)
|
|
session.add(client)
|
|
await session.commit()
|
|
await session.refresh(client)
|
|
await sync_server_config(session, server_id)
|
|
return client
|
|
|
|
|
|
async def set_client_enabled(session: AsyncSession, client_id: int, enabled: bool) -> VpnClient:
|
|
client = await session.get(VpnClient, client_id)
|
|
if not client:
|
|
raise ValueError("Client not found")
|
|
client.is_enabled = enabled
|
|
await session.commit()
|
|
await session.refresh(client)
|
|
await sync_server_config(session, client.server_id)
|
|
return client
|
|
|
|
|
|
async def delete_client(session: AsyncSession, client_id: int) -> None:
|
|
client = await session.get(VpnClient, client_id)
|
|
if not client:
|
|
raise ValueError("Client not found")
|
|
server_id = client.server_id
|
|
await session.delete(client)
|
|
await session.commit()
|
|
await sync_server_config(session, server_id)
|
|
|
|
|
|
async def get_client_config(session: AsyncSession, client_id: int) -> str:
|
|
result = await session.execute(
|
|
select(VpnClient).options(selectinload(VpnClient.server)).where(VpnClient.id == client_id)
|
|
)
|
|
client = result.scalar_one_or_none()
|
|
if not client:
|
|
raise ValueError("Client not found")
|
|
server = client.server
|
|
dns = server.dns
|
|
if "," not in dns:
|
|
dns = f"{dns}, 1.0.0.1"
|
|
return render_client_config(
|
|
protocol=server.protocol,
|
|
client_private_key=client.private_key,
|
|
client_address=client.address,
|
|
dns=dns,
|
|
server_public_key=server.server_public_key,
|
|
preshared_key=client.preshared_key,
|
|
endpoint_host=server.public_host,
|
|
endpoint_port=server.public_port,
|
|
allowed_ips=client.allowed_ips,
|
|
mtu=server.mtu or (1280 if server.protocol == Protocol.AWG2.value else 1420),
|
|
jc=server.jc,
|
|
jmin=server.jmin,
|
|
jmax=server.jmax,
|
|
s1=server.s1,
|
|
s2=server.s2,
|
|
s3=server.s3,
|
|
s4=server.s4,
|
|
h1=server.h1,
|
|
h2=server.h2,
|
|
h3=server.h3,
|
|
h4=server.h4,
|
|
)
|
|
|
|
|
|
def _build_server_config_text(server: VpnServer) -> str:
|
|
peers = [
|
|
{
|
|
"name": c.name,
|
|
"public_key": c.public_key,
|
|
"preshared_key": c.preshared_key,
|
|
"address": c.address if c.address.endswith("/32") else f"{c.address.split('/')[0]}/32",
|
|
}
|
|
for c in server.clients
|
|
if c.is_enabled
|
|
]
|
|
return render_server_config(
|
|
protocol=server.protocol,
|
|
interface_name=server.interface_name,
|
|
private_key=server.server_private_key,
|
|
address=server_address(server.subnet),
|
|
listen_port=server.public_port,
|
|
peers=peers,
|
|
jc=server.jc,
|
|
jmin=server.jmin,
|
|
jmax=server.jmax,
|
|
s1=server.s1,
|
|
s2=server.s2,
|
|
s3=server.s3,
|
|
s4=server.s4,
|
|
h1=server.h1,
|
|
h2=server.h2,
|
|
h3=server.h3,
|
|
h4=server.h4,
|
|
)
|
|
|
|
|
|
async def sync_server_config(session: AsyncSession, server_id: int, push_remote: bool = True) -> Path:
|
|
result = await session.execute(
|
|
select(VpnServer).options(selectinload(VpnServer.clients)).where(VpnServer.id == server_id)
|
|
)
|
|
server = result.scalar_one()
|
|
content = _build_server_config_text(server)
|
|
|
|
base = Path(settings.wg_config_dir if server.protocol == Protocol.WIREGUARD.value else settings.awg_config_dir)
|
|
base.mkdir(parents=True, exist_ok=True)
|
|
path = base / f"{server.interface_name}.conf"
|
|
path.write_text(content, encoding="utf-8")
|
|
|
|
if push_remote and server.vpn_installed and server.ssh_host:
|
|
await asyncio.to_thread(deploy_config, server, content)
|
|
|
|
return path
|
|
|
|
|
|
async def install_vpn_on_server(session: AsyncSession, server_id: int) -> list[str]:
|
|
result = await session.execute(
|
|
select(VpnServer).options(selectinload(VpnServer.clients)).where(VpnServer.id == server_id)
|
|
)
|
|
server = result.scalar_one_or_none()
|
|
if not server:
|
|
raise ValueError("Server not found")
|
|
content = _build_server_config_text(server)
|
|
# Always save locally first
|
|
await sync_server_config(session, server_id, push_remote=False)
|
|
logs = await asyncio.to_thread(install_vpn, server, content)
|
|
server.vpn_installed = True
|
|
await session.commit()
|
|
return logs
|