diff --git a/tui/screens/thread.py b/tui/screens/thread.py index 251f691..bb5e576 100644 --- a/tui/screens/thread.py +++ b/tui/screens/thread.py @@ -1,7 +1,5 @@ import asyncio -from pathlib import Path -from platformdirs import user_downloads_dir from textual import work from textual.app import ComposeResult from textual.binding import Binding @@ -19,7 +17,13 @@ from core.records import ( from core.slingshot import get_record, resolve_identity from core.util import attachment_cid, blob_url from tui.screens.compose import ComposeReplyScreen -from tui.util import ban_user, hide_post, require_session, require_sysop +from tui.util import ( + ban_user, + download_blob, + hide_post, + require_session, + require_sysop, +) from tui.widgets.breadcrumb import Breadcrumb from tui.widgets.post import Post @@ -302,9 +306,6 @@ class ThreadScreen(Screen): @work(exclusive=True) async def _do_save(self, post: Post) -> None: - downloads = Path(user_downloads_dir()) - downloads.mkdir(parents=True, exist_ok=True) - client = self.app.http_client for attachment in post.attachments: name = attachment.get("name", "file") @@ -314,16 +315,7 @@ class ThreadScreen(Screen): url = blob_url(post.author_pds, post.author_did, cid) try: - resp = await client.get(url) - resp.raise_for_status() - path = downloads / name - if path.exists(): - stem, suffix = path.stem, path.suffix - i = 1 - while path.exists(): - path = downloads / f"{stem}_{i}{suffix}" - i += 1 - path.write_bytes(resp.content) + path = await download_blob(client, url, name, downloads) self.notify(f"Saved to {path}") except Exception: self.notify(f"Failed to download {name}.", severity="error") diff --git a/tui/util.py b/tui/util.py index 55dfe30..412c969 100644 --- a/tui/util.py +++ b/tui/util.py @@ -1,11 +1,42 @@ """TUI utilities.""" +from pathlib import Path + +import httpx +from platformdirs import user_downloads_dir + from core.auth.session import SessionStore from core.models import AuthError, BBS from core.records import create_ban_record, create_hidden_record from core.resolver import invalidate_bbs_cache +def unique_path(path: Path) -> Path: + """Return path, or path with a `_N` suffix if it already exists.""" + if not path.exists(): + return path + stem, suffix = path.stem, path.suffix + counter = 1 + while True: + candidate = path.parent / f"{stem}_{counter}{suffix}" + if not candidate.exists(): + return candidate + counter += 1 + + +async def download_blob( + client: httpx.AsyncClient, url: str, filename: str +) -> Path: + """Fetch a blob URL and save it to the user's Downloads folder.""" + downloads = Path(user_downloads_dir()) + downloads.mkdir(parents=True, exist_ok=True) + resp = await client.get(url) + resp.raise_for_status() + path = unique_path(downloads / filename) + path.write_bytes(resp.content) + return path + + def require_session(screen) -> dict | None: """Return the user session if logged in, else notify and return None.""" session = screen.app.user_session diff --git a/tui/widgets/post.py b/tui/widgets/post.py index 50e6408..cf45ab1 100644 --- a/tui/widgets/post.py +++ b/tui/widgets/post.py @@ -2,6 +2,7 @@ import re import webbrowser from urllib.parse import unquote +from textual import work from textual.app import ComposeResult from textual.widget import Widget from textual.widgets import Markdown, Static @@ -12,6 +13,7 @@ from core.util import ( blob_url, format_datetime_local as format_datetime, ) +from tui.util import download_blob ATTACHMENT_LINK_RE = re.compile(r"!?\[([^\]]*)\]\(attachment:([^)\s]+)\)") ATTACHMENT_REF_RE = re.compile(r"attachment:([^)\s>\"']+)") @@ -72,15 +74,40 @@ class AttachmentLink(Static, can_focus=True): } """ - def __init__(self, display: str, url: str, **kwargs) -> None: + def __init__( + self, + display: str, + url: str, + filename: str | None = None, + **kwargs, + ) -> None: super().__init__(display, markup=False, **kwargs) self._url = url + self._filename = filename def on_click(self) -> None: - webbrowser.open(self._url) + self._activate() def key_enter(self) -> None: - webbrowser.open(self._url) + self._activate() + + def _activate(self) -> None: + if self._filename: + self._save() + else: + webbrowser.open(self._url) + + @work + async def _save(self) -> None: + try: + path = await download_blob( + self.app.http_client, self._url, self._filename + ) + self.notify(f"Saved to {path}") + except Exception: + self.notify( + f"Failed to download {self._filename}.", severity="error" + ) class Post(Widget, can_focus=True): @@ -163,16 +190,36 @@ class Post(Widget, can_focus=True): self._body, self.attachments, self.author_pds, self.author_did ) yield Markdown(body, classes="post-body") + + downloadable, undownloadable = self._partition_attachments() + filename_by_url = {url: name for name, url in downloadable} + for number, (label, url) in enumerate(extract_body_links(body), 1): - yield AttachmentLink(f"[{number}] {label}", url) + yield AttachmentLink( + f"[{number}] {label}", url, filename=filename_by_url.get(url) + ) + embedded = referenced_attachment_names(self._body) + for name, url in downloadable: + if name not in embedded: + yield AttachmentLink(f"[{name}]", url, filename=name) + for name in undownloadable: + if name not in embedded: + yield Static(f"[{name}]", classes="post-attachment", markup=False) + + def _partition_attachments( + self, + ) -> tuple[list[tuple[str, str]], list[str]]: + """Split attachments into (name, url) we can fetch and names we can't.""" + downloadable: list[tuple[str, str]] = [] + undownloadable: list[str] = [] for attachment in self.attachments: name = attachment.get("name", "file") - if name in embedded: - continue cid = attachment_cid(attachment) if cid and self.author_pds and self.author_did: - url = blob_url(self.author_pds, self.author_did, cid) - yield AttachmentLink(f"[{name}]", url) + downloadable.append( + (name, blob_url(self.author_pds, self.author_did, cid)) + ) else: - yield Static(f"[{name}]", classes="post-attachment", markup=False) + undownloadable.append(name) + return downloadable, undownloadable