vpn: require an active WireGuard handshake before downloading (refuse a dead tunnel)
This commit is contained in:
+10
-1
@@ -11,6 +11,7 @@ from pathlib import Path
|
||||
from . import __version__, vpn
|
||||
from .config import Config, QUALITY_PRESETS, default_config_path, load_config, write_default_config
|
||||
from .downloader import cookie_opts, download_video
|
||||
from .progress import DownloadProgress
|
||||
|
||||
|
||||
def log_setup(verbose: bool) -> None:
|
||||
@@ -74,7 +75,8 @@ def cmd_download(args: argparse.Namespace) -> int:
|
||||
extra = cookie_opts(cfg)
|
||||
|
||||
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:
|
||||
logging.error("download failed: %s", args.url)
|
||||
return 1
|
||||
@@ -91,6 +93,13 @@ def cmd_vpn(args: argparse.Namespace) -> int:
|
||||
else:
|
||||
vpn.up(cfg)
|
||||
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":
|
||||
vpn.down(cfg)
|
||||
print("tunnel down")
|
||||
|
||||
@@ -8,6 +8,8 @@ from typing import Optional
|
||||
|
||||
import yt_dlp
|
||||
|
||||
from .progress import DownloadProgress
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
_BASE_OPTS = {
|
||||
@@ -16,6 +18,7 @@ _BASE_OPTS = {
|
||||
"noplaylist": True,
|
||||
"ignoreerrors": 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
|
||||
@@ -54,17 +57,20 @@ def download_video(
|
||||
quality: str = "best",
|
||||
format_override: Optional[str] = None,
|
||||
extra_opts: dict | None = None,
|
||||
progress: DownloadProgress | None = None,
|
||||
) -> Optional[Path]:
|
||||
"""Download a single video into out_dir. Returns the resulting file path."""
|
||||
out_dir = out_dir.expanduser()
|
||||
out_dir.mkdir(parents=True, exist_ok=True)
|
||||
fmt = _select_format(quality, format_override)
|
||||
bar = progress or DownloadProgress()
|
||||
opts = {
|
||||
**_BASE_OPTS,
|
||||
**(extra_opts or {}),
|
||||
"format": fmt,
|
||||
"outtmpl": str(out_dir / "%(title).150B [%(id)s].%(ext)s"),
|
||||
"nocheckcertificate": True,
|
||||
"progress_hooks": [bar.hook],
|
||||
}
|
||||
if quality != "audio" and not format_override:
|
||||
opts["merge_output_format"] = "mp4"
|
||||
|
||||
@@ -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
@@ -137,6 +137,31 @@ def is_up(cfg) -> bool:
|
||||
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:
|
||||
ns, iface = netns_name(cfg), iface_name(cfg)
|
||||
host_link = _host_link(iface)
|
||||
@@ -195,12 +220,29 @@ def down(cfg) -> None:
|
||||
_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:
|
||||
if is_up(cfg):
|
||||
if has_handshake(cfg):
|
||||
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
|
||||
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:
|
||||
@@ -210,6 +252,9 @@ def status_text(cfg) -> str:
|
||||
return "tunnel: DOWN (namespace or WireGuard interface missing)"
|
||||
lines = ["tunnel: UP", f"namespace : {ns}"]
|
||||
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"])
|
||||
lines.append("default : " + (route or "(none)"))
|
||||
_, dns = _sudo_out(["cat", f"/etc/netns/{ns}/resolv.conf"])
|
||||
|
||||
Reference in New Issue
Block a user