Files
projects/Dockers/gluetun-pia-wireguard-rotator/rotate.py
T
Bram bc9d79de02
Build and Push Docker Images / build-and-push (push) Successful in 17s
optimize
2026-08-14 20:03:05 +02:00

412 lines
15 KiB
Python

#!/usr/bin/env python3
"""Pick a PIA region/server and write a WireGuard config for Gluetun."""
from __future__ import annotations
import json
import os
import random
import re
import socket
import subprocess
import sys
import tempfile
import time
from concurrent.futures import ThreadPoolExecutor, as_completed
from datetime import datetime
from pathlib import Path
from typing import Any
import pia
def log(msg: str) -> None:
print(f"[{datetime.now().astimezone().isoformat(timespec='seconds')}] {msg}", file=sys.stderr)
def require_env(name: str) -> str:
value = os.environ.get(name, "")
if not value:
raise SystemExit(f"Missing required env var: {name}")
return value
def parse_list(raw: str) -> list[str]:
raw = raw.strip()
if not raw:
return []
if raw.startswith("["):
data = json.loads(raw)
if not isinstance(data, list):
raise SystemExit("Expected a JSON array")
return [str(item).strip() for item in data if str(item).strip()]
return [part.strip() for part in raw.split(",") if part.strip()]
def parse_regions() -> list[str]:
regions = parse_list(require_env("PIA_REGIONS"))
if not regions:
raise SystemExit("PIA_REGIONS is empty after parsing")
return regions
def read_state(state_path: Path) -> dict[str, Any]:
if not state_path.is_file():
return {}
try:
data = json.loads(state_path.read_text(encoding="utf-8"))
except (OSError, json.JSONDecodeError):
return {}
return data if isinstance(data, dict) else {}
def tcp_latency_ms(ip: str, port: int, timeout: float) -> float | None:
sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
sock.settimeout(timeout)
started = time.perf_counter()
try:
sock.connect((ip, port))
return (time.perf_counter() - started) * 1000.0
except OSError:
return None
finally:
sock.close()
def average_tcp_latency_ms(ip: str, port: int, timeout: float, samples: int) -> float | None:
readings: list[float] = []
for _ in range(samples):
value = tcp_latency_ms(ip, port, timeout)
if value is None:
return None
readings.append(value)
return sum(readings) / len(readings)
def pick_fastest(
candidates: list[str],
previous_region: str | None,
previous_server_ip: str | None,
) -> tuple[pia.WgServer, list[dict[str, Any]]]:
port = int(os.environ.get("LATENCY_PORT", "1337"))
timeout = float(os.environ.get("LATENCY_TIMEOUT_SECONDS", "2"))
samples = max(1, int(os.environ.get("LATENCY_SAMPLES", "2")))
margin = float(os.environ.get("LATENCY_SWITCH_MARGIN_MS", "15"))
serverlist = pia.fetch_serverlist()
servers: list[pia.WgServer] = []
for region_id in candidates:
servers.extend(pia.region_wg_servers(serverlist, region_id))
if not servers:
raise SystemExit("No reachable WireGuard servers found for configured regions")
best_by_region: dict[str, dict[str, Any]] = {}
with ThreadPoolExecutor(max_workers=min(32, len(servers))) as pool:
futures = {
pool.submit(average_tcp_latency_ms, server.ip, port, timeout, samples): server
for server in servers
}
for future in as_completed(futures):
server = futures[future]
value = future.result()
current = best_by_region.setdefault(
server.region,
{
"region": server.region,
"latency_ms": None,
"server_ip": None,
"server_cn": None,
"servers": 0,
"failures": 0,
},
)
current["servers"] += 1
if value is None:
current["failures"] += 1
continue
if current["latency_ms"] is None or value < current["latency_ms"]:
current["latency_ms"] = round(value, 2)
current["server_ip"] = server.ip
current["server_cn"] = server.cn
results = sorted(
best_by_region.values(),
key=lambda item: (
item["latency_ms"] is None,
item["latency_ms"] if item["latency_ms"] is not None else float("inf"),
item["region"],
),
)
for item in results:
latency = "timeout" if item["latency_ms"] is None else f"{item['latency_ms']:.2f}ms"
log(
f"Latency {item['region']}: {latency} "
f"(best={item['server_cn']}/{item['server_ip']}, "
f"servers={item['servers']}, failures={item['failures']})"
)
reachable = [item for item in results if item["latency_ms"] is not None]
if not reachable:
raise SystemExit("All latency probes failed; cannot select fastest region")
winner = reachable[0]
if previous_region:
previous_result = next((item for item in reachable if item["region"] == previous_region), None)
if previous_result is not None:
improvement = previous_result["latency_ms"] - winner["latency_ms"]
same_server = (
previous_server_ip
and previous_result["server_ip"] == previous_server_ip
)
if winner["region"] != previous_region and improvement < margin:
log(
f"Keeping current region {previous_region} "
f"({previous_result['latency_ms']:.2f}ms); "
f"best {winner['region']} only {improvement:.2f}ms faster "
f"(margin {margin:g}ms)"
)
winner = previous_result
elif winner["region"] == previous_region and same_server:
log(f"Current endpoint {previous_server_ip} is still fastest in {previous_region}")
elif winner["region"] == previous_region:
log(f"Current region {previous_region} is still fastest")
server = pia.WgServer(
region=winner["region"],
ip=winner["server_ip"],
cn=winner["server_cn"],
)
log(f"Selected {server.region} via {server.cn} ({server.ip}) at {winner['latency_ms']:.2f}ms")
return server, results
def pick_random(candidates: list[str], previous_region: str | None) -> pia.WgServer:
serverlist = pia.fetch_serverlist()
pool = [region for region in candidates if region != previous_region] or list(candidates)
if previous_region and previous_region not in pool and len(candidates) == 1:
log(f"Only one region configured; reusing previous: {previous_region}")
elif previous_region and previous_region not in pool:
log(f"Excluding previous region: {previous_region}")
random.shuffle(pool)
for region_id in pool:
servers = pia.region_wg_servers(serverlist, region_id)
if servers:
server = random.choice(servers)
log(f"Randomly selected {server.region} via {server.cn} ({server.ip})")
return server
raise SystemExit("No WireGuard servers found for configured regions")
def pick_server(state_path: Path) -> tuple[pia.WgServer, str, list[dict[str, Any]]]:
regions = parse_regions()
state = read_state(state_path)
previous_region = state.get("region") if isinstance(state.get("region"), str) else None
previous_server_ip = state.get("server_ip") if isinstance(state.get("server_ip"), str) else None
mode = os.environ.get("REGION_SELECT", "fastest").strip().lower() or "fastest"
if mode == "random":
return pick_random(regions, previous_region), mode, []
if mode != "fastest":
raise SystemExit(f"Invalid REGION_SELECT '{mode}' (expected fastest|random)")
server, results = pick_fastest(regions, previous_region, previous_server_ip)
return server, mode, results
def parse_restart_containers() -> list[str]:
gluetun = os.environ.get("GLUETUN_CONTAINER", "m3u-filter-vpn").strip() or "m3u-filter-vpn"
ordered = [gluetun]
seen = {gluetun}
for container in parse_list(os.environ.get("RESTART_CONTAINERS", "")):
if container in seen:
continue
ordered.append(container)
seen.add(container)
return ordered
def restart_containers(containers: list[str]) -> None:
if not containers:
raise SystemExit("No containers configured to restart")
for container in containers:
log(f"Restarting container: {container}")
result = subprocess.run(["docker", "restart", container], check=False)
if result.returncode != 0:
raise SystemExit(f"docker restart failed for container={container}")
def config_age_seconds(wg_path: Path) -> float | None:
if not wg_path.is_file():
return None
try:
return time.time() - wg_path.stat().st_mtime
except OSError:
return None
def should_skip_regeneration(
server: pia.WgServer,
state: dict[str, Any],
wg_path: Path,
latency_results: list[dict[str, Any]],
) -> bool:
if pia.env_bool("FORCE_ROTATE", False):
log("FORCE_ROTATE=true; regenerating WireGuard config")
return False
previous_region = state.get("region")
previous_server_ip = state.get("server_ip")
if previous_region != server.region:
return False
age = config_age_seconds(wg_path)
if age is None:
return False
max_age = pia.parse_duration_seconds(os.environ.get("WG_CONFIG_MAX_AGE", "7d"), 604800)
if age > max_age:
log(f"Existing WireGuard config is stale (age {int(age)}s > max {max_age}s); regenerating")
return False
if not re.search(
r"^\[Interface\]",
wg_path.read_text(encoding="utf-8", errors="replace"),
flags=re.MULTILINE,
):
return False
# Same region, different server: only regenerate if clearly faster.
if previous_server_ip and server.ip != previous_server_ip:
margin = float(os.environ.get("LATENCY_SWITCH_MARGIN_MS", "15"))
previous_latency = state.get("latency_ms")
winner = next((item for item in latency_results if item.get("server_ip") == server.ip), None)
new_latency = winner.get("latency_ms") if winner else None
if (
isinstance(previous_latency, (int, float))
and isinstance(new_latency, (int, float))
and (previous_latency - new_latency) < margin
):
log(
f"Keeping current server {previous_server_ip}; "
f"{server.ip} only {previous_latency - new_latency:.2f}ms faster "
f"(margin {margin:g}ms)"
)
return True
log(f"Switching server within {server.region}: {previous_server_ip} -> {server.ip}")
return False
log(
f"Skipping PIA token/addKey; endpoint unchanged "
f"({server.region}/{server.ip}) and config age {int(age)}s <= {max_age}s"
)
return True
def write_state(
state_path: Path,
server: pia.WgServer,
mode: str,
restarted: list[str],
latency_results: list[dict[str, Any]],
*,
skipped: bool = False,
) -> None:
previous = read_state(state_path)
winner = next(
(
item
for item in latency_results
if item.get("region") == server.region and item.get("server_ip") == server.ip
),
next((item for item in latency_results if item.get("region") == server.region), None),
)
now = datetime.now().astimezone().isoformat(timespec="seconds")
payload: dict[str, Any] = {
"region": server.region,
"server_ip": server.ip,
"server_cn": server.cn,
"selection": mode,
"wg_config": os.environ.get("WG_CONFIG_PATH", "/config/wireguard/wg0.conf"),
"restarted_containers": restarted,
"last_checked_at": now,
"skipped_regeneration": skipped,
}
if skipped and isinstance(previous.get("rotated_at"), str):
payload["rotated_at"] = previous["rotated_at"]
else:
payload["rotated_at"] = now
if winner and winner.get("latency_ms") is not None:
payload["latency_ms"] = winner["latency_ms"]
elif skipped and previous.get("latency_ms") is not None:
payload["latency_ms"] = previous["latency_ms"]
if latency_results:
payload["latency_results"] = latency_results
state_path.parent.mkdir(parents=True, exist_ok=True)
state_path.write_text(json.dumps(payload, indent=2) + "\n", encoding="utf-8")
state_path.chmod(0o644)
def rotate_once() -> None:
require_env("PIA_USER")
require_env("PIA_PASS")
require_env("PIA_REGIONS")
state_path = Path(os.environ.get("ROTATOR_STATE_PATH", "/config/rotator-state.json"))
wg_path = Path(os.environ.get("WG_CONFIG_PATH", "/config/wireguard/wg0.conf"))
state = read_state(state_path)
server, mode, latency_results = pick_server(state_path)
log(f"Selected endpoint: {server.region} / {server.cn} / {server.ip} (mode={mode})")
if should_skip_regeneration(server, state, wg_path, latency_results):
# If we decided to keep the previous server IP, persist that identity.
keep_ip = state.get("server_ip") if isinstance(state.get("server_ip"), str) else server.ip
keep_cn = state.get("server_cn") if isinstance(state.get("server_cn"), str) else server.cn
if keep_ip != server.ip:
server = pia.WgServer(region=server.region, ip=keep_ip, cn=str(keep_cn or server.cn))
write_state(state_path, server, mode, [], latency_results, skipped=True)
log(f"Rotation skipped for {server.region}/{server.ip}")
return
with tempfile.TemporaryDirectory(prefix="pia-rotate-") as tmp:
tmp_conf = Path(tmp) / "wg0.conf"
log("Generating WireGuard config via native PIA client")
pia.generate_wg_config(
require_env("PIA_USER"),
require_env("PIA_PASS"),
server,
tmp_conf,
)
text = tmp_conf.read_text(encoding="utf-8", errors="replace")
if not re.search(r"^\[Interface\]", text, flags=re.MULTILINE):
raise SystemExit("Generated config missing [Interface] section")
wg_path.parent.mkdir(parents=True, exist_ok=True)
os.replace(tmp_conf, wg_path)
wg_path.chmod(0o600)
log(f"Wrote {wg_path}")
restarted = parse_restart_containers()
restart_containers(restarted)
write_state(state_path, server, mode, restarted, latency_results, skipped=False)
log(f"Rotation complete for {server.region}/{server.ip}")
def list_regions_main() -> None:
for region in pia.list_regions():
print(
f"{region['id']:<35} {region['name']:<35} "
f"country={region['country']:<4} port-forward={region['port_forward']!s:<5} "
f"offline={region['offline']} wg={region['wg_servers']}"
)
if __name__ == "__main__":
if len(sys.argv) > 1 and sys.argv[1] in {"--list-regions", "list-regions"}:
list_regions_main()
else:
rotate_once()