big upgrade UwU
This commit is contained in:
@@ -0,0 +1,84 @@
|
||||
"""Retry helper shared by the Telegram API calls and the file downloads."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from collections.abc import Awaitable, Callable
|
||||
from typing import TypeVar
|
||||
|
||||
import aiohttp
|
||||
from telegram.error import BadRequest, Forbidden, InvalidToken, NetworkError, RetryAfter
|
||||
|
||||
from sticker_downloader.errors import FetchError
|
||||
|
||||
T = TypeVar("T")
|
||||
|
||||
#: Transient failures. ``BadRequest`` subclasses ``NetworkError`` in
|
||||
#: python-telegram-bot, so it has to be excluded explicitly below — retrying a
|
||||
#: "sticker set is invalid" response would just waste three round-trips.
|
||||
_RETRIABLE = (NetworkError, aiohttp.ClientError, asyncio.TimeoutError, OSError)
|
||||
_NEVER_RETRIABLE = (BadRequest, Forbidden, InvalidToken)
|
||||
|
||||
|
||||
def is_retriable(error: BaseException) -> bool:
|
||||
"""Whether another attempt could plausibly succeed."""
|
||||
if isinstance(error, RetryAfter):
|
||||
return True
|
||||
if isinstance(error, _NEVER_RETRIABLE):
|
||||
return False
|
||||
if isinstance(error, FetchError):
|
||||
return error.retriable
|
||||
return isinstance(error, _RETRIABLE)
|
||||
|
||||
|
||||
def retry_delay(error: BaseException, attempt: int, base_delay: float) -> float:
|
||||
"""Seconds to wait before ``attempt`` (1-based) is retried.
|
||||
|
||||
Honours Telegram's own ``retry_after`` hint when it sends one, and otherwise
|
||||
backs off exponentially.
|
||||
"""
|
||||
if isinstance(error, RetryAfter):
|
||||
retry_after = getattr(error, "retry_after", None)
|
||||
if isinstance(retry_after, (int, float)):
|
||||
return float(retry_after)
|
||||
# Newer python-telegram-bot versions may hand back a timedelta.
|
||||
seconds = getattr(retry_after, "total_seconds", None)
|
||||
if callable(seconds):
|
||||
return float(seconds())
|
||||
return base_delay * (2 ** (attempt - 1))
|
||||
|
||||
|
||||
async def with_retries(
|
||||
operation: Callable[[], Awaitable[T]],
|
||||
*,
|
||||
attempts: int = 3,
|
||||
base_delay: float = 0.5,
|
||||
on_retry: Callable[[BaseException, int, float], None] | None = None,
|
||||
sleep: Callable[[float], Awaitable[None]] = asyncio.sleep,
|
||||
) -> T:
|
||||
"""Await ``operation``, retrying transient failures up to ``attempts`` times.
|
||||
|
||||
Args:
|
||||
operation: A zero-argument coroutine function; called afresh per attempt.
|
||||
attempts: Total number of tries, including the first one.
|
||||
base_delay: Seed for the exponential backoff, in seconds.
|
||||
on_retry: Called with ``(error, attempt, delay)`` before each retry.
|
||||
sleep: Injectable sleep, so tests do not have to wait.
|
||||
|
||||
Raises:
|
||||
The last error, once the attempts are exhausted or it is not retriable.
|
||||
"""
|
||||
last_error: BaseException
|
||||
for attempt in range(1, attempts + 1):
|
||||
try:
|
||||
return await operation()
|
||||
except Exception as error:
|
||||
last_error = error
|
||||
if attempt == attempts or not is_retriable(error):
|
||||
raise
|
||||
delay = retry_delay(error, attempt, base_delay)
|
||||
if on_retry is not None:
|
||||
on_retry(error, attempt, delay)
|
||||
await sleep(delay)
|
||||
|
||||
raise last_error # pragma: no cover — loop always returns or raises
|
||||
Reference in New Issue
Block a user