vpn: require an active WireGuard handshake before downloading (refuse a dead tunnel)

This commit is contained in:
Konstantin Passig PC
2026-08-25 17:06:27 +02:00
parent b5d56d4d34
commit 85ae1fb184
4 changed files with 170 additions and 3 deletions
+10 -1
View File
@@ -11,6 +11,7 @@ from pathlib import Path
from . import __version__, vpn from . import __version__, vpn
from .config import Config, QUALITY_PRESETS, default_config_path, load_config, write_default_config from .config import Config, QUALITY_PRESETS, default_config_path, load_config, write_default_config
from .downloader import cookie_opts, download_video from .downloader import cookie_opts, download_video
from .progress import DownloadProgress
def log_setup(verbose: bool) -> None: def log_setup(verbose: bool) -> None:
@@ -74,7 +75,8 @@ def cmd_download(args: argparse.Namespace) -> int:
extra = cookie_opts(cfg) extra = cookie_opts(cfg)
print(f"downloading: {args.url} (quality={quality}, out={out_dir.expanduser()})") print(f"downloading: {args.url} (quality={quality}, out={out_dir.expanduser()})")
path = download_video(args.url, out_dir, quality, format_override, extra) bar = DownloadProgress()
path = download_video(args.url, out_dir, quality, format_override, extra, progress=bar)
if not path: if not path:
logging.error("download failed: %s", args.url) logging.error("download failed: %s", args.url)
return 1 return 1
@@ -91,6 +93,13 @@ def cmd_vpn(args: argparse.Namespace) -> int:
else: else:
vpn.up(cfg) vpn.up(cfg)
print("tunnel up") print("tunnel up")
if not vpn.wait_for_handshake(cfg):
logging.error(
"tunnel is up but no WireGuard handshake was established — "
"downloads will be refused until the tunnel is active"
)
return 1
print("handshake established")
elif args.vpn_command == "down": elif args.vpn_command == "down":
vpn.down(cfg) vpn.down(cfg)
print("tunnel down") print("tunnel down")
+6
View File
@@ -8,6 +8,8 @@ from typing import Optional
import yt_dlp import yt_dlp
from .progress import DownloadProgress
log = logging.getLogger(__name__) log = logging.getLogger(__name__)
_BASE_OPTS = { _BASE_OPTS = {
@@ -16,6 +18,7 @@ _BASE_OPTS = {
"noplaylist": True, "noplaylist": True,
"ignoreerrors": True, "ignoreerrors": True,
"no_color": True, "no_color": True,
"noprogress": True, # we draw our own progress bar via progress_hooks
} }
# named presets -> yt-dlp format strings. bestvideo+bestaudio requires ffmpeg # named presets -> yt-dlp format strings. bestvideo+bestaudio requires ffmpeg
@@ -54,17 +57,20 @@ def download_video(
quality: str = "best", quality: str = "best",
format_override: Optional[str] = None, format_override: Optional[str] = None,
extra_opts: dict | None = None, extra_opts: dict | None = None,
progress: DownloadProgress | None = None,
) -> Optional[Path]: ) -> Optional[Path]:
"""Download a single video into out_dir. Returns the resulting file path.""" """Download a single video into out_dir. Returns the resulting file path."""
out_dir = out_dir.expanduser() out_dir = out_dir.expanduser()
out_dir.mkdir(parents=True, exist_ok=True) out_dir.mkdir(parents=True, exist_ok=True)
fmt = _select_format(quality, format_override) fmt = _select_format(quality, format_override)
bar = progress or DownloadProgress()
opts = { opts = {
**_BASE_OPTS, **_BASE_OPTS,
**(extra_opts or {}), **(extra_opts or {}),
"format": fmt, "format": fmt,
"outtmpl": str(out_dir / "%(title).150B [%(id)s].%(ext)s"), "outtmpl": str(out_dir / "%(title).150B [%(id)s].%(ext)s"),
"nocheckcertificate": True, "nocheckcertificate": True,
"progress_hooks": [bar.hook],
} }
if quality != "audio" and not format_override: if quality != "audio" and not format_override:
opts["merge_output_format"] = "mp4" opts["merge_output_format"] = "mp4"
+107
View File
@@ -0,0 +1,107 @@
"""Minimal terminal progress bar for yt-dlp downloads (stdlib only)."""
from __future__ import annotations
import sys
import time
from pathlib import Path
def _fmt_bytes(n: float) -> str:
n = float(n or 0)
for unit in ("B", "KiB", "MiB", "GiB", "TiB"):
if n < 1024 or unit == "TiB":
return f"{int(n)} B" if unit == "B" else f"{n:.1f} {unit}"
n /= 1024
return "0 B"
def _fmt_eta(sec: float) -> str:
sec = int(sec or 0)
if sec < 60:
return f"{sec:02d}s"
m, s = divmod(sec, 60)
if m < 60:
return f"{m:02d}:{s:02d}"
h, m = divmod(m, 60)
return f"{h:02d}:{m:02d}:{s:02d}"
class DownloadProgress:
"""Draws a one-line progress bar from yt-dlp `progress_hooks` data.
On a TTY the bar redraws in place; when piped to a file it falls back to
a sparse log line every few seconds (like yt-dlp's own ``[download]``).
A bestvideo+bestaudio download reports each stream as its own file, so the
bar resets per stream and shows the stream's filename.
"""
_BAR = 24
def __init__(self, stream=None):
self._stream = stream or sys.stderr
self._tty = self._stream.isatty()
self._file: str | None = None
self._last = 0.0
self._last_log = 0.0
self._line_open = False
self._min_interval = 0.05
def hook(self, data: dict) -> None:
status = data.get("status")
if status == "downloading":
self._on_downloading(data)
elif status == "finished":
self._on_finished()
def _on_downloading(self, data: dict) -> None:
fname = Path(data.get("filename", "")).name
now = time.monotonic()
if fname != self._file:
if self._tty and self._line_open:
self._stream.write("\n")
self._file = fname
self._line_open = False
if not self._tty:
if now - self._last_log < 5.0:
return
self._last_log = now
self._stream.write(self._text(data) + "\n")
self._stream.flush()
return
if now - self._last < self._min_interval:
return
self._last = now
self._stream.write("\r\x1b[K" + self._text(data))
self._stream.flush()
self._line_open = True
def _on_finished(self) -> None:
if self._tty and self._line_open:
self._stream.write("\n")
self._stream.flush()
self._file = None
self._line_open = False
def _text(self, data: dict) -> str:
done = data.get("downloaded_bytes") or 0
total = data.get("total_bytes") or data.get("total_bytes_estimate") or 0
speed = data.get("speed")
eta = data.get("eta")
pct = (done / total * 100) if total else None
parts = []
if pct is not None:
filled = int(self._BAR * pct / 100)
bar = "\u2588" * filled + "\u2591" * (self._BAR - filled)
parts.append(f"[{bar}] {pct:5.1f}%")
else:
parts.append(f"[{'\u2591' * self._BAR}] {_fmt_bytes(done)}")
if total:
parts.append(f"{_fmt_bytes(done)}/{_fmt_bytes(total)}")
if speed:
parts.append(f"{_fmt_bytes(speed)}/s")
if eta is not None:
parts.append(f"ETA {_fmt_eta(eta)}")
if self._file:
parts.append(self._file)
return " ".join(parts)
+47 -2
View File
@@ -137,6 +137,31 @@ def is_up(cfg) -> bool:
return rc == 0 and "no such device" not in out.lower() return rc == 0 and "no such device" not in out.lower()
def handshake_age(cfg) -> int | None:
"""Seconds since the last WireGuard handshake, or None if none/unknown."""
ns, iface = netns_name(cfg), iface_name(cfg)
rc, out = _sudo_out(["ip", "netns", "exec", ns, "wg", "show", iface])
if rc != 0:
return None
for line in out.splitlines():
m = re.search(r"latest handshake:\s*(.+)", line)
if not m:
continue
text = m.group(1).strip().lower()
if text == "no handshake":
return None
m2 = re.match(r"(\d+)\s+seconds? ago", text)
if m2:
return int(m2.group(1))
return None
def has_handshake(cfg, max_age: int = 180) -> bool:
"""True when the tunnel has a recent handshake (i.e. is actually usable)."""
age = handshake_age(cfg)
return age is not None and age <= max_age
def up(cfg) -> None: def up(cfg) -> None:
ns, iface = netns_name(cfg), iface_name(cfg) ns, iface = netns_name(cfg), iface_name(cfg)
host_link = _host_link(iface) host_link = _host_link(iface)
@@ -195,12 +220,29 @@ def down(cfg) -> None:
_sudo(["rm", "-rf", f"/etc/netns/{ns}"]) _sudo(["rm", "-rf", f"/etc/netns/{ns}"])
def wait_for_handshake(cfg, timeout: float = 15.0) -> bool:
import time
deadline = time.monotonic() + timeout
while time.monotonic() < deadline:
if has_handshake(cfg):
return True
time.sleep(0.5)
return has_handshake(cfg)
def ensure_up(cfg) -> None: def ensure_up(cfg) -> None:
if is_up(cfg): if has_handshake(cfg):
return return
log.info("Mullvad tunnel is down — bringing it up") log.info("Mullvad tunnel is down or not handshaken — bringing it up")
down(cfg) # converge from any stale half-configured state down(cfg) # converge from any stale half-configured state
up(cfg) up(cfg)
if not wait_for_handshake(cfg):
down(cfg)
raise RuntimeError(
"WireGuard tunnel came up but no handshake was established — "
"check the endpoint/reachability and that the tunnel is actually "
"active. Refusing to download over a dead tunnel."
)
def status_text(cfg) -> str: def status_text(cfg) -> str:
@@ -210,6 +252,9 @@ def status_text(cfg) -> str:
return "tunnel: DOWN (namespace or WireGuard interface missing)" return "tunnel: DOWN (namespace or WireGuard interface missing)"
lines = ["tunnel: UP", f"namespace : {ns}"] lines = ["tunnel: UP", f"namespace : {ns}"]
lines.extend(line for line in out.splitlines()) lines.extend(line for line in out.splitlines())
age = handshake_age(cfg)
lines.append("handshake : " + ("usable" if has_handshake(cfg)
else ("none" if age is None else f"{age}s ago (stale)")))
_, route = _sudo_out(["ip", "netns", "exec", ns, "ip", "route", "show", "default"]) _, route = _sudo_out(["ip", "netns", "exec", ns, "ip", "route", "show", "default"])
lines.append("default : " + (route or "(none)")) lines.append("default : " + (route or "(none)"))
_, dns = _sudo_out(["cat", f"/etc/netns/{ns}/resolv.conf"]) _, dns = _sudo_out(["cat", f"/etc/netns/{ns}/resolv.conf"])