diff --git a/yt_downloader/cli.py b/yt_downloader/cli.py index c5f8ac9..1268523 100644 --- a/yt_downloader/cli.py +++ b/yt_downloader/cli.py @@ -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") diff --git a/yt_downloader/downloader.py b/yt_downloader/downloader.py index 60c0e78..fd718f0 100644 --- a/yt_downloader/downloader.py +++ b/yt_downloader/downloader.py @@ -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" diff --git a/yt_downloader/progress.py b/yt_downloader/progress.py new file mode 100644 index 0000000..86e856f --- /dev/null +++ b/yt_downloader/progress.py @@ -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) \ No newline at end of file diff --git a/yt_downloader/vpn.py b/yt_downloader/vpn.py index 84bf0d5..4f8cde8 100644 --- a/yt_downloader/vpn.py +++ b/yt_downloader/vpn.py @@ -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"])