big upgrade UwU
This commit is contained in:
@@ -0,0 +1,359 @@
|
||||
"""The async sticker-pack downloader."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
from collections.abc import Awaitable, Callable, Iterable, Sequence
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timezone
|
||||
from pathlib import Path
|
||||
from typing import Any, TypeVar
|
||||
|
||||
from telegram.error import BadRequest, TelegramError
|
||||
|
||||
from sticker_downloader.config import (
|
||||
ALL_FILE_TYPES,
|
||||
ANIMATED,
|
||||
METADATA_FILE_NAME,
|
||||
PACK_THUMBNAIL_STEM,
|
||||
STATIC,
|
||||
THUMBNAIL,
|
||||
VIDEO,
|
||||
DownloadConfig,
|
||||
)
|
||||
from sticker_downloader.errors import InvalidPackReference, StickerDownloaderError
|
||||
from sticker_downloader.fetcher import Fetcher
|
||||
from sticker_downloader.progress import NullProgress, ProgressReporter
|
||||
from sticker_downloader.results import FileOutcome, FileStatus, PackResult
|
||||
from sticker_downloader.retry import with_retries
|
||||
from sticker_downloader.urls import parse_pack_name
|
||||
|
||||
T = TypeVar("T")
|
||||
|
||||
_UNSAFE_FILENAME_CHARS = re.compile(r'[<>:"/\\|?*\x00-\x1f]')
|
||||
_PART_SUFFIX = ".part"
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _FileJob:
|
||||
"""One file to put on disk."""
|
||||
|
||||
file_id: str
|
||||
stem: Path
|
||||
#: ``None`` means "derive the extension from Telegram's own file path",
|
||||
#: which is how pack thumbnails are handled (they can be webp, tgs or webm).
|
||||
suffix: str | None
|
||||
|
||||
|
||||
class StickerDownloader:
|
||||
"""Downloads sticker packs concurrently into per-pack directories."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
bot: Any,
|
||||
fetcher: Fetcher,
|
||||
config: DownloadConfig | None = None,
|
||||
progress: ProgressReporter | None = None,
|
||||
*,
|
||||
sleep: Callable[[float], Awaitable[None]] = asyncio.sleep,
|
||||
now: Callable[[], datetime] | None = None,
|
||||
) -> None:
|
||||
self._bot = bot
|
||||
self._fetcher = fetcher
|
||||
self._config = config or DownloadConfig()
|
||||
self._progress: ProgressReporter = progress or NullProgress()
|
||||
self._sleep = sleep
|
||||
self._now = now or (lambda: datetime.now(timezone.utc))
|
||||
|
||||
@property
|
||||
def config(self) -> DownloadConfig:
|
||||
return self._config
|
||||
|
||||
# -- public API -------------------------------------------------------
|
||||
|
||||
async def download_all(self, references: Iterable[str]) -> list[PackResult]:
|
||||
"""Download every reference in turn.
|
||||
|
||||
Packs run one after another so progress stays readable; the files inside a
|
||||
pack are downloaded concurrently.
|
||||
"""
|
||||
results: list[PackResult] = []
|
||||
for reference in references:
|
||||
results.append(await self.download_pack(reference))
|
||||
self._progress.run_finished(results)
|
||||
return results
|
||||
|
||||
async def download_pack(self, reference: str) -> PackResult:
|
||||
"""Download a single pack, never raising for an expected failure."""
|
||||
try:
|
||||
pack_name = parse_pack_name(reference)
|
||||
except InvalidPackReference as error:
|
||||
result = PackResult(reference=reference, error=str(error))
|
||||
self._progress.pack_finished(result)
|
||||
return result
|
||||
|
||||
try:
|
||||
sticker_set = await self._call(lambda: self._bot.get_sticker_set(pack_name))
|
||||
except Exception as error:
|
||||
result = PackResult(
|
||||
reference=reference, name=pack_name, error=_describe(error)
|
||||
)
|
||||
self._progress.pack_finished(result)
|
||||
return result
|
||||
|
||||
stickers: Sequence[Any] = list(getattr(sticker_set, "stickers", None) or [])
|
||||
name = getattr(sticker_set, "name", None) or pack_name
|
||||
directory = self._config.output_dir / name
|
||||
|
||||
jobs = self._plan_jobs(sticker_set, stickers, directory)
|
||||
self._progress.pack_started(reference, name, len(stickers), len(jobs))
|
||||
|
||||
if not self._config.dry_run:
|
||||
await asyncio.to_thread(directory.mkdir, parents=True, exist_ok=True)
|
||||
|
||||
semaphore = asyncio.Semaphore(self._config.concurrency)
|
||||
outcomes = list(
|
||||
await asyncio.gather(*(self._run_job(job, semaphore) for job in jobs))
|
||||
)
|
||||
|
||||
result = PackResult(
|
||||
reference=reference,
|
||||
name=name,
|
||||
title=getattr(sticker_set, "title", None),
|
||||
directory=directory,
|
||||
total_stickers=len(stickers),
|
||||
outcomes=outcomes,
|
||||
)
|
||||
|
||||
if self._config.write_metadata and not self._config.dry_run:
|
||||
try:
|
||||
await self._write_metadata(sticker_set, stickers, directory, result)
|
||||
except OSError as error:
|
||||
self._progress.warn(f"could not write {METADATA_FILE_NAME}: {error}")
|
||||
|
||||
self._progress.pack_finished(result)
|
||||
return result
|
||||
|
||||
# -- planning ---------------------------------------------------------
|
||||
|
||||
def _plan_jobs(
|
||||
self, sticker_set: Any, stickers: Sequence[Any], directory: Path
|
||||
) -> list[_FileJob]:
|
||||
jobs: list[_FileJob] = []
|
||||
|
||||
set_thumbnail = getattr(sticker_set, "thumbnail", None)
|
||||
if self._config.pack_thumbnail and set_thumbnail is not None:
|
||||
jobs.append(
|
||||
_FileJob(
|
||||
file_id=set_thumbnail.file_id,
|
||||
stem=directory / PACK_THUMBNAIL_STEM,
|
||||
suffix=None,
|
||||
)
|
||||
)
|
||||
|
||||
for index, sticker in enumerate(stickers, start=1):
|
||||
base = self._sticker_stem(index, sticker)
|
||||
|
||||
primary = primary_file_type(sticker)
|
||||
if self._config.wants(primary):
|
||||
jobs.append(
|
||||
_FileJob(
|
||||
file_id=sticker.file_id,
|
||||
stem=directory / base,
|
||||
suffix=f".{primary}",
|
||||
)
|
||||
)
|
||||
|
||||
thumbnail = getattr(sticker, "thumbnail", None)
|
||||
if self._config.wants(THUMBNAIL) and thumbnail is not None:
|
||||
jobs.append(
|
||||
_FileJob(
|
||||
file_id=thumbnail.file_id,
|
||||
stem=directory / f"{base}_thumb",
|
||||
suffix=f".{THUMBNAIL}",
|
||||
)
|
||||
)
|
||||
|
||||
return jobs
|
||||
|
||||
def _sticker_stem(self, index: int, sticker: Any) -> str:
|
||||
base = f"{index:03d}"
|
||||
if not self._config.emoji_names:
|
||||
return base
|
||||
emoji = sanitize_filename_part(getattr(sticker, "emoji", None) or "")
|
||||
return f"{base}_{emoji}" if emoji else base
|
||||
|
||||
# -- execution --------------------------------------------------------
|
||||
|
||||
async def _run_job(self, job: _FileJob, semaphore: asyncio.Semaphore) -> FileOutcome:
|
||||
async with semaphore:
|
||||
outcome = await self._process_job(job)
|
||||
self._progress.file_finished(outcome)
|
||||
return outcome
|
||||
|
||||
async def _process_job(self, job: _FileJob) -> FileOutcome:
|
||||
planned_path = self._planned_path(job)
|
||||
try:
|
||||
existing = None if self._config.overwrite else self._existing_path(job)
|
||||
if existing is not None:
|
||||
return FileOutcome(existing, FileStatus.SKIPPED_EXISTING)
|
||||
|
||||
if self._config.dry_run:
|
||||
return FileOutcome(planned_path, FileStatus.PLANNED)
|
||||
|
||||
file = await self._call(lambda: self._bot.get_file(job.file_id))
|
||||
remote_path = getattr(file, "file_path", None)
|
||||
if not remote_path:
|
||||
raise StickerDownloaderError("Telegram returned no file path")
|
||||
|
||||
destination = self._resolve_path(job, remote_path)
|
||||
data = await self._call(lambda: self._fetcher.fetch(remote_path))
|
||||
await asyncio.to_thread(_write_atomically, destination, data)
|
||||
return FileOutcome(destination, FileStatus.DOWNLOADED, size=len(data))
|
||||
except Exception as error:
|
||||
return FileOutcome(planned_path, FileStatus.FAILED, error=_describe(error))
|
||||
|
||||
async def _call(self, operation: Callable[[], Awaitable[T]]) -> T:
|
||||
"""Run one network operation with this run's retry policy."""
|
||||
return await with_retries(
|
||||
operation,
|
||||
attempts=self._config.retries,
|
||||
on_retry=self._on_retry,
|
||||
sleep=self._sleep,
|
||||
)
|
||||
|
||||
def _on_retry(self, error: BaseException, attempt: int, delay: float) -> None:
|
||||
self._progress.warn(
|
||||
f"{_describe(error)} — retrying in {delay:.1f}s (attempt {attempt + 1})"
|
||||
)
|
||||
|
||||
# -- paths ------------------------------------------------------------
|
||||
|
||||
def _planned_path(self, job: _FileJob) -> Path:
|
||||
"""Best guess at the destination, used for dry runs and error messages."""
|
||||
return self._resolve_path(job, f"unknown.{STATIC}")
|
||||
|
||||
def _resolve_path(self, job: _FileJob, remote_path: str) -> Path:
|
||||
suffix = job.suffix if job.suffix is not None else _suffix_from_url(remote_path)
|
||||
return job.stem.parent / f"{job.stem.name}{suffix}"
|
||||
|
||||
def _existing_path(self, job: _FileJob) -> Path | None:
|
||||
"""The already-downloaded file for ``job``, if there is one.
|
||||
|
||||
Checking before calling ``get_file`` means resuming a large batch costs no
|
||||
API calls for the parts that are already on disk.
|
||||
"""
|
||||
if job.suffix is not None:
|
||||
candidate = job.stem.parent / f"{job.stem.name}{job.suffix}"
|
||||
return candidate if candidate.is_file() else None
|
||||
for file_type in ALL_FILE_TYPES:
|
||||
candidate = job.stem.parent / f"{job.stem.name}.{file_type}"
|
||||
if candidate.is_file():
|
||||
return candidate
|
||||
return None
|
||||
|
||||
# -- metadata ---------------------------------------------------------
|
||||
|
||||
async def _write_metadata(
|
||||
self,
|
||||
sticker_set: Any,
|
||||
stickers: Sequence[Any],
|
||||
directory: Path,
|
||||
result: PackResult,
|
||||
) -> None:
|
||||
present = {
|
||||
outcome.path
|
||||
for outcome in result.outcomes
|
||||
if outcome.status in (FileStatus.DOWNLOADED, FileStatus.SKIPPED_EXISTING)
|
||||
}
|
||||
|
||||
entries = []
|
||||
for index, sticker in enumerate(stickers, start=1):
|
||||
base = self._sticker_stem(index, sticker)
|
||||
files = sorted(
|
||||
path.name
|
||||
for path in present
|
||||
if path.stem == base or path.stem == f"{base}_thumb"
|
||||
)
|
||||
entries.append(
|
||||
{
|
||||
"index": index,
|
||||
"emoji": getattr(sticker, "emoji", None),
|
||||
"file_unique_id": getattr(sticker, "file_unique_id", None),
|
||||
"type": _enum_value(getattr(sticker, "type", None)),
|
||||
"width": getattr(sticker, "width", None),
|
||||
"height": getattr(sticker, "height", None),
|
||||
"is_animated": bool(getattr(sticker, "is_animated", False)),
|
||||
"is_video": bool(getattr(sticker, "is_video", False)),
|
||||
"files": files,
|
||||
}
|
||||
)
|
||||
|
||||
from sticker_downloader import __version__
|
||||
|
||||
payload = {
|
||||
"pack": {
|
||||
"reference": result.reference,
|
||||
"name": result.name,
|
||||
"title": result.title,
|
||||
"sticker_type": _enum_value(getattr(sticker_set, "sticker_type", None)),
|
||||
},
|
||||
"generated_at": self._now().isoformat(),
|
||||
"downloader_version": __version__,
|
||||
"requested_file_types": self._config.sorted_file_types,
|
||||
"sticker_count": len(stickers),
|
||||
"stickers": entries,
|
||||
}
|
||||
|
||||
text = json.dumps(payload, indent=2, ensure_ascii=False) + "\n"
|
||||
await asyncio.to_thread(
|
||||
(directory / METADATA_FILE_NAME).write_text, text, encoding="utf-8"
|
||||
)
|
||||
|
||||
|
||||
# -- module-level helpers -------------------------------------------------
|
||||
|
||||
|
||||
def primary_file_type(sticker: Any) -> str:
|
||||
"""The file type Telegram stores this sticker as."""
|
||||
if getattr(sticker, "is_animated", False):
|
||||
return ANIMATED
|
||||
if getattr(sticker, "is_video", False):
|
||||
return VIDEO
|
||||
return STATIC
|
||||
|
||||
|
||||
def sanitize_filename_part(text: str) -> str:
|
||||
"""Strip characters that are illegal or awkward inside a file name."""
|
||||
cleaned = _UNSAFE_FILENAME_CHARS.sub("", text).strip(" .")
|
||||
return cleaned[:32]
|
||||
|
||||
|
||||
def _suffix_from_url(remote_path: str) -> str:
|
||||
suffix = Path(remote_path.split("?")[0]).suffix.lower()
|
||||
return suffix if suffix else f".{STATIC}"
|
||||
|
||||
|
||||
def _write_atomically(destination: Path, data: bytes) -> None:
|
||||
"""Write via a temporary file so an interrupted run leaves no partial file."""
|
||||
destination.parent.mkdir(parents=True, exist_ok=True)
|
||||
temporary = destination.with_name(destination.name + _PART_SUFFIX)
|
||||
temporary.write_bytes(data)
|
||||
os.replace(temporary, destination)
|
||||
|
||||
|
||||
def _enum_value(value: Any) -> Any:
|
||||
return getattr(value, "value", value)
|
||||
|
||||
|
||||
def _describe(error: BaseException) -> str:
|
||||
"""A short, user-facing description of a failure."""
|
||||
if isinstance(error, BadRequest) and "stickerset_invalid" in str(error).lower():
|
||||
return "sticker pack not found"
|
||||
if isinstance(error, TelegramError):
|
||||
return f"{type(error).__name__}: {error}"
|
||||
message = str(error).strip()
|
||||
return message or type(error).__name__
|
||||
Reference in New Issue
Block a user