From d1fd93f18082e7a515e1e3a380dd1500dd0b31ec Mon Sep 17 00:00:00 2001
From: Alex <25013571+alexhb1@users.noreply.github.com>
Date: Sun, 10 May 2026 15:20:43 +0100
Subject: [PATCH] Harden Welib URL validation (#979)
---
shelfmark/download/http.py | 184 +++++++--
shelfmark/release_sources/direct_download.py | 93 +++--
.../test_welib_url_allowlist.py | 366 +++++++++++++++++-
3 files changed, 585 insertions(+), 58 deletions(-)
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)