202 lines
6.4 KiB
Python
202 lines
6.4 KiB
Python
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
|