"""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("上游重定向次数过多")