@@ -0,0 +1,201 @@
|
||||
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
|
||||
Reference in New Issue
Block a user