Files
wireguard-admin/app/wireguard.py
T
lofyerandfactory-droid[bot] <138933559+factory-droid[bot]@users.noreply.github.com> 192960ad3e Advanced networking: per-peer site subnets (site-to-site), per-peer client routes, peer isolation option, client setup guide
Co-authored-by: factory-droid[bot] <138933559+factory-droid[bot]@users.noreply.github.com>
2026-07-04 08:21:17 +08:00

263 lines
8.3 KiB
Python

import ipaddress
import subprocess
from dataclasses import dataclass, field
from datetime import datetime, timezone
from .config import settings
from .crypto import decrypt
from .models import Peer
ONLINE_THRESHOLD_SECONDS = 180
def _run(args: list[str], input_text: str | None = None) -> str:
result = subprocess.run(
args, input=input_text, capture_output=True, text=True, check=True
)
return result.stdout.strip()
def genkey() -> str:
return _run(["wg", "genkey"])
def genpsk() -> str:
return _run(["wg", "genpsk"])
def pubkey(private_key: str) -> str:
return _run(["wg", "pubkey"], input_text=private_key)
def generate_keypair() -> tuple[str, str]:
private = genkey()
return private, pubkey(private)
@dataclass
class PeerStatus:
public_key: str
endpoint: str | None = None
latest_handshake: datetime | None = None
rx_bytes: int = 0
tx_bytes: int = 0
@property
def online(self) -> bool:
if self.latest_handshake is None:
return False
delta = datetime.now(timezone.utc) - self.latest_handshake
return delta.total_seconds() < ONLINE_THRESHOLD_SECONDS
@dataclass
class InterfaceStatus:
name: str
public_key: str | None = None
listen_port: int | None = None
up: bool = False
peers: dict[str, PeerStatus] = field(default_factory=dict)
@property
def total_rx(self) -> int:
return sum(p.rx_bytes for p in self.peers.values())
@property
def total_tx(self) -> int:
return sum(p.tx_bytes for p in self.peers.values())
def get_status() -> InterfaceStatus:
status = InterfaceStatus(name=settings.wg_interface)
try:
output = _run(["wg", "show", settings.wg_interface, "dump"])
except (subprocess.CalledProcessError, FileNotFoundError):
return status
status.up = True
lines = output.splitlines()
if lines:
fields = lines[0].split("\t")
if len(fields) >= 3:
status.public_key = fields[1]
status.listen_port = int(fields[2])
for line in lines[1:]:
fields = line.split("\t")
if len(fields) < 8:
continue
peer = PeerStatus(public_key=fields[0])
if fields[2] != "(none)":
peer.endpoint = fields[2]
handshake = int(fields[4])
if handshake:
peer.latest_handshake = datetime.fromtimestamp(handshake, tz=timezone.utc)
peer.rx_bytes = int(fields[5])
peer.tx_bytes = int(fields[6])
status.peers[peer.public_key] = peer
return status
def _server_key_paths() -> tuple:
key_dir = settings.data_dir / "server"
return key_dir / "privatekey", key_dir / "publickey"
def ensure_server_keys() -> tuple[str, str]:
private_path, public_path = _server_key_paths()
if private_path.exists() and public_path.exists():
return private_path.read_text().strip(), public_path.read_text().strip()
private_path.parent.mkdir(parents=True, exist_ok=True)
private, public = generate_keypair()
private_path.touch(mode=0o600)
private_path.write_text(private + "\n")
public_path.write_text(public + "\n")
return private, public
def server_address() -> str:
network = ipaddress.ip_network(settings.wg_subnet)
return f"{next(network.hosts())}/{network.prefixlen}"
def validate_address(address: str, taken: list[str]) -> str:
network = ipaddress.ip_network(settings.wg_subnet)
try:
ip = ipaddress.ip_address(address.split("/")[0].strip())
except ValueError:
raise ValueError(f"Invalid IP address: {address}")
if ip not in network:
raise ValueError(f"{ip} is not in subnet {settings.wg_subnet}")
if ip in (network.network_address, network.broadcast_address):
raise ValueError(f"{ip} is not a usable host address")
if ip == next(network.hosts()):
raise ValueError(f"{ip} is reserved for the server")
if ip in {ipaddress.ip_interface(a).ip for a in taken}:
raise ValueError(f"{ip} is already assigned to another peer")
return f"{ip}/32"
def next_free_address(taken: list[str]) -> str:
network = ipaddress.ip_network(settings.wg_subnet)
used = {ipaddress.ip_interface(a).ip for a in taken}
hosts = network.hosts()
used.add(next(hosts))
for host in hosts:
if host not in used:
return f"{host}/32"
raise RuntimeError(f"No free addresses left in {settings.wg_subnet}")
def validate_cidr_list(value: str) -> str:
networks = []
for part in value.split(","):
part = part.strip()
if not part:
continue
try:
networks.append(str(ipaddress.ip_network(part, strict=False)))
except ValueError:
raise ValueError(f"Invalid CIDR: {part}")
return ", ".join(networks)
def server_allowed_ips(peer: Peer) -> str:
allowed = peer.address
if peer.extra_allowed_ips:
allowed += f", {peer.extra_allowed_ips}"
return allowed
def render_server_config(private_key: str, peers: list[Peer]) -> str:
lines = [
"[Interface]",
f"PrivateKey = {private_key}",
f"Address = {server_address()}",
f"ListenPort = {settings.wg_port}",
]
# wg-quick adds routes for /32 peer addresses automatically; extra
# site subnets need explicit routes so return traffic enters the tunnel.
for peer in peers:
if peer.enabled and peer.extra_allowed_ips:
for subnet in peer.extra_allowed_ips.split(","):
subnet = subnet.strip()
lines.append(f"PostUp = ip route replace {subnet} dev %i")
for peer in peers:
if not peer.enabled:
continue
lines += [
"",
"[Peer]",
f"# {peer.name}",
f"PublicKey = {peer.public_key}",
f"PresharedKey = {decrypt(peer.preshared_key_enc)}",
f"AllowedIPs = {server_allowed_ips(peer)}",
]
return "\n".join(lines) + "\n"
def render_client_config(peer: Peer, server_public_key: str) -> str:
return "\n".join(
[
"[Interface]",
f"PrivateKey = {decrypt(peer.private_key_enc)}",
f"Address = {peer.address}",
f"DNS = {peer.dns or settings.wg_dns}",
"",
"[Peer]",
f"PublicKey = {server_public_key}",
f"PresharedKey = {decrypt(peer.preshared_key_enc)}",
f"Endpoint = {settings.wg_host}:{settings.wg_port}",
f"AllowedIPs = {peer.client_allowed_ips or settings.wg_allowed_ips}",
f"PersistentKeepalive = {settings.wg_persistent_keepalive}",
]
) + "\n"
def write_server_config(private_key: str, peers: list[Peer]) -> None:
settings.wg_config_dir.mkdir(parents=True, exist_ok=True)
config_path = settings.wg_config_dir / f"{settings.wg_interface}.conf"
config_path.touch(mode=0o600)
config_path.write_text(render_server_config(private_key, peers))
def _is_managed(status: InterfaceStatus) -> bool:
_, server_public = ensure_server_keys()
return status.public_key == server_public
def interface_up() -> None:
status = get_status()
if status.up and not _is_managed(status):
raise RuntimeError(
f"Interface {settings.wg_interface} is up but uses a foreign key; "
"refusing to manage it. Set WG_INTERFACE to a dedicated interface."
)
if not status.up:
_run(["wg-quick", "up", settings.wg_interface])
def sync_routes(peers: list[Peer]) -> None:
for peer in peers:
if not peer.extra_allowed_ips:
continue
for subnet in peer.extra_allowed_ips.split(","):
subnet = subnet.strip()
args = ["ip", "route", "replace" if peer.enabled else "del", subnet]
if peer.enabled:
args += ["dev", settings.wg_interface]
try:
_run(args)
except subprocess.CalledProcessError:
pass
def sync_peers(private_key: str, peers: list[Peer]) -> None:
write_server_config(private_key, peers)
status = get_status()
if status.up and _is_managed(status):
stripped = _run(
["wg-quick", "strip", str(settings.wg_config_dir / f"{settings.wg_interface}.conf")]
)
_run(["wg", "syncconf", settings.wg_interface, "/dev/stdin"], input_text=stripped)
sync_routes(peers)