diff --git a/multiurl/http.py b/multiurl/http.py index 4a6e345..d01fd39 100644 --- a/multiurl/http.py +++ b/multiurl/http.py @@ -452,7 +452,9 @@ def iterate_requests(chunk_size): ) -def robust(call, maximum_tries=500, retry_after=120, mirrors=None): +def robust( + call, maximum_tries=500, retry_after=120, mirrors=None, respect_retry_header=False +): def retriable(code): return code in RETRIABLE @@ -509,6 +511,29 @@ def wrapped(url, *args, **kwargs): LOG.warning("Retrying using mirror %s", mirror) main_url = f"{mirror}{url[replace:]}" else: + nonlocal retry_after + retry_header = getattr(r, "headers", {}).get("Retry-After") + if respect_retry_header and retry_header is not None: + try: + # seconds + retry_after = int(retry_header) + except ValueError: + try: + # http date + retry_date = datetime.datetime.strptime( + retry_header, "%a, %d %b %Y %H:%M:%S GMT" + ) + + gmt = pytz.timezone("GMT") + now = datetime.datetime.now(gmt) + retry_date = gmt.localize(retry_date) + retry_sec = (retry_date - now).total_seconds() + + if retry_sec > 0: + retry_after = retry_sec + except Exception: + pass + LOG.warning("Retrying in %s seconds", retry_after) time.sleep(retry_after) LOG.info("Retrying now...") diff --git a/tests/test_robust.py b/tests/test_robust.py index 55bd2a5..a0a56b3 100644 --- a/tests/test_robust.py +++ b/tests/test_robust.py @@ -7,6 +7,7 @@ # nor does it submit to any jurisdiction. # +import datetime import logging import os import random @@ -14,9 +15,11 @@ from contextlib import contextmanager import pytest +import pytz +import requests from multiurl import download -from multiurl.http import RETRIABLE +from multiurl.http import RETRIABLE, robust def handler(signum, frame): @@ -47,6 +50,32 @@ def test_robust(): ) +def test_retry_header(): + # patch requests.get to add a Retry-After header + def patched_get(retry, *args, **kwargs): + r = requests.get(*args, **kwargs) + if callable(retry): + retry = retry() + r.headers.update({"Retry-After": retry}) + return r + + def http_date(): + gmt = pytz.timezone("GMT") + now = datetime.datetime.now(gmt) + return (now + datetime.timedelta(seconds=10)).strftime( + "%a, %d %b %Y %H:%M:%S GMT" + ) + + # test with seconds and http date format + for retry in ["5", http_date]: + with timeout(60): + code = random.choice(RETRIABLE) + r = robust( + patched_get, maximum_tries=2, retry_after=120, respect_retry_header=True + )(retry, f"http://httpbin.org/status/{code}") + assert r.status_code == code + + @pytest.mark.skipif(True, reason="Mirror disabled") def test_mirror(): download(