Search / upstream_http.py
vomebook
feat: support additional Reader formats
1787788
Raw History Blame Contribute Delete
10.1 kB
"""Validated upstream requests and bounded redirect reuse, independent of API state."""
import re
import time
from urllib.parse import unquote, urljoin, urlparse
import aiohttp
ALLOWED_DOWNLOAD_HOSTS = {"huggingface.co", "hf-mirror.com", "hf.co"}
ALLOWED_DOWNLOAD_HOST_SUFFIXES = (".hf.co", ".huggingface.co", ".xethub.hf.co")
MAX_REDIRECTS = 5
READER_ASSET_PATH_RE = re.compile(
r"objects/[0-9a-f]{2}/[0-9a-f]{64}/(?:[0-9a-f]{16}/)?(?:page-manifest\.json|ocr-manifest\.json|pages/page-[0-9]{6}\.(?:webp|jxl)|ocr/(?:page-[0-9]{6}\.json\.gz|book-text\.json\.gz)|(?:[a-z0-9-]+/)?(?:linearized\.pdf|chapter-manifest\.json|document\.(?:pdf|epub|mobi|azw|azw3|fb2|docx|html|txt|md|webp|jpg|jpeg|png|gif|bmp|swf)|book\.epub|audio\.(?:mp3|wav|m4a|flac|mpga)|video\.(?:mp4|mov)|epub-chapters/(?:chapter-manifest\.json|chapters/chapter-[0-9]{4}\.xhtml|resources/[A-Za-z0-9._~%+\-/]+|epub-search-index\.json\.gz)))"
)
READER_ASSET_SOURCE_RE = re.compile(
r"/datasets/vomebook/Reader-Assets/resolve/main/(?:pdf_manifest\.json|" + READER_ASSET_PATH_RE.pattern + r")"
)
VOICEOFML_READER_SOURCE_RE = re.compile(r"^/datasets/VoiceOfML/[A-Za-z0-9._-]+/(?:resolve|raw)/main/.+$")
UPSTREAM_REDIRECT_TTL_SECONDS = 300
UPSTREAM_REDIRECT_MAX_ENTRIES = 2000
upstream_redirect_cache: dict[tuple[str, str], tuple[float, str]] = {}
def is_allowed_download_host(hostname: str | None) -> bool:
return bool(hostname) and (
hostname in ALLOWED_DOWNLOAD_HOSTS or hostname.endswith(ALLOWED_DOWNLOAD_HOST_SUFFIXES)
)
def validate_download_url(url: str) -> str:
parsed = urlparse(url)
if (
parsed.scheme != "https"
or not is_allowed_download_host(parsed.hostname)
or parsed.port not in (None, 443)
or parsed.username is not None
or parsed.password is not None
):
raise ValueError(f"不允许的下载跳转地址: {parsed.hostname or 'unknown'}")
return url
def normalize_download_url(url: str) -> str:
parsed = urlparse(url)
if parsed.hostname == "hf-mirror.com":
return parsed._replace(netloc="huggingface.co").geturl()
return url
def decode_url_path(path: str) -> str:
decoded_path = path
for _ in range(8):
next_path = unquote(decoded_path)
if next_path == decoded_path:
return decoded_path
decoded_path = next_path
raise ValueError("不允许的下载来源路径")
def validate_path_safety(path: str) -> str:
decoded_path = decode_url_path(path)
if (
"\\" in decoded_path
or "\x00" in decoded_path
or any(part in (".", "..") for part in decoded_path.split("/"))
):
raise ValueError("不允许的下载来源路径")
return decoded_path
def validate_voiceofml_source_url(url: str) -> str:
url = normalize_download_url(url)
validate_download_url(url)
parsed = urlparse(url)
decoded_path = validate_path_safety(parsed.path)
if (
parsed.hostname != "huggingface.co"
or parsed.fragment
or parsed.query
or not decoded_path.startswith("/datasets/VoiceOfML/")
):
raise ValueError("只允许 VoiceOfML 数据集文件")
return url
def validate_reader_source_url(url: str) -> str:
url = normalize_download_url(url)
validate_download_url(url)
parsed = urlparse(url)
if parsed.hostname != "huggingface.co" or parsed.fragment or parsed.query:
raise ValueError("不允许的阅读来源")
decoded_path = validate_path_safety(parsed.path)
if (VOICEOFML_READER_SOURCE_RE.fullmatch(decoded_path)
or READER_ASSET_SOURCE_RE.fullmatch(decoded_path)):
return url
raise ValueError("不允许的阅读来源")
def is_immutable_reader_asset(url: str) -> bool:
parsed = urlparse(url)
return (
parsed.hostname == "huggingface.co"
and not parsed.query
and READER_ASSET_SOURCE_RE.fullmatch(validate_path_safety(parsed.path)) is not None
and "/objects/" in parsed.path
)
def validate_voiceofml_redirect_url(url: str) -> str:
normalized = normalize_download_url(url)
validate_download_url(normalized)
if urlparse(normalized).hostname == "huggingface.co":
parsed = urlparse(normalized)
cache_source = reader_cache_source_path(parsed.path)
if cache_source is not None and not parsed.fragment:
# Hugging Face redirects existing files to commit-pinned cache
# paths. Validate the mapped dataset source while preserving the
# redirect URL and its cache metadata for the upstream request.
validate_voiceofml_source_url("https://huggingface.co" + cache_source)
return normalized
return validate_voiceofml_source_url(normalized)
return normalized
def reader_cache_source_path(path: str) -> str | None:
"""Map HF's commit-pinned cache route back to its dataset source for validation."""
match = re.fullmatch(
r"/api/resolve-cache/datasets/([^/]+/[^/]+)/[0-9a-f]{40}/(.+)",
validate_path_safety(path),
)
if match:
return f"/datasets/{match[1]}/resolve/main/{match[2]}"
return None
def validate_reader_redirect_url(url: str) -> str:
normalized = normalize_download_url(url)
validate_download_url(normalized)
if urlparse(normalized).hostname == "huggingface.co":
parsed = urlparse(normalized)
cache_source = reader_cache_source_path(parsed.path)
if cache_source is not None and not parsed.fragment:
# Only redirects may use this route and its HF query metadata.
# Apply the same dataset and asset-path restrictions as the source.
validate_reader_source_url("https://huggingface.co" + cache_source)
return normalized
return validate_reader_source_url(normalized)
return normalized
def source_url_scope(url: str) -> str | None:
parsed = urlparse(url)
if parsed.hostname != "huggingface.co":
return None
decoded_path = reader_cache_source_path(parsed.path) or validate_path_safety(parsed.path)
if decoded_path.startswith("/datasets/VoiceOfML/"):
return "voiceofml"
if decoded_path.startswith("/datasets/vomebook/Reader-Assets/"):
return "reader-assets"
return None
def _redirect_cache_get(method: str, validated_url: str, url_validator, initial_scope: str | None) -> str | None:
entry = upstream_redirect_cache.get((method, validated_url))
if not entry:
return None
cached_at, target = entry
if time.monotonic() - cached_at >= UPSTREAM_REDIRECT_TTL_SECONDS:
upstream_redirect_cache.pop((method, validated_url), None)
return None
try:
checked = url_validator(target)
except ValueError:
upstream_redirect_cache.pop((method, validated_url), None)
return None
next_scope = source_url_scope(checked)
if initial_scope is not None and next_scope is not None and next_scope != initial_scope:
upstream_redirect_cache.pop((method, validated_url), None)
return None
return checked
def _redirect_cache_put(method: str, validated_url: str, target: str) -> None:
upstream_redirect_cache[(method, validated_url)] = (time.monotonic(), target)
while len(upstream_redirect_cache) > UPSTREAM_REDIRECT_MAX_ENTRIES:
oldest = min(upstream_redirect_cache, key=lambda key: upstream_redirect_cache[key][0])
upstream_redirect_cache.pop(oldest, None)
async def _request(session, method, url, **options):
request = getattr(session, "request", None)
if request is not None:
return await request(method, url, **options)
request = getattr(session, method.lower(), None)
if request is None:
raise ValueError(f"上游会话不支持 {method} 请求")
return await request(url, **options)
async def open_download_response(
session,
url: str,
timeout: aiohttp.ClientTimeout,
*,
method: str = "GET",
request_headers: dict[str, str] | None = None,
url_validator=validate_download_url,
auto_decompress: bool | None = None,
):
method = method.upper()
if method not in ("GET", "HEAD"):
raise ValueError("不支持的上游请求方法")
validated_url = url_validator(url)
initial_scope = source_url_scope(validated_url)
headers = {"Accept-Encoding": "identity"}
headers.update(request_headers or {})
request_options = {} if auto_decompress is None else {"auto_decompress": auto_decompress}
for attempt in range(2):
current_url = validated_url
if attempt == 0:
cached = _redirect_cache_get(method, validated_url, url_validator, initial_scope)
if cached is not None:
current_url = cached
walked = current_url != validated_url
for _ in range(MAX_REDIRECTS + 1):
response = await _request(
session, method, current_url, allow_redirects=False,
timeout=timeout, headers=headers, **request_options,
)
if response.status not in (301, 302, 303, 307, 308):
if response.status in (401, 403) and attempt == 0 and walked:
upstream_redirect_cache.pop((method, validated_url), None)
response.release()
break
if walked:
_redirect_cache_put(method, validated_url, current_url)
return response
location = response.headers.get("Location")
response.release()
if not location:
raise ValueError("上游重定向缺少地址")
next_url = url_validator(urljoin(current_url, location))
next_scope = source_url_scope(next_url)
if initial_scope is not None and next_scope is not None and next_scope != initial_scope:
raise ValueError("上游重定向超出允许的数据集范围")
current_url = next_url
walked = True
else:
raise ValueError("上游重定向次数过多")