from __future__ import annotations from pathlib import Path from sqlalchemy import 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_keypair, generate_preshared_key, next_client_ip, render_client_config, render_server_config, server_address, ) 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() 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 return render_client_config( protocol=server.protocol, client_private_key=client.private_key, client_address=client.address, dns=server.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, jc=server.jc, jmin=server.jmin, jmax=server.jmax, s1=server.s1, s2=server.s2, h1=server.h1, h2=server.h2, h3=server.h3, h4=server.h4, ) async def sync_server_config(session: AsyncSession, server_id: int) -> Path: result = await session.execute( select(VpnServer).options(selectinload(VpnServer.clients)).where(VpnServer.id == server_id) ) server = result.scalar_one() 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 ] content = 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, h1=server.h1, h2=server.h2, h3=server.h3, h4=server.h4, ) 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") return path