Files
MediaCrawler/tests/test_media_downloader.py
T
程序员阿江(Relakkes) 0ca7b29cf0 feat(media): 重构媒体下载,支持 xhs/dy/ks/bili/wb 五平台
旧实现只覆盖 4 个平台,且把整个文件读进内存、无重试与完整性校验,
代码按平台复制粘贴了 4 份。本次用统一下载器替换:

- 新增 media_downloader/:流式写入、Range 续传、指数退避重试、大小校验、
  路径穿越防护;B 站 DASH 音视频分轨下载后交由 ffmpeg 无损合流
- 新增 media_platform/<平台>/media.py:从平台原始响应提取媒体地址,
  与下载器解耦;快手首次接入下载能力
- 开关:config.ENABLE_GET_MEDIA 与 --get_media,并打通 API/WebUI;
  同时修正旧配置项 ENABLE_GET_MEIDAS 的拼写
- 落盘按帖子聚合:{SAVE_DATA_PATH 或 data}/{platform}/media/{内容ID}/
- B 站装好 ffmpeg 时走 DASH 最高画质,否则降级 mp4 直链(产物 video-durl.mp4,
  避免低清文件阻塞后续的高清路径)
- 删除 4 个 *_store_media.py、AbstractStoreImage/Video 及各 client 的媒体 GET 方法

媒体下载失败只记录日志,不中断爬取主流程。
2026-09-17 22:54:22 +08:00

891 lines
30 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
# -*- coding: utf-8 -*-
# Copyright (c) 2025 relakkes@gmail.com
#
# This file is part of MediaCrawler project.
# Repository: https://github.com/NanmiCoder/MediaCrawler/blob/main/tests/test_media_downloader.py
# GitHub: https://github.com/NanmiCoder
# Licensed under NON-COMMERCIAL LEARNING LICENSE 1.1
#
"""MediaDownloader 单测。
全部基于 ``tests/media_server.py`` 提供的本地 HTTP 服务器,真实走 socket 与 httpx
不使用任何 mock,也不访问外网。
"""
from __future__ import annotations
import os
import shutil
import subprocess
import httpx
from pathlib import Path
import pytest
from media_downloader import MediaDownloader, MediaItem, MediaType
from media_downloader.paths import (
build_media_dir,
build_media_path,
ensure_within,
guess_extension,
redact_url,
sanitize_component,
url_fingerprint,
)
from tests.media_server import DEFAULT_CONTENT, MediaTestServer, start_server
@pytest.fixture(scope="module")
def server():
srv, _thread = start_server()
yield srv
srv.shutdown()
srv.server_close()
@pytest.fixture(autouse=True)
def _reset_server(server: MediaTestServer):
server.reset()
yield
@pytest.fixture
def downloader(server: MediaTestServer, tmp_path: Path) -> MediaDownloader:
"""重试间隔调到毫秒级,避免测试因退避等待变慢"""
return MediaDownloader(
platform="xhs",
base_dir=tmp_path,
max_retries=2,
retry_base_delay=0.001,
retry_max_delay=0.005,
)
def make_item(server: MediaTestServer, path: str, **overrides) -> MediaItem:
params = {
"url": server.url(path),
"media_type": MediaType.IMAGE,
"content_id": "note-1",
"stem": "001",
}
params.update(overrides)
return MediaItem(**params)
def part_path_for(tmp_path: Path, item: MediaItem, url: str) -> Path:
directory = build_media_dir(tmp_path, "xhs", item.content_id)
directory.mkdir(parents=True, exist_ok=True)
return directory / f"{item.stem}.{url_fingerprint(url)}.{os.getpid()}.part"
# --------------------------------------------------------------------------- 基础
@pytest.mark.asyncio
async def test_download_success_and_layout(server, downloader, tmp_path):
item = make_item(server, "/ok/photo")
result = await downloader.download(item)
assert result == tmp_path / "xhs" / "media" / "note-1" / "001.jpg"
assert result.read_bytes() == DEFAULT_CONTENT
assert not list(result.parent.glob("*.part")), "下载完成后不应残留临时文件"
@pytest.mark.asyncio
async def test_download_all_shares_one_client(server, downloader, monkeypatch):
"""同一帖子的多个媒体必须复用同一个 AsyncClient(而不是每个文件新建一个连接池)"""
import media_downloader.downloader as downloader_module
created_clients = []
original = downloader_module.make_async_client
def counting_make_async_client(**kwargs):
client = original(**kwargs)
created_clients.append(client)
return client
monkeypatch.setattr(downloader_module, "make_async_client", counting_make_async_client)
items = [
make_item(server, "/ok/a", stem="001"),
make_item(server, "/ok/b", stem="002"),
make_item(server, "/ok/c", stem="003"),
]
paths = await downloader.download_all(items)
assert len(paths) == 3
assert [path.name for path in paths] == ["001.jpg", "002.jpg", "003.jpg"]
assert server.request_count == 3
assert len(created_clients) == 1, f"3 个文件应共用一个 client,实际创建了 {len(created_clients)} 个"
@pytest.mark.asyncio
async def test_download_all_empty(downloader):
assert await downloader.download_all([]) == []
# --------------------------------------------------------------------------- 跳过与覆盖
@pytest.mark.asyncio
async def test_existing_file_is_skipped_without_request(server, downloader):
item = make_item(server, "/ok/photo")
first = await downloader.download(item)
assert first is not None
server.reset()
second = await downloader.download(item)
assert second == first
assert server.request_count == 0, "已下载完成的文件不应再次发起请求"
@pytest.mark.asyncio
async def test_overwrite_forces_redownload(server, tmp_path):
downloader = MediaDownloader(
platform="xhs",
base_dir=tmp_path,
max_retries=0,
overwrite=True,
)
item = make_item(server, "/ok/photo")
await downloader.download(item)
server.reset()
await downloader.download(item)
assert server.request_count == 1
@pytest.mark.asyncio
async def test_zero_byte_file_is_treated_as_broken(server, downloader):
item = make_item(server, "/ok/photo")
target = downloader.base_dir / "xhs" / "media" / "note-1" / "001.jpg"
target.parent.mkdir(parents=True, exist_ok=True)
target.write_bytes(b"")
server.reset()
result = await downloader.download(item)
assert result is not None
assert result.read_bytes() == DEFAULT_CONTENT
assert server.request_count == 1
# --------------------------------------------------------------------------- 续传
@pytest.mark.asyncio
async def test_resume_from_partial_file(server, downloader, tmp_path):
"""预置半个片段 -> 必须带 Range 请求,且大小校验以 Content-Range 的总长为准。
这是关键回归:206 响应的 Content-Length 只是剩余长度(3072),
若被误当成文件总长,4096 != 3072 会误判为失败。
"""
item = make_item(server, "/ok/photo")
part = part_path_for(tmp_path, item, item.url)
part.write_bytes(DEFAULT_CONTENT[:1024])
result = await downloader.download(item)
assert result is not None
assert result.read_bytes() == DEFAULT_CONTENT
assert server.request_count == 1
assert server.requests[0].headers.get("range") == "bytes=1024-"
@pytest.mark.asyncio
async def test_server_ignoring_range_falls_back_to_full_download(server, downloader, tmp_path):
item = make_item(server, "/no-range/photo")
part = part_path_for(tmp_path, item, item.url)
part.write_bytes(DEFAULT_CONTENT[:1024])
result = await downloader.download(item)
assert result is not None
assert result.read_bytes() == DEFAULT_CONTENT, "服务端不支持 Range 时应截断重写,而不是拼接"
@pytest.mark.asyncio
async def test_truncated_response_retries_with_range(server, downloader, tmp_path):
"""服务端中途断流 -> 重试时带 Range 续传 -> 最终拼接出完整文件"""
item = make_item(server, "/truncated/photo?fail=1")
result = await downloader.download(item)
assert result is not None
assert result.read_bytes() == DEFAULT_CONTENT
assert server.request_count == 2
assert server.requests[1].headers.get("range") == f"bytes={len(DEFAULT_CONTENT) // 2}-"
@pytest.mark.asyncio
async def test_stale_part_larger_than_remote_triggers_416_restart(server, downloader, tmp_path):
"""本地残留片段比远端还大 -> 416 -> 清空片段重来"""
item = make_item(server, "/ok/photo")
oversized = len(DEFAULT_CONTENT) + 5000
part = part_path_for(tmp_path, item, item.url)
part.write_bytes(b"z" * oversized)
result = await downloader.download(item)
assert result is not None
assert result.read_bytes() == DEFAULT_CONTENT
# 第一次带越界的 Range 拿到 416,片段被丢弃后第二次必须不带 Range 重新下载
assert server.request_count == 2
assert server.requests[0].headers.get("range") == f"bytes={oversized}-"
assert server.requests[1].headers.get("range") is None
@pytest.mark.asyncio
async def test_mismatched_content_range_start_discards_partial_file(server, downloader, tmp_path):
"""服务端返回的 Range 起点与请求不符时必须丢弃片段重下,而不是把错位内容拼进去"""
item = make_item(server, "/wrong-range/photo?delta=100&fail=1")
part = part_path_for(tmp_path, item, item.url)
part.write_bytes(DEFAULT_CONTENT[:1024])
result = await downloader.download(item)
assert result is not None
assert result.read_bytes() == DEFAULT_CONTENT
assert server.request_count == 2
assert server.requests[0].headers.get("range") == "bytes=1024-"
assert server.requests[1].headers.get("range") is None, "片段被丢弃后应重新完整下载"
@pytest.mark.asyncio
async def test_mismatched_content_range_never_produces_corrupt_file(server, downloader, tmp_path):
"""服务端持续返回错位 Range 时应放弃下载,绝不落盘内容错位的文件"""
item = make_item(server, "/wrong-range/broken?delta=100")
part = part_path_for(tmp_path, item, item.url)
part.write_bytes(DEFAULT_CONTENT[:1024])
result = await downloader.download(item)
assert result is None
directory = build_media_dir(downloader.base_dir, "xhs", item.content_id)
assert not list(directory.glob("*.jpg")), "内容错位的响应不能落盘"
@pytest.mark.asyncio
@pytest.mark.parametrize("bad_content_type", ["text/html", "text/plain", "application/json"])
async def test_error_page_with_200_is_rejected(server, downloader, bad_content_type):
"""CDN 防盗链页常以 200 + text/html 返回,不能当成媒体文件落盘"""
server.set_content_type("errpage", bad_content_type)
item = make_item(server, "/ctyped/errpage")
result = await downloader.download(item)
assert result is None
directory = build_media_dir(downloader.base_dir, "xhs", item.content_id)
assert not list(directory.glob("*.jpg")), "错误页不应被保存为媒体文件"
@pytest.mark.asyncio
async def test_content_type_mismatch_preserves_partial_progress(server, downloader, tmp_path):
"""错误页在读 body 之前就被拦下,本地已有的片段必须保留以便续传"""
server.set_content_type("limited", "text/html")
item = make_item(server, "/ctyped/limited")
part = part_path_for(tmp_path, item, item.url)
part.write_bytes(DEFAULT_CONTENT[:3000])
result = await downloader.download(item)
assert result is None
# 每一次尝试都要基于已有片段续传,而不是丢掉 3000 字节从头再来
ranges = [record.headers.get("range") for record in server.requests]
assert ranges == ["bytes=3000-"] * server.request_count, (
f"被拦截的响应没有消耗任何字节,已下载的进度不应被丢弃,实际: {ranges}"
)
@pytest.mark.asyncio
async def test_unsolicited_partial_response_is_rejected(server, downloader):
"""未发 Range 却收到起点非 0 的 206 时必须拒绝。
这里用"长度自洽但内容错位"的毒响应,确保拦住它的是起点守卫本身,
而不是被大小校验代偿(后者在长度恰好吻合时就失效了)。
"""
item = make_item(server, "/poison-206/photo?start=100")
result = await downloader.download(item)
assert result is None
directory = build_media_dir(downloader.base_dir, "xhs", item.content_id)
assert not list(directory.glob("*.jpg")), "起点错位的内容不能落盘"
@pytest.mark.asyncio
async def test_partial_206_with_unknown_total_is_rejected(server, downloader):
"""206 的 total 为 * 时无法确认完整性(服务端可能只给了区间的一部分),必须重下"""
item = make_item(server, "/unknown-total/photo")
result = await downloader.download(item)
assert result is None
directory = build_media_dir(downloader.base_dir, "xhs", item.content_id)
assert not list(directory.glob("*.jpg")), "无法确认完整性的内容不能落盘"
@pytest.mark.asyncio
async def test_partial_response_without_content_range_is_rejected(server, downloader, tmp_path):
"""206 缺少 Content-Range 时服务端行为不可预期(可能回全量体),必须丢弃片段重下"""
item = make_item(server, "/no-content-range/photo")
part = part_path_for(tmp_path, item, item.url)
part.write_bytes(DEFAULT_CONTENT[:1024])
result = await downloader.download(item)
assert result is None
assert not part.exists(), "无法确认区间语义的片段必须删除"
def test_total_size_semantics():
"""总长只能来自 Content-Rangetotal 未知时不反推,chunked 时忽略 Content-Length"""
# 206:以 Content-Range 的 total 为准
response = httpx.Response(206, headers={"content-range": "bytes 100-999/1000"})
content_range = MediaDownloader._parse_content_range(response)
assert content_range == (100, 999, 1000)
assert MediaDownloader._resolve_total_size(content_range, response, 206) == 1000
# total 为 *(服务端只返回部分区间):不能拿 剩余长度+已下载量 反推出"总长"
response = httpx.Response(
206, headers={"content-range": "bytes 100-999/*", "content-length": "900"}
)
content_range = MediaDownloader._parse_content_range(response)
assert content_range == (100, 999, None)
assert MediaDownloader._resolve_total_size(content_range, response, 206) is None
# 缺少 Content-Range 的 206 无法判定区间语义
assert MediaDownloader._resolve_total_size(None, httpx.Response(206), 206) is None
# chunked 响应即使带了 Content-Length 也必须忽略(RFC 7230
response = httpx.Response(200, headers={"transfer-encoding": "chunked", "content-length": "999"})
assert MediaDownloader._resolve_total_size(None, response, 200) is None
# 普通 200 用 Content-Length
response = httpx.Response(200, headers={"content-length": "4096"})
assert MediaDownloader._resolve_total_size(None, response, 200) == 4096
@pytest.mark.asyncio
async def test_chunked_response_ignores_bogus_content_length(server, downloader):
"""RFC 7230:有 Transfer-Encoding 时必须忽略 Content-Length,否则完整响应会被误判"""
item = make_item(server, "/chunked-bogus-cl/photo")
result = await downloader.download(item)
assert result is not None
assert result.read_bytes() == DEFAULT_CONTENT
@pytest.mark.asyncio
async def test_gzip_response_still_lands_decoded_bytes(server, downloader):
"""服务端无视 identity 强制压缩时,落盘的应是解码后的原始字节"""
item = make_item(server, "/gzip/photo")
result = await downloader.download(item)
assert result is not None
assert result.read_bytes() == DEFAULT_CONTENT
@pytest.mark.asyncio
async def test_malformed_url_in_backup_chain_does_not_abort(server, downloader):
"""畸形备用地址不能终止整条回退链(httpx.InvalidURL 不是 HTTPError 的子类)"""
item = make_item(
server,
"/status/500",
backup_urls=("http://h/a\nb.jpg", server.url("/ok/backup-after-bad")),
)
result = await downloader.download(item)
assert result is not None, "畸形地址应被跳过,继续尝试后面的候选"
@pytest.mark.asyncio
async def test_backup_candidates_are_capped(server, tmp_path):
"""备用地址数量必须有上限,否则超长 url_list 会放大成几十次请求"""
downloader = MediaDownloader("dy", base_dir=tmp_path, max_retries=0, max_candidates=2)
item = make_item(
server,
"/status/500",
backup_urls=tuple(server.url(f"/status/50{i}") for i in range(10)),
)
result = await downloader.download(item)
assert result is None
assert server.request_count == 2, f"候选总数应被裁剪到 2,实际请求 {server.request_count} 次"
@pytest.mark.asyncio
async def test_non_string_url_never_escapes_contract(server, downloader):
"""MediaItem 是公开契约,非字符串 url 必须收敛为 None 而不是抛异常"""
for bad_url in (123, ["http://x"], {"url": "http://x"}, None):
item = MediaItem(url=bad_url, media_type=MediaType.IMAGE, content_id="c")
assert await downloader.download(item) is None
@pytest.mark.asyncio
async def test_part_file_name_is_sanitized(server, downloader, tmp_path):
"""stem 未清洗时 .part 会写到 base_dir 之外,临时文件名同样要清洗"""
item = make_item(server, "/ok/photo", stem="../../../../tmp/ESCAPED")
result = await downloader.download(item)
assert result is not None
assert result.is_relative_to(tmp_path)
assert not list(Path("/tmp").glob("ESCAPED*")), "临时文件不得落在 base_dir 之外"
@pytest.mark.asyncio
async def test_update_credentials_switches_proxy(server, tmp_path, monkeypatch):
"""代理池就地刷新后,下载器必须跟着换代理,否则长跑时媒体下载会一直走失效代理"""
import media_downloader.downloader as downloader_module
used_proxies = []
original = downloader_module.make_async_client
def spy(**kwargs):
used_proxies.append(kwargs.get("proxy"))
return original(**kwargs)
monkeypatch.setattr(downloader_module, "make_async_client", spy)
downloader = MediaDownloader("xhs", base_dir=tmp_path, max_retries=0)
await downloader.download(make_item(server, "/ok/cred-1"))
downloader.update_credentials(proxy="http://new-proxy:8080")
await downloader.download(make_item(server, "/ok/cred-2", stem="002"))
assert used_proxies == [None, "http://new-proxy:8080"]
@pytest.mark.asyncio
async def test_binary_content_type_is_accepted(server, downloader):
"""通用二进制类型无法判定内容,应当放行(CDN 常见)"""
server.set_content_type("blob", "application/octet-stream")
item = make_item(server, "/ctyped/blob")
result = await downloader.download(item)
assert result is not None
assert result.read_bytes() == DEFAULT_CONTENT
@pytest.mark.asyncio
async def test_video_content_type_accepted_for_audio_stream(server, downloader):
"""DASH 音频轨的 Content-Type 是 audio/*,属于视频任务的一部分,必须放行"""
server.set_content_type("audio-track", "audio/mp4")
item = make_item(server, "/ctyped/audio-track", media_type=MediaType.VIDEO, stem="audio")
result = await downloader.download(item)
assert result is not None
assert result.suffix == ".m4a"
@pytest.mark.asyncio
async def test_stale_part_from_other_url_is_discarded(server, downloader, tmp_path):
"""URL 变化时旧的临时片段必须作废(不能把新内容拼到旧半成品上),且下载成功后清理干净"""
item = make_item(server, "/ok/photo")
stale_part = part_path_for(tmp_path, item, server.url("/ok/other"))
stale_part.write_bytes(b"y" * 1024)
result = await downloader.download(item)
assert result is not None
assert result.read_bytes() == DEFAULT_CONTENT
assert server.requests[0].headers.get("range") is None, "不同 URL 的片段不应被当作续传基础"
assert not stale_part.exists(), "成功后应清理同 stem 的陈旧片段"
assert not list(result.parent.glob("*.part"))
# --------------------------------------------------------------------------- 错误处理
@pytest.mark.asyncio
@pytest.mark.parametrize("status", [403, 404, 410])
async def test_fatal_status_is_not_retried(server, downloader, status):
item = make_item(server, f"/status/{status}")
result = await downloader.download(item)
assert result is None
assert server.request_count == 1, f"HTTP {status} 不应重试"
@pytest.mark.asyncio
async def test_server_error_is_retried_until_success(server, tmp_path):
downloader = MediaDownloader(
platform="xhs",
base_dir=tmp_path,
max_retries=3,
retry_base_delay=0.001,
retry_max_delay=0.005,
)
item = make_item(server, "/flaky/photo?fail=2")
result = await downloader.download(item)
assert result is not None
assert result.read_bytes() == DEFAULT_CONTENT
assert server.request_count == 3
@pytest.mark.asyncio
async def test_retry_exhausted_returns_none(server, downloader):
item = make_item(server, "/status/500")
result = await downloader.download(item)
assert result is None
assert server.request_count == 3 # 1 次 + 2 次重试
@pytest.mark.asyncio
async def test_empty_body_is_rejected(server, downloader):
item = make_item(server, "/empty/photo")
result = await downloader.download(item)
assert result is None
@pytest.mark.asyncio
async def test_timeout_returns_none(server, tmp_path):
downloader = MediaDownloader(
platform="xhs",
base_dir=tmp_path,
timeout=0.2,
max_retries=0,
)
item = make_item(server, "/slow/photo?s=2")
result = await downloader.download(item)
assert result is None
@pytest.mark.asyncio
@pytest.mark.parametrize(
"bad_url",
[
"",
"//cdn.example.com/x.jpg", # 协议相对地址,httpx 无法处理
"ftp://cdn.example.com/x.jpg",
"not-a-url",
],
)
async def test_invalid_url_is_rejected_without_request(server, downloader, bad_url):
item = MediaItem(url=bad_url, media_type=MediaType.IMAGE, content_id="note-1")
assert await downloader.download(item) is None
assert server.request_count == 0
# --------------------------------------------------------------------------- 备用地址
@pytest.mark.asyncio
async def test_backup_url_used_after_primary_fails(server, downloader):
item = make_item(
server,
"/status/403",
backup_urls=(server.url("/ok/backup"),),
)
result = await downloader.download(item)
assert result is not None
assert result.read_bytes() == DEFAULT_CONTENT
assert server.request_count == 2
assert server.paths() == ["/status/403", "/ok/backup"]
# --------------------------------------------------------------------------- 响应头行为
@pytest.mark.asyncio
async def test_chunked_response_without_content_length(server, downloader):
item = make_item(server, "/chunked/photo")
result = await downloader.download(item)
assert result is not None
assert result.read_bytes() == DEFAULT_CONTENT
@pytest.mark.asyncio
async def test_redirect_is_followed(server, downloader):
item = make_item(server, "/redirect/photo")
result = await downloader.download(item)
assert result is not None
assert result.read_bytes() == DEFAULT_CONTENT
@pytest.mark.asyncio
async def test_extension_inferred_from_content_type(server, downloader):
server.set_content_type("anon", "image/webp")
item = make_item(server, "/ctyped/anon")
result = await downloader.download(item)
assert result is not None
assert result.suffix == ".webp"
@pytest.mark.asyncio
async def test_explicit_extension_wins(server, downloader):
item = make_item(server, "/ctyped/anon", extension=".png")
result = await downloader.download(item)
assert result is not None
assert result.suffix == ".png"
@pytest.mark.asyncio
async def test_request_headers(server, downloader):
item = make_item(server, "/ok/photo")
await downloader.download(item)
headers = server.requests[0].headers
assert headers.get("accept-encoding") == "identity", (
"必须显式声明不压缩,否则解码后字节数与 Content-Length 不等,大小校验会误判"
)
assert "user-agent" in headers
# 下载器不认识平台,Referer 必须由平台侧注入(见 test_xhs_core_sends_referer
assert headers.get("referer") is None
@pytest.mark.asyncio
async def test_item_headers_override_defaults(server, downloader):
item = make_item(server, "/ok/photo", headers={"Cookie": "session=abc"})
await downloader.download(item)
assert server.requests[0].headers.get("cookie") == "session=abc"
@pytest.mark.asyncio
async def test_extra_headers_are_applied(server, tmp_path):
"""平台侧注入的 Referer 等反爬头必须真正发出去"""
downloader = MediaDownloader(
platform="bili",
base_dir=tmp_path,
max_retries=0,
extra_headers={"Referer": "https://www.bilibili.com/"},
)
item = make_item(server, "/ok/photo")
await downloader.download(item)
headers = server.requests[0].headers
assert headers.get("referer") == "https://www.bilibili.com/"
@pytest.mark.asyncio
async def test_proxy_is_passed_to_client(server, tmp_path, monkeypatch):
import media_downloader.downloader as downloader_module
captured: dict = {}
original = downloader_module.make_async_client
def spy(**kwargs):
captured.update(kwargs)
return original(**kwargs)
monkeypatch.setattr(downloader_module, "make_async_client", spy)
downloader = MediaDownloader(
platform="xhs",
base_dir=tmp_path,
proxy="http://127.0.0.1:9",
max_retries=0,
)
item = make_item(server, "/ok/photo")
# 代理不可用会失败,但我们要断言的是参数透传
await downloader.download(item)
assert captured.get("proxy") == "http://127.0.0.1:9"
assert captured.get("follow_redirects") is True
# --------------------------------------------------------------------------- 路径安全
@pytest.mark.asyncio
async def test_malicious_content_id_stays_inside_base_dir(server, downloader, tmp_path):
item = make_item(
server,
"/ok/photo",
content_id="../../../etc",
stem="../../passwd",
)
result = await downloader.download(item)
assert result is not None
assert result.resolve().is_relative_to(tmp_path.resolve())
# --------------------------------------------------------------------------- DASH 合流
def _make_dash_streams(tmp_path: Path) -> tuple[bytes, bytes]:
"""用 ffmpeg 生成一对真实的 DASH 分轨素材(视频轨 / 音频轨)"""
ffmpeg = shutil.which("ffmpeg")
if not ffmpeg:
pytest.skip("本机未安装 ffmpeg")
video = tmp_path / "_src_video.m4s"
audio = tmp_path / "_src_audio.m4s"
subprocess.run(
[ffmpeg, "-hide_banner", "-loglevel", "error", "-y", "-f", "lavfi",
"-i", "color=c=blue:s=160x120:d=1", "-c:v", "mpeg4", "-f", "mp4", str(video)],
check=True, capture_output=True,
)
subprocess.run(
[ffmpeg, "-hide_banner", "-loglevel", "error", "-y", "-f", "lavfi",
"-i", "sine=f=440:d=1", "-c:a", "aac", "-f", "mp4", str(audio)],
check=True, capture_output=True,
)
return video.read_bytes(), audio.read_bytes()
@pytest.mark.asyncio
async def test_dash_streams_are_downloaded_and_merged(server, tmp_path):
video_bytes, audio_bytes = _make_dash_streams(tmp_path)
server.set_content("dash-v", video_bytes)
server.set_content("dash-a", audio_bytes)
downloader = MediaDownloader("bili", base_dir=tmp_path / "out", max_retries=0)
item = MediaItem(
url=server.url("/ok/dash-v?ct=video/mp4"),
audio_url=server.url("/ok/dash-a?ct=audio/mp4"),
media_type=MediaType.VIDEO,
content_id="BV1test",
stem="video",
)
result = await downloader.download(item)
assert result is not None
assert result.suffix == ".mp4"
assert result.stat().st_size > 0
assert server.request_count == 2, "分轨应各下载一次"
assert not list(result.parent.glob(".tmp-*")), "合流完成后必须清理临时目录"
probe = shutil.which("ffprobe")
if probe:
stream_types = subprocess.run(
[probe, "-v", "error", "-show_entries", "stream=codec_type",
"-of", "csv=p=0", str(result)],
check=True, capture_output=True, text=True,
).stdout
assert "video" in stream_types and "audio" in stream_types
@pytest.mark.asyncio
async def test_dash_failure_cleans_up_temp_dir(server, tmp_path):
video_bytes, _ = _make_dash_streams(tmp_path)
server.set_content("dash-v2", video_bytes)
downloader = MediaDownloader("bili", base_dir=tmp_path / "out", max_retries=0)
item = MediaItem(
url=server.url("/ok/dash-v2?ct=video/mp4"),
audio_url=server.url("/status/403"),
media_type=MediaType.VIDEO,
content_id="BV1fail",
stem="video",
)
result = await downloader.download(item)
assert result is None
media_dir = tmp_path / "out" / "bili" / "media" / "BV1fail"
assert media_dir.is_dir()
assert not list(media_dir.glob(".tmp-*")), "失败的 DASH 下载同样要清理临时目录"
# --------------------------------------------------------------------------- paths 纯函数
@pytest.mark.parametrize(
("raw", "expected"),
[
("note-123_ab", "note-123_ab"),
("../../etc/passwd", "___etc_passwd"),
("", "unknown"),
("...", "unknown"),
("a" * 100, "a" * 64),
("CON", "_CON"),
("con.txt", "_con.txt"),
],
)
def test_sanitize_component(raw, expected):
assert sanitize_component(raw) == expected
@pytest.mark.parametrize("raw", ["..", ".", "", " "])
def test_sanitize_component_fallback_used_for_empty(raw):
assert sanitize_component(raw, fallback="fb") == "fb"
@pytest.mark.parametrize(
("url", "content_type", "explicit", "expected"),
[
("https://cdn.com/a/b.jpg", None, None, ".jpg"),
("https://cdn.com/a/b.jpeg?x=1", None, None, ".jpeg"),
("https://cdn.com/a/b.mp4!large", None, None, ".mp4"),
("https://cdn.com/a/b", "image/webp; charset=utf-8", None, ".webp"),
("https://cdn.com/a/b.unknownext", "video/mp4", None, ".mp4"),
("https://cdn.com/a/b.jpg", "image/png", ".png", ".png"),
("https://cdn.com/a/b.exe", None, None, ".jpg"),
("", None, None, ".jpg"),
],
)
def test_guess_extension(url, content_type, explicit, expected):
assert guess_extension(url, content_type=content_type, explicit=explicit, default=".jpg") == expected
def test_guess_extension_ignores_explicit_outside_whitelist():
assert guess_extension("https://cdn.com/a.jpg", explicit=".exe", default=".jpg") == ".jpg"
@pytest.mark.parametrize(
("platform", "content_id", "expected"),
[
("xhs", "note-1", "xhs/media/note-1"),
("dy", "../../evil", "dy/media/___evil"),
],
)
def test_build_media_dir(tmp_path, platform, content_id, expected):
assert build_media_dir(tmp_path, platform, content_id) == tmp_path / expected
def test_build_media_path_normalizes_extension():
path = build_media_path("/base", "xhs", "note-1", "cover", "JPG")
assert path == Path("/base/xhs/media/note-1/cover.jpg")
def test_ensure_within_rejects_escape(tmp_path):
with pytest.raises(ValueError):
ensure_within(tmp_path, tmp_path / ".." / "outside.mp4")
def test_url_fingerprint_is_stable_and_short():
fingerprint = url_fingerprint("https://cdn.com/a.jpg?sig=1")
assert fingerprint == url_fingerprint("https://cdn.com/a.jpg?sig=1")
assert fingerprint != url_fingerprint("https://cdn.com/a.jpg?sig=2")
assert len(fingerprint) == 8
def test_redact_url_drops_query():
assert redact_url("https://cdn.com/a.jpg?sign=secret&t=1") == "https://cdn.com/a.jpg"
assert redact_url("not-a-url") == "<invalid-url>"