Files
telegram-sticker-downloader/sticker_downloader/downloader.py
T
2026-07-27 15:16:01 +02:00

360 lines
13 KiB
Python

"""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__