diff --git a/shelfmark/download/http.py b/shelfmark/download/http.py index 8460e28e..0220c4e9 100644 --- a/shelfmark/download/http.py +++ b/shelfmark/download/http.py @@ -54,6 +54,62 @@ def _raise_too_many_redirects(message: str) -> NoReturn: raise requests.exceptions.TooManyRedirects(message) +def _close_response(response: requests.Response) -> None: + close = getattr(response, "close", None) + if callable(close): + close() + + +def _get_redirect_url(current_url: str, response: requests.Response) -> str: + location = response.headers.get("Location", "") + if not location: + _raise_too_many_redirects(f"Redirect with no Location header: {current_url}") + return urljoin(current_url, location) + + +def _redirect_blocked( + redirect_url: str, + redirect_allowed: Callable[[str], bool] | None, +) -> bool: + return redirect_allowed is not None and not redirect_allowed(redirect_url) + + +def _preflight_redirects_for_bypasser( + url: str, + headers: dict[str, str], + redirect_allowed: Callable[[str], bool], + session: requests.Session | None, +) -> tuple[bool, str]: + current_url = url + redirects_followed = 0 + request_client = session or requests + + while True: + response = request_client.get( + current_url, + proxies=get_proxies(current_url), + timeout=REQUEST_TIMEOUT, + headers=headers, + allow_redirects=False, + verify=get_ssl_verify(current_url), + ) + try: + if not response.is_redirect: + return True, current_url + + redirect_url = _get_redirect_url(current_url, response) + if _redirect_blocked(redirect_url, redirect_allowed): + logger.warning("Blocked redirect from %s to %s", current_url, redirect_url) + return False, current_url + + redirects_followed += 1 + if redirects_followed > _MAX_REDIRECTS: + _raise_too_many_redirects(f"Too many redirects for {current_url}") + current_url = redirect_url + finally: + _close_response(response) + + def _get_internal_bypasser() -> ModuleType: """Lazy import of internal bypasser module.""" global _internal_bypasser @@ -229,6 +285,7 @@ def html_get_page( include_response_url: bool = False, success_delay: float = 1.0, session: requests.Session | None = None, + redirect_allowed: Callable[[str], bool] | None = None, ) -> str | tuple[str, str]: """Fetch HTML content from a URL with retry mechanism. @@ -245,6 +302,8 @@ def html_get_page( resolved response URL after redirects. success_delay: Optional delay (seconds) after successful fetch. session: Optional requests session to reuse across attempts. + redirect_allowed: Optional callback used to validate each redirect target. + When omitted, existing requests redirect behavior is unchanged. """ @@ -271,6 +330,25 @@ def html_get_page( cookies: dict[str, str] = {} try: if use_bypasser_now and _is_cf_bypass_enabled(): + headers = {"User-Agent": DOWNLOAD_HEADERS["User-Agent"]} + if redirect_allowed is not None: + try: + redirects_ok, current_url = _preflight_redirects_for_bypasser( + current_url, + headers, + redirect_allowed, + session, + ) + except requests.exceptions.RequestException as e: + logger.debug( + "Bypasser redirect preflight failed for %s: %s", + current_url, + e, + ) + else: + if not redirects_ok: + return _result("", current_url) + if status_callback: status_callback("resolving", "Bypassing protection...") heartbeat_stop = Event() @@ -311,7 +389,7 @@ def html_get_page( # requests follow those redirects, the request fails on DNS and we rotate away # from an otherwise working mirror. Handle AA redirects manually instead. is_aa_url = network.should_rotate_dns_for_url(current_url) - allow_redirects = not is_aa_url + allow_redirects = not is_aa_url and redirect_allowed is None redirects_followed = 0 while True: @@ -328,14 +406,13 @@ def html_get_page( verify=get_ssl_verify(current_url), ) - if is_aa_url and response.is_redirect: - location = response.headers.get("Location", "") - if not location: - _raise_too_many_redirects( - f"Redirect with no Location header: {current_url}" - ) + if (is_aa_url or redirect_allowed is not None) and response.is_redirect: + redirect_url = _get_redirect_url(current_url, response) + if _redirect_blocked(redirect_url, redirect_allowed): + logger.warning("Blocked redirect from %s to %s", current_url, redirect_url) + _close_response(response) + return _result("", current_url) - redirect_url = urljoin(current_url, location) current_host = urlparse(current_url).hostname or "" redirect_host = urlparse(redirect_url).hostname or "" @@ -357,7 +434,7 @@ def html_get_page( # Reset per-request state for the new host. headers = {"User-Agent": DOWNLOAD_HEADERS["User-Agent"]} is_aa_url = network.should_rotate_dns_for_url(current_url) - allow_redirects = not is_aa_url + allow_redirects = not is_aa_url and redirect_allowed is None redirects_followed = 0 continue @@ -374,6 +451,8 @@ def html_get_page( if redirects_followed > _MAX_REDIRECTS: _raise_too_many_redirects(f"Too many redirects for {current_url}") current_url = redirect_url + is_aa_url = network.should_rotate_dns_for_url(current_url) + allow_redirects = not is_aa_url and redirect_allowed is None continue response.raise_for_status() @@ -452,6 +531,7 @@ def download_url( _selector: network.AAMirrorSelector | None = None, status_callback: Callable[[str, str | None], None] | None = None, referer: str | None = None, + redirect_allowed: Callable[[str], bool] | None = None, ) -> BytesIO | None: """Download content from URL with automatic retry and resume support.""" selector = _selector or network.AAMirrorSelector() @@ -487,16 +567,39 @@ def download_url( MAX_DOWNLOAD_RETRIES, ) # Try with CF cookies/UA if available - cookies = _apply_cf_bypass(current_url, headers) - response = requests.get( - current_url, - stream=True, - proxies=get_proxies(current_url), - timeout=REQUEST_TIMEOUT, - cookies=cookies, - headers=headers, - verify=get_ssl_verify(current_url), - ) + redirects_followed = 0 + while True: + cookies = _apply_cf_bypass(current_url, headers) + response = requests.get( + current_url, + stream=True, + proxies=get_proxies(current_url), + timeout=REQUEST_TIMEOUT, + cookies=cookies, + headers=headers, + allow_redirects=redirect_allowed is None, + verify=get_ssl_verify(current_url), + ) + + if redirect_allowed is not None and response.is_redirect: + redirect_url = _get_redirect_url(current_url, response) + if _redirect_blocked(redirect_url, redirect_allowed): + logger.warning( + "Blocked download redirect from %s to %s", + current_url, + redirect_url, + ) + _close_response(response) + return None + redirects_followed += 1 + if redirects_followed > _MAX_REDIRECTS: + _raise_too_many_redirects(f"Too many redirects for {current_url}") + _close_response(response) + current_url = redirect_url + continue + + break + response.raise_for_status() if status_callback: @@ -580,6 +683,7 @@ def download_url( progress_callback, cancel_flag, headers, + redirect_allowed, ) if resumed: return resumed @@ -629,6 +733,7 @@ def _try_resume( progress_callback: Callable[[float], None] | None, cancel_flag: Event | None, base_headers: dict | None = None, + redirect_allowed: Callable[[str], bool] | None = None, ) -> BytesIO | None: """Try to resume an interrupted download.""" for attempt in range(MAX_RESUME_ATTEMPTS): @@ -646,16 +751,37 @@ def _try_resume( **(base_headers or DOWNLOAD_HEADERS), "Range": f"bytes={start_byte}-", } - cookies = _apply_cf_bypass(url, resume_headers) - response = requests.get( - url, - stream=True, - proxies=get_proxies(url), - timeout=REQUEST_TIMEOUT, - headers=resume_headers, - cookies=cookies, - verify=get_ssl_verify(url), - ) + current_url = url + redirects_followed = 0 + while True: + cookies = _apply_cf_bypass(current_url, resume_headers) + response = requests.get( + current_url, + stream=True, + proxies=get_proxies(current_url), + timeout=REQUEST_TIMEOUT, + headers=resume_headers, + cookies=cookies, + allow_redirects=redirect_allowed is None, + verify=get_ssl_verify(current_url), + ) + if redirect_allowed is not None and response.is_redirect: + redirect_url = _get_redirect_url(current_url, response) + if _redirect_blocked(redirect_url, redirect_allowed): + logger.warning( + "Blocked resume redirect from %s to %s", + current_url, + redirect_url, + ) + _close_response(response) + return None + redirects_followed += 1 + if redirects_followed > _MAX_REDIRECTS: + _raise_too_many_redirects(f"Too many redirects for {current_url}") + _close_response(response) + current_url = redirect_url + continue + break # Check resume support if response.status_code == _HTTP_STATUS_OK: # Server doesn't support resume diff --git a/shelfmark/release_sources/direct_download.py b/shelfmark/release_sources/direct_download.py index baf9c9e4..d57a39af 100644 --- a/shelfmark/release_sources/direct_download.py +++ b/shelfmark/release_sources/direct_download.py @@ -222,8 +222,27 @@ def _is_configured_zlib_link(url: str) -> bool: return False -def _get_direct_download_allowed_hosts() -> set[str]: - """Return configured hosts allowed for direct-download fallback fetches.""" +def _direct_download_origin(url: str) -> tuple[str, str, int] | None: + """Return normalized origin tuple for allowlist comparisons.""" + parsed = urlparse(url) + scheme = parsed.scheme.lower() + if scheme not in {"http", "https"} or not parsed.hostname: + return None + + try: + port = parsed.port + except ValueError: + return None + + hostname = parsed.hostname.lower().rstrip(".") + if not hostname: + return None + + return (scheme, hostname, port or (443 if scheme == "https" else 80)) + + +def _get_direct_download_allowed_origins() -> set[tuple[str, str, int]]: + """Return configured origins allowed for direct-download fallback fetches.""" from shelfmark.core import mirrors candidate_urls = [ @@ -231,24 +250,17 @@ def _get_direct_download_allowed_hosts() -> set[str]: *mirrors.get_aa_mirrors(), *mirrors.get_welib_mirrors(), ] - hosts: set[str] = set() + origins: set[tuple[str, str, int]] = set() for candidate_url in candidate_urls: - parsed = urlparse(candidate_url) - if parsed.scheme not in {"http", "https"}: - continue - if parsed.hostname: - hosts.add(parsed.hostname.lower().rstrip(".")) - return hosts + if origin := _direct_download_origin(candidate_url): + origins.add(origin) + return origins def _is_allowed_direct_download_url(url: str) -> bool: - """Return True when URL scheme and host match configured direct-download mirrors.""" - parsed = urlparse(url) - if parsed.scheme not in {"http", "https"} or not parsed.hostname: - return False - - hostname = parsed.hostname.lower().rstrip(".") - return hostname in _get_direct_download_allowed_hosts() + """Return True when URL origin matches configured direct-download mirrors.""" + origin = _direct_download_origin(url) + return bool(origin and origin in _get_direct_download_allowed_origins()) def _get_md5_url_template(source_id: str) -> str | None: @@ -933,8 +945,15 @@ def _try_download_url( if status_callback: status_callback("resolving", f"Trying {source_context}") + redirect_allowed = _is_allowed_direct_download_url if source_id == "welib" else None download_url = _get_download_url( - url, book_info.title, cancel_flag, status_callback, selector, source_context + url, + book_info.title, + cancel_flag, + status_callback, + selector, + source_context, + redirect_allowed=redirect_allowed, ) if not download_url: _raise_runtime_error("No download URL resolved") @@ -955,6 +974,7 @@ def _try_download_url( selector, status_callback, referer=url, + redirect_allowed=redirect_allowed, ) if not data: @@ -1008,6 +1028,7 @@ def _get_download_urls_from_welib( selector=selector or network.AAMirrorSelector(), cancel_flag=cancel_flag, status_callback=status_callback, + redirect_allowed=_is_allowed_direct_download_url, ) except ( SearchUnavailableError, @@ -1201,6 +1222,7 @@ def _get_download_url( status_callback: Callable[[str, str | None], None] | None = None, selector: network.AAMirrorSelector | None = None, source_context: str | None = None, + redirect_allowed: Callable[[str], bool] | None = None, ) -> str: """Extract actual download URL from various source pages. @@ -1211,6 +1233,7 @@ def _get_download_url( status_callback: Optional callback for status updates selector: Optional AA mirror selector source_context: Optional context string like "Welib (1/12)" for status messages + redirect_allowed: Optional callback used to validate redirect targets. """ sel = selector or network.AAMirrorSelector() @@ -1218,7 +1241,11 @@ def _get_download_url( # AA fast download API (JSON response) if link.startswith(f"{network.get_aa_base_url()}/dyn/api/fast_download.json"): page = downloader.html_get_page( - link, selector=sel, cancel_flag=cancel_flag, status_callback=status_callback + link, + selector=sel, + cancel_flag=cancel_flag, + status_callback=status_callback, + redirect_allowed=redirect_allowed, ) page_data = json.loads(_html_response_text(page)) download_url = page_data.get("download_url", "") @@ -1230,7 +1257,11 @@ def _get_download_url( return _extract_libgen_download_url(link, cancel_flag) html = downloader.html_get_page( - link, selector=sel, cancel_flag=cancel_flag, status_callback=status_callback + link, + selector=sel, + cancel_flag=cancel_flag, + status_callback=status_callback, + redirect_allowed=redirect_allowed, ) if not html: return "" @@ -1245,7 +1276,11 @@ def _get_download_url( # Retry after delay if page not fully loaded time.sleep(2) html = downloader.html_get_page( - link, selector=sel, cancel_flag=cancel_flag, status_callback=status_callback + link, + selector=sel, + cancel_flag=cancel_flag, + status_callback=status_callback, + redirect_allowed=redirect_allowed, ) if html: soup = BeautifulSoup(_html_response_text(html), "html.parser") @@ -1255,7 +1290,14 @@ def _get_download_url( # AA slow download / partner servers elif "/slow_download/" in link: url = _extract_slow_download_url( - soup, link, title, cancel_flag, status_callback, sel, source_context + soup, + link, + title, + cancel_flag, + status_callback, + sel, + source_context, + redirect_allowed=redirect_allowed, ) else: @@ -1279,6 +1321,7 @@ def _extract_slow_download_url( status_callback: Callable[[str, str | None], None] | None, selector: network.AAMirrorSelector, source_context: str | None = None, + redirect_allowed: Callable[[str], bool] | None = None, ) -> str: """Extract download URL from AA slow download pages.""" html_str = str(soup) @@ -1373,7 +1416,13 @@ def _extract_slow_download_url( status_callback("resolving", f"{source_context} - Fetching") return _get_download_url( - link, title, cancel_flag, status_callback, selector, source_context + link, + title, + cancel_flag, + status_callback, + selector, + source_context, + redirect_allowed=redirect_allowed, ) link_texts = [a.get_text(strip=True)[:50] for a in soup.find_all("a", href=True)[:10]] diff --git a/tests/direct_download/test_welib_url_allowlist.py b/tests/direct_download/test_welib_url_allowlist.py index 2d640133..10da694d 100644 --- a/tests/direct_download/test_welib_url_allowlist.py +++ b/tests/direct_download/test_welib_url_allowlist.py @@ -1,5 +1,8 @@ from io import BytesIO +import pytest +import requests + from shelfmark.release_sources import BrowseRecord @@ -12,22 +15,187 @@ def _book() -> BrowseRecord: ) -def _enable_welib_only(monkeypatch, dd): +def _enable_welib_only( + monkeypatch, + dd, + *, + template: str = "https://welib.example/md5/{md5}", + mirrors: list[str] | None = None, +): monkeypatch.setattr(dd, "_get_source_priority", lambda: [{"id": "welib", "enabled": True}]) monkeypatch.setattr(dd, "_is_source_enabled", lambda source_id: source_id == "welib") monkeypatch.setattr(dd.config, "USE_CF_BYPASS", True) monkeypatch.setattr( "shelfmark.core.mirrors.get_welib_url_template", - lambda: "https://welib.example/md5/{md5}", + lambda: template, ) monkeypatch.setattr( "shelfmark.core.mirrors.get_welib_mirrors", - lambda: ["https://welib.example"], + lambda: mirrors or ["https://welib.example"], ) monkeypatch.setattr("shelfmark.core.mirrors.get_aa_mirrors", lambda: []) monkeypatch.setattr(dd.network, "get_aa_base_url", lambda: "https://annas.example") +class _FakeResponse: + def __init__( + self, + status_code: int, + *, + headers: dict[str, str] | None = None, + text: str = "", + chunks: list[bytes] | None = None, + url: str = "", + ) -> None: + self.status_code = status_code + self.headers = headers or {} + self.text = text + self._chunks = chunks or [] + self.url = url + + @property + def is_redirect(self) -> bool: + return self.status_code in (301, 302, 303, 307, 308) and bool(self.headers.get("Location")) + + def raise_for_status(self) -> None: + return None + + def iter_content(self, chunk_size: int = 8192): + del chunk_size + yield from self._chunks + + def close(self) -> None: + return None + + +class _DummyProgressBar: + def __init__(self, *args, **kwargs) -> None: + del args, kwargs + + def update(self, amount: int) -> None: + del amount + + def close(self) -> None: + return None + + +def _use_direct_welib_http(monkeypatch, dd) -> None: + monkeypatch.setattr(dd.downloader, "_is_cf_bypass_enabled", lambda: False) + monkeypatch.setattr(dd.downloader, "get_proxies", lambda _url: {}) + monkeypatch.setattr(dd.downloader, "get_ssl_verify", lambda _url: True) + monkeypatch.setattr(dd.downloader.time, "sleep", lambda _seconds: None) + monkeypatch.setattr(dd.downloader, "tqdm", _DummyProgressBar) + + +@pytest.mark.parametrize("preflight_result", ["non_redirect", "request_error"]) +def test_html_get_page_with_allowlist_still_uses_bypasser_after_preflight( + monkeypatch, preflight_result +): + from shelfmark.download import http as downloader + + preflighted_urls: list[str] = [] + bypassed_urls: list[str] = [] + + def fake_get(url: str, **kwargs): + preflighted_urls.append(url) + assert kwargs["allow_redirects"] is False + if preflight_result == "request_error": + raise requests.exceptions.ConnectionError("preflight failed") + return _FakeResponse(200, text="direct response should not be used", url=url) + + def fake_get_bypassed_page(url: str, *_args, **_kwargs): + bypassed_urls.append(url) + return "bypassed" + + monkeypatch.setattr(downloader, "_is_cf_bypass_enabled", lambda: True) + monkeypatch.setattr(downloader, "get_proxies", lambda _url: {}) + monkeypatch.setattr(downloader, "get_ssl_verify", lambda _url: True) + monkeypatch.setattr(downloader.requests, "get", fake_get) + monkeypatch.setattr(downloader, "get_bypassed_page", fake_get_bypassed_page) + + result = downloader.html_get_page( + "https://welib.example/md5/abc123", + use_bypasser=True, + redirect_allowed=lambda url: url.startswith("https://welib.example/"), + ) + + assert result == "bypassed" + assert preflighted_urls == ["https://welib.example/md5/abc123"] + assert bypassed_urls == ["https://welib.example/md5/abc123"] + + +def test_html_get_page_with_allowlist_blocks_redirect_before_bypasser(monkeypatch): + from shelfmark.download import http as downloader + + preflighted_urls: list[str] = [] + + def fake_get(url: str, **kwargs): + preflighted_urls.append(url) + assert kwargs["allow_redirects"] is False + return _FakeResponse( + 302, + headers={"Location": "https://untrusted.example/md5/abc123"}, + url=url, + ) + + def unexpected_bypasser(*_args, **_kwargs): + raise AssertionError("disallowed redirect must block before bypasser") + + monkeypatch.setattr(downloader, "_is_cf_bypass_enabled", lambda: True) + monkeypatch.setattr(downloader, "get_proxies", lambda _url: {}) + monkeypatch.setattr(downloader, "get_ssl_verify", lambda _url: True) + monkeypatch.setattr(downloader.requests, "get", fake_get) + monkeypatch.setattr(downloader, "get_bypassed_page", unexpected_bypasser) + + result = downloader.html_get_page( + "https://welib.example/md5/abc123", + use_bypasser=True, + redirect_allowed=lambda url: url.startswith("https://welib.example/"), + ) + + assert result == "" + assert preflighted_urls == ["https://welib.example/md5/abc123"] + + +def test_html_get_page_with_allowlist_passes_preflight_redirect_url_to_bypasser(monkeypatch): + from shelfmark.download import http as downloader + + preflighted_urls: list[str] = [] + bypassed_urls: list[str] = [] + + def fake_get(url: str, **kwargs): + preflighted_urls.append(url) + assert kwargs["allow_redirects"] is False + if url == "https://welib.example/md5/abc123": + return _FakeResponse(302, headers={"Location": "/landing/abc123"}, url=url) + if url == "https://welib.example/landing/abc123": + return _FakeResponse(200, text="direct response should not be used", url=url) + raise AssertionError(f"unexpected preflight URL: {url}") + + def fake_get_bypassed_page(url: str, *_args, **_kwargs): + bypassed_urls.append(url) + return "bypassed redirect target" + + monkeypatch.setattr(downloader, "_is_cf_bypass_enabled", lambda: True) + monkeypatch.setattr(downloader, "get_proxies", lambda _url: {}) + monkeypatch.setattr(downloader, "get_ssl_verify", lambda _url: True) + monkeypatch.setattr(downloader.requests, "get", fake_get) + monkeypatch.setattr(downloader, "get_bypassed_page", fake_get_bypassed_page) + + result = downloader.html_get_page( + "https://welib.example/md5/abc123", + use_bypasser=True, + redirect_allowed=lambda url: url.startswith("https://welib.example/"), + ) + + assert result == "bypassed redirect target" + assert preflighted_urls == [ + "https://welib.example/md5/abc123", + "https://welib.example/landing/abc123", + ] + assert bypassed_urls == ["https://welib.example/landing/abc123"] + + def test_welib_rejects_hostile_returned_url_before_fetch(monkeypatch, tmp_path): import shelfmark.release_sources.direct_download as dd @@ -52,7 +220,7 @@ def test_welib_rejects_hostile_returned_url_before_fetch(monkeypatch, tmp_path): assert fetched_urls == ["https://welib.example/md5/abc123"] -def test_welib_allows_configured_host_returned_url(monkeypatch, tmp_path): +def test_welib_allows_configured_origin_with_default_https_port(monkeypatch, tmp_path): import shelfmark.release_sources.direct_download as dd _enable_welib_only(monkeypatch, dd) @@ -62,6 +230,79 @@ def test_welib_allows_configured_host_returned_url(monkeypatch, tmp_path): def fake_html_get_page(url: str, **_kwargs): fetched_pages.append(url) if url == "https://welib.example/md5/abc123": + return 'Download' + raise AssertionError(f"unexpected fetch: {url}") + + def fake_download_url(url: str, *_args, referer: str | None = None, **_kwargs): + downloaded.append((url, referer)) + payload = BytesIO(b"x" * (11 * 1024)) + payload.seek(0, 2) + return payload + + book_path = tmp_path / "book.epub" + + monkeypatch.setattr(dd.downloader, "html_get_page", fake_html_get_page) + monkeypatch.setattr(dd.downloader, "download_url", fake_download_url) + + result = dd._download_book(_book(), book_path) + + assert result == "https://welib.example:443/files/book.epub" + assert fetched_pages == ["https://welib.example/md5/abc123"] + assert downloaded == [ + ("https://welib.example:443/files/book.epub", "https://welib.example/md5/abc123") + ] + assert book_path.read_bytes() == b"x" * (11 * 1024) + + +@pytest.mark.parametrize( + "returned_url", + [ + "http://welib.example/files/book.epub", + "https://welib.example:444/files/book.epub", + "http://welib.example:8080/files/book.epub", + ], +) +def test_welib_rejects_same_host_different_origin_before_download( + monkeypatch, tmp_path, returned_url +): + import shelfmark.release_sources.direct_download as dd + + _enable_welib_only(monkeypatch, dd) + fetched_pages: list[str] = [] + + def fake_html_get_page(url: str, **_kwargs): + fetched_pages.append(url) + if url == "https://welib.example/md5/abc123": + return f'Download' + raise AssertionError(f"unexpected fetch: {url}") + + def unexpected_download(*_args, **_kwargs): + raise AssertionError("different-origin URL must not reach file download") + + monkeypatch.setattr(dd.downloader, "html_get_page", fake_html_get_page) + monkeypatch.setattr(dd.downloader, "download_url", unexpected_download) + + result = dd._download_book(_book(), tmp_path / "book.epub") + + assert result is None + assert fetched_pages == ["https://welib.example/md5/abc123"] + + +def test_welib_allows_explicit_non_default_origin(monkeypatch, tmp_path): + import shelfmark.release_sources.direct_download as dd + + _enable_welib_only( + monkeypatch, + dd, + template="http://welib.example:8080/md5/{md5}", + mirrors=["http://welib.example:8080"], + ) + fetched_pages: list[str] = [] + downloaded: list[tuple[str, str | None]] = [] + + def fake_html_get_page(url: str, **_kwargs): + fetched_pages.append(url) + if url == "http://welib.example:8080/md5/abc123": return 'Download' raise AssertionError(f"unexpected fetch: {url}") @@ -78,9 +319,120 @@ def test_welib_allows_configured_host_returned_url(monkeypatch, tmp_path): result = dd._download_book(_book(), book_path) - assert result == "https://welib.example/files/book.epub" - assert fetched_pages == ["https://welib.example/md5/abc123"] + assert result == "http://welib.example:8080/files/book.epub" + assert fetched_pages == ["http://welib.example:8080/md5/abc123"] assert downloaded == [ - ("https://welib.example/files/book.epub", "https://welib.example/md5/abc123") + ("http://welib.example:8080/files/book.epub", "http://welib.example:8080/md5/abc123") + ] + assert book_path.read_bytes() == b"x" * (11 * 1024) + + +@pytest.mark.parametrize( + "redirect_url", + [ + "http://169.254.169.254/latest", + "https://welib.example:444/md5/abc123", + ], +) +def test_welib_rejects_page_resolution_redirect_outside_allowed_origin( + monkeypatch, tmp_path, redirect_url +): + import shelfmark.release_sources.direct_download as dd + + _enable_welib_only(monkeypatch, dd) + _use_direct_welib_http(monkeypatch, dd) + fetched_urls: list[str] = [] + + def fake_get(url: str, **_kwargs): + fetched_urls.append(url) + if url == "https://welib.example/md5/abc123": + return _FakeResponse(302, headers={"Location": redirect_url}, url=url) + raise AssertionError(f"unexpected fetch: {url}") + + monkeypatch.setattr(dd.downloader.requests, "get", fake_get) + + result = dd._download_book(_book(), tmp_path / "book.epub") + + assert result is None + assert fetched_urls == ["https://welib.example/md5/abc123"] + + +def test_welib_rejects_final_file_redirect_outside_allowed_origin(monkeypatch, tmp_path): + import shelfmark.release_sources.direct_download as dd + + _enable_welib_only(monkeypatch, dd) + _use_direct_welib_http(monkeypatch, dd) + fetched_urls: list[str] = [] + + def fake_get(url: str, **_kwargs): + fetched_urls.append(url) + if url == "https://welib.example/md5/abc123": + return _FakeResponse( + 200, + text='Download', + url=url, + ) + if url == "https://welib.example/files/book.epub": + return _FakeResponse( + 302, + headers={"Location": "https://untrusted.example/files/book.epub"}, + url=url, + ) + raise AssertionError(f"unexpected fetch: {url}") + + monkeypatch.setattr(dd.downloader.requests, "get", fake_get) + + result = dd._download_book(_book(), tmp_path / "book.epub") + + assert result is None + assert fetched_urls == [ + "https://welib.example/md5/abc123", + "https://welib.example/files/book.epub", + ] + assert not (tmp_path / "book.epub").exists() + + +def test_welib_allows_same_origin_redirects_during_resolution_and_download(monkeypatch, tmp_path): + import shelfmark.release_sources.direct_download as dd + + _enable_welib_only(monkeypatch, dd) + _use_direct_welib_http(monkeypatch, dd) + fetched_urls: list[str] = [] + + def fake_get(url: str, **_kwargs): + fetched_urls.append(url) + if url == "https://welib.example/md5/abc123": + return _FakeResponse(302, headers={"Location": "/landing/abc123"}, url=url) + if url == "https://welib.example/landing/abc123": + return _FakeResponse( + 200, + text='Download', + url=url, + ) + if url == "https://welib.example/files/book.epub": + return _FakeResponse(302, headers={"Location": "/files/book-v2.epub"}, url=url) + if url == "https://welib.example/files/book-v2.epub": + return _FakeResponse( + 200, + headers={ + "content-length": str(11 * 1024), + "content-type": "application/epub+zip", + }, + chunks=[b"x" * (11 * 1024)], + url=url, + ) + raise AssertionError(f"unexpected fetch: {url}") + + book_path = tmp_path / "book.epub" + monkeypatch.setattr(dd.downloader.requests, "get", fake_get) + + result = dd._download_book(_book(), book_path) + + assert result == "https://welib.example/files/book.epub" + assert fetched_urls == [ + "https://welib.example/md5/abc123", + "https://welib.example/landing/abc123", + "https://welib.example/files/book.epub", + "https://welib.example/files/book-v2.epub", ] assert book_path.read_bytes() == b"x" * (11 * 1024)