foxygit / ytmdl Log in
commits tags

/src/ytmdl/app.py · 10.82 KB

raw
from __future__ import annotations

import queue as sync_queue
from pathlib import Path
from typing import Iterable

from textual.app import App, ComposeResult, SystemCommand
from textual.binding import Binding
from textual.containers import Horizontal, Vertical
from textual.screen import Screen
from textual.widgets import Button, DataTable, Footer, Header, Input, Label, Switch
from textual.worker import get_current_worker

from . import downloader
from .queue import STATUS_GLYPH, QueueItem, Status
from .screens import DirectoryPickerScreen, NerdCommandPalette

DEFAULT_OUTPUT_DIR = Path.home() / "Music" / "ytmdl"
CONCURRENT_DOWNLOADS = 2


def _format_speed(bytes_per_sec: float | None) -> str:
    if not bytes_per_sec:
        return ""
    value = float(bytes_per_sec)
    for unit in ("B/s", "KiB/s", "MiB/s", "GiB/s"):
        if value < 1024:
            return f"{value:.1f}{unit}"
        value /= 1024
    return f"{value:.1f}TiB/s"


def _format_eta(seconds: float | None) -> str:
    if seconds is None:
        return ""
    seconds = int(seconds)
    minutes, secs = divmod(seconds, 60)
    hours, minutes = divmod(minutes, 60)
    if hours:
        return f"{hours:d}:{minutes:02d}:{secs:02d}"
    return f"{minutes:d}:{secs:02d}"


class YtmdlApp(App):
    CSS = """
    #url-input {
        margin: 1 2 0 2;
    }

    #output-dir-row {
        height: auto;
        margin: 0 2 1 2;
    }

    #output-dir-label {
        width: auto;
        color: $text-muted;
        padding: 1 1 0 0;
    }

    #output-dir-button {
        min-width: 0;
        height: 1;
        margin-top: 1;
        border: none;
        background: transparent;
        color: $text;
        text-style: underline;
    }

    #output-dir-button:hover {
        color: $accent;
    }

    #playlist-folder-row {
        height: auto;
        margin: 0 2 1 2;
    }

    #playlist-folder-row Label {
        color: $text-muted;
        padding: 1 0 0 1;
    }

    DataTable {
        margin: 0 2 1 2;
    }
    """

    BINDINGS = [
        Binding("q", "quit", "Quit"),
        Binding("d,delete", "remove_selected", "Remove"),
        Binding("r", "retry_selected", "Retry"),
        Binding("ctrl+o", "edit_output_dir", "Change save dir", priority=True),
        Binding("ctrl+n", "clear_finished", "Clear finished", priority=True),
    ]

    def __init__(
        self,
        output_dir: Path = DEFAULT_OUTPUT_DIR,
        audio_format: str = "mp3",
    ) -> None:
        super().__init__()
        self.output_dir = output_dir
        self.audio_format = audio_format
        self.playlist_folders = True
        self.items: dict[str, QueueItem] = {}
        self.download_queue: sync_queue.Queue[QueueItem | None] = sync_queue.Queue()

    def compose(self) -> ComposeResult:
        yield Header()
        yield Vertical(
            Input(
                placeholder="Paste a YouTube / YouTube Music video, playlist, or album URL and press Enter",
                id="url-input",
            ),
            Horizontal(
                Label("Saving to:", id="output-dir-label"),
                Button(str(self.output_dir), id="output-dir-button"),
                id="output-dir-row",
            ),
            Horizontal(
                Switch(value=self.playlist_folders, id="playlist-folder-switch"),
                Label("Create a folder per playlist/album (named after it)"),
                id="playlist-folder-row",
            ),
        )
        table = DataTable(id="queue-table")
        table.cursor_type = "row"
        table.zebra_stripes = True
        yield table
        yield Footer()

    def on_mount(self) -> None:
        table = self.query_one(DataTable)
        self.columns = table.add_columns("Status", "Title", "Folder", "Progress", "Speed", "ETA")
        self.query_one("#url-input", Input).focus()
        for _ in range(CONCURRENT_DOWNLOADS):
            self.run_worker(self._download_worker, thread=True, exclusive=False)

    def on_input_submitted(self, event: Input.Submitted) -> None:
        url = event.value.strip()
        if not url:
            return
        event.input.value = ""
        self.notify(f"Resolving {url}...", timeout=3)
        self.run_worker(lambda: self._resolve_worker(url), thread=True, exclusive=False)

    def on_button_pressed(self, event: Button.Pressed) -> None:
        if event.button.id == "output-dir-button":
            self.action_edit_output_dir()

    def on_switch_changed(self, event: Switch.Changed) -> None:
        if event.switch.id == "playlist-folder-switch":
            self.playlist_folders = event.value

    def action_command_palette(self) -> None:
        if self.use_command_palette and not NerdCommandPalette.is_open(self):
            self.push_screen(NerdCommandPalette(id="--command-palette"))

    def get_system_commands(self, screen: Screen) -> Iterable[SystemCommand]:
        yield from super().get_system_commands(screen)
        yield SystemCommand(
            "Clear finished",
            "Remove completed/errored entries from the queue to start a new round",
            self.action_clear_finished,
        )

    def action_edit_output_dir(self) -> None:
        self.push_screen(DirectoryPickerScreen(self.output_dir), self._on_output_dir_picked)

    def _on_output_dir_picked(self, path: Path | None) -> None:
        if path is None:
            return
        self.output_dir = path
        self.query_one("#output-dir-button", Button).label = str(path)
        self.notify(f"Now saving to {path}", timeout=3)

    # --- background workers -------------------------------------------------

    def _resolve_worker(self, url: str) -> None:
        try:
            result = downloader.resolve(url)
        except downloader.ResolveError as exc:
            self.call_from_thread(self.notify, f"Failed to resolve: {exc}", severity="error", timeout=6)
            return

        subdir = ""
        if result.playlist_title and self.playlist_folders:
            subdir = downloader.sanitize_dirname(result.playlist_title)

        for entry in result.entries:
            item = QueueItem(url=entry["url"], title=entry["title"], subdir=subdir)
            self.call_from_thread(self._enqueue_item, item)

    def _download_worker(self) -> None:
        worker = get_current_worker()
        while not worker.is_cancelled:
            try:
                item = self.download_queue.get(timeout=0.5)
            except sync_queue.Empty:
                continue
            if item is None:
                break
            self._process_item(item)

    def _process_item(self, item: QueueItem) -> None:
        if item.status is Status.DONE:
            return
        self.call_from_thread(self._update_item, item.id, status=Status.DOWNLOADING, error="")

        def hook(d: dict) -> None:
            if d.get("status") == "downloading":
                total = d.get("total_bytes") or d.get("total_bytes_estimate")
                downloaded = d.get("downloaded_bytes") or 0
                percent = (downloaded / total * 100) if total else 0.0
                speed = _format_speed(d.get("speed"))
                eta = _format_eta(d.get("eta"))
                self.call_from_thread(
                    self._update_item, item.id, percent=percent, speed=speed, eta=eta
                )

        target_dir = self.output_dir / item.subdir if item.subdir else self.output_dir
        try:
            downloader.download(
                item.url,
                target_dir,
                audio_format=self.audio_format,
                progress_hook=hook,
            )
        except downloader.DownloadError as exc:
            self.call_from_thread(
                self._update_item, item.id, status=Status.ERROR, error=str(exc), speed="", eta=""
            )
        else:
            self.call_from_thread(
                self._update_item, item.id, status=Status.DONE, percent=100.0, speed="", eta=""
            )

    # --- UI-thread state mutation -------------------------------------------

    def _enqueue_item(self, item: QueueItem) -> None:
        self.items[item.id] = item
        table = self.query_one(DataTable)
        table.add_row(
            STATUS_GLYPH[item.status],
            item.title,
            item.subdir,
            "0%",
            "",
            "",
            key=item.id,
        )
        self.download_queue.put(item)

    def _update_item(self, item_id: str, **changes) -> None:
        item = self.items.get(item_id)
        if item is None:
            return
        for field, value in changes.items():
            setattr(item, field, value)

        table = self.query_one(DataTable)
        if item_id not in table.rows:
            return
        status_col, title_col, folder_col, progress_col, speed_col, eta_col = self.columns
        table.update_cell(item_id, status_col, STATUS_GLYPH[item.status])
        if item.status is Status.ERROR:
            table.update_cell(item_id, progress_col, item.error[:40])
        else:
            table.update_cell(item_id, progress_col, f"{item.percent:.0f}%")
        table.update_cell(item_id, speed_col, item.speed)
        table.update_cell(item_id, eta_col, item.eta)

    # --- actions -------------------------------------------------------------

    def _selected_item_id(self) -> str | None:
        table = self.query_one(DataTable)
        if table.row_count == 0:
            return None
        try:
            row_key, _ = table.coordinate_to_cell_key(table.cursor_coordinate)
        except Exception:
            return None
        return row_key.value

    def action_clear_finished(self) -> None:
        finished_ids = [
            item.id for item in self.items.values() if item.status in (Status.DONE, Status.ERROR)
        ]
        table = self.query_one(DataTable)
        for item_id in finished_ids:
            table.remove_row(item_id)
            del self.items[item_id]
        if finished_ids:
            self.notify(f"Cleared {len(finished_ids)} finished item(s)", timeout=3)
        else:
            self.notify("Nothing to clear", timeout=2)

    def action_remove_selected(self) -> None:
        item_id = self._selected_item_id()
        if item_id is None:
            return
        item = self.items.get(item_id)
        if item is None or item.status is Status.DOWNLOADING:
            return
        table = self.query_one(DataTable)
        table.remove_row(item_id)
        del self.items[item_id]

    def action_retry_selected(self) -> None:
        item_id = self._selected_item_id()
        if item_id is None:
            return
        item = self.items.get(item_id)
        if item is None or item.status is not Status.ERROR:
            return
        item.status = Status.PENDING
        item.percent = 0.0
        item.error = ""
        self._update_item(item_id, status=Status.PENDING)
        self.download_queue.put(item)

    def action_quit(self) -> None:
        for _ in range(CONCURRENT_DOWNLOADS):
            self.download_queue.put(None)
        self.exit()