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 . 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")
|
||||||
|
|||||||
@@ -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"
|
||||||
|
|||||||
@@ -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()
|
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"])
|
||||||
|
|||||||
Reference in New Issue
Block a user