"""P3 优化单元测试 — generation.py 三项优化. 覆盖: P3-1: _download_library_assets strict 模式 P3-2: 归属校验合并到同一 DB session P3-3: _verify_url_accessible HEAD 重试 """ from __future__ import annotations import sys from pathlib import Path from types import SimpleNamespace from unittest.mock import MagicMock, patch import pytest sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "worker")) # ── P3-3: _verify_url_accessible 重试 ─────────────────────────────────────── class TestVerifyUrlAccessibleRetry: """_verify_url_accessible 重试逻辑.""" @patch("time.sleep") @patch("urllib.request.OpenerDirector.open") def test_first_attempt_success(self, mock_open, mock_sleep): """首次成功,不重试.""" from worker_app.tasks.generation import _verify_url_accessible mock_resp = MagicMock() mock_resp.status = 200 mock_resp.__enter__ = MagicMock(return_value=mock_resp) mock_resp.__exit__ = MagicMock(return_value=False) mock_open.return_value = mock_resp assert _verify_url_accessible("https://example.com/file.mp4") is True assert mock_open.call_count == 1 mock_sleep.assert_not_called() @patch("time.sleep") @patch("urllib.request.OpenerDirector.open") def test_retry_then_success(self, mock_open, mock_sleep): """首次失败,重试后成功.""" from worker_app.tasks.generation import _verify_url_accessible # 第一次失败(网络异常),第二次成功 mock_resp_ok = MagicMock() mock_resp_ok.status = 200 mock_resp_ok.__enter__ = MagicMock(return_value=mock_resp_ok) mock_resp_ok.__exit__ = MagicMock(return_value=False) mock_open.side_effect = [ OSError("connection reset"), mock_resp_ok, ] assert _verify_url_accessible("https://example.com/file.mp4") is True assert mock_open.call_count == 2 mock_sleep.assert_called_once_with(1) @patch("time.sleep") @patch("urllib.request.OpenerDirector.open") def test_all_retries_exhausted(self, mock_open, mock_sleep): """全部重试耗尽,返回 False.""" from worker_app.tasks.generation import _verify_url_accessible mock_open.side_effect = OSError("connection refused") assert _verify_url_accessible("https://example.com/file.mp4") is False # 1 首次 + 2 重试 = 3 次 assert mock_open.call_count == 3 assert mock_sleep.call_count == 2 @patch("time.sleep") @patch("urllib.request.OpenerDirector.open") def test_http_500_then_success(self, mock_open, mock_sleep): """HTTP 500 后重试成功.""" from worker_app.tasks.generation import _verify_url_accessible mock_resp_500 = MagicMock() mock_resp_500.status = 500 mock_resp_500.__enter__ = MagicMock(return_value=mock_resp_500) mock_resp_500.__exit__ = MagicMock(return_value=False) mock_resp_200 = MagicMock() mock_resp_200.status = 200 mock_resp_200.__enter__ = MagicMock(return_value=mock_resp_200) mock_resp_200.__exit__ = MagicMock(return_value=False) mock_open.side_effect = [mock_resp_500, mock_resp_200] assert _verify_url_accessible("https://example.com/file.mp4") is True assert mock_open.call_count == 2 @patch("time.sleep") @patch("urllib.request.OpenerDirector.open") def test_custom_retries_zero(self, mock_open, mock_sleep): """retries=0 时不重试.""" from worker_app.tasks.generation import _verify_url_accessible mock_open.side_effect = OSError("timeout") assert _verify_url_accessible("https://example.com/file.mp4", retries=0) is False assert mock_open.call_count == 1 mock_sleep.assert_not_called() # ── P3-1: _download_library_assets strict 模式 ────────────────────────────── def _make_mock_asset( asset_id: str, name: str, file_url: str | None, asset_library_id: str = "lib-1", project_id: str = "", ): """构造 mock AssetModel 实例.""" return SimpleNamespace( id=asset_id, name=name, file_url=file_url, asset_library_id=asset_library_id, project_id=project_id, status="ready", file_type="video", created_at="2026-01-01", ) def _setup_mock_session(assets): """构造 mock session,返回 (mock_session, mock_query_chain).""" mock_session = MagicMock() mock_query = MagicMock() # chain: session.query().filter().filter().order_by().all() mock_session.query.return_value = mock_query mock_query.filter.return_value = mock_query mock_query.order_by.return_value = mock_query mock_query.all.return_value = assets return mock_session class TestDownloadLibraryAssetsStrictMode: """_download_library_assets strict 模式.""" @patch("worker_app.tasks.generation.SessionLocal") @patch("worker_app.tasks.generation.download_asset") def test_strict_mode_raises_on_download_failure(self, mock_download, mock_session_factory): """strict=True 时,单个素材下载失败立即抛 RuntimeError.""" from worker_app.tasks.generation import _download_library_assets assets = [ _make_mock_asset("a1", "video1.mp4", "uploads/video1.mp4"), _make_mock_asset("a2", "video2.mp4", "uploads/video2.mp4"), ] mock_session = _setup_mock_session(assets) mock_session_factory.return_value = mock_session # 第一个成功,第二个失败 mock_download.side_effect = [True, False] with pytest.raises(RuntimeError, match="素材下载失败"): _download_library_assets( Path("/tmp"), asset_library_id="lib-1", asset_ids=["a1", "a2"], strict=True, ) @patch("worker_app.tasks.generation.SessionLocal") @patch("worker_app.tasks.generation.download_asset") def test_non_strict_mode_returns_partial_results(self, mock_download, mock_session_factory): """strict=False 时,跳过失败素材,返回成功列表.""" from worker_app.tasks.generation import _download_library_assets assets = [ _make_mock_asset("a1", "video1.mp4", "uploads/video1.mp4"), _make_mock_asset("a2", "video2.mp4", "uploads/video2.mp4"), _make_mock_asset("a3", "video3.mp4", "uploads/video3.mp4"), ] mock_session = _setup_mock_session(assets) mock_session_factory.return_value = mock_session # 第一个成功,第二个失败,第三个成功 mock_download.side_effect = [True, False, True] result = _download_library_assets( Path("/tmp"), asset_library_id="lib-1", asset_ids=["a1", "a2", "a3"], strict=False, ) assert len(result) == 2 @patch("worker_app.tasks.generation.SessionLocal") @patch("worker_app.tasks.generation.download_asset") def test_non_strict_all_fail_raises(self, mock_download, mock_session_factory): """strict=False 但全部失败时仍抛 RuntimeError.""" from worker_app.tasks.generation import _download_library_assets assets = [ _make_mock_asset("a1", "video1.mp4", "uploads/video1.mp4"), ] mock_session = _setup_mock_session(assets) mock_session_factory.return_value = mock_session mock_download.return_value = False with pytest.raises(RuntimeError, match="全部下载失败"): _download_library_assets( Path("/tmp"), asset_library_id="lib-1", asset_ids=["a1"], strict=False, ) @patch("worker_app.tasks.generation.SessionLocal") @patch("worker_app.tasks.generation.download_asset") def test_strict_mode_raises_on_missing_file_url(self, mock_download, mock_session_factory): """strict=True 时,素材缺少 file_url 立即抛异常.""" from worker_app.tasks.generation import _download_library_assets assets = [ _make_mock_asset("a1", "video1.mp4", None), # file_url 为空 ] mock_session = _setup_mock_session(assets) mock_session_factory.return_value = mock_session with pytest.raises(RuntimeError, match="素材缺少 file_url"): _download_library_assets( Path("/tmp"), asset_library_id="lib-1", strict=True, ) # download_asset 不应被调用 mock_download.assert_not_called() @patch("worker_app.tasks.generation.SessionLocal") @patch("worker_app.tasks.generation.download_asset") def test_default_is_strict(self, mock_download, mock_session_factory): """默认 strict=True.""" from worker_app.tasks.generation import _download_library_assets assets = [ _make_mock_asset("a1", "video1.mp4", "uploads/video1.mp4"), ] mock_session = _setup_mock_session(assets) mock_session_factory.return_value = mock_session mock_download.return_value = False # 不传 strict 参数,默认严格模式 with pytest.raises(RuntimeError, match="素材下载失败"): _download_library_assets( Path("/tmp"), asset_library_id="lib-1", ) # ── P3-2: 归属校验合并到同一 session ──────────────────────────────────────── class TestDownloadLibraryAssetsOwnershipValidation: """归属校验合并到 _download_library_assets 同一 session.""" @patch("worker_app.tasks.generation.SessionLocal") @patch("worker_app.tasks.generation.download_asset") def test_ownership_mismatch_raises_value_error(self, mock_download, mock_session_factory): """asset_ids 不属于指定素材库时抛 ValueError.""" from worker_app.tasks.generation import _download_library_assets # asset 属于 lib-2,但请求的是 lib-1 assets = [ _make_mock_asset("a1", "video1.mp4", "uploads/video1.mp4", asset_library_id="lib-2"), ] mock_session = _setup_mock_session(assets) mock_session_factory.return_value = mock_session with pytest.raises(ValueError, match="素材不属于指定素材库"): _download_library_assets( Path("/tmp"), asset_library_id="lib-1", asset_ids=["a1"], ) # 不应调用 download_asset(校验在下载前) mock_download.assert_not_called() @patch("worker_app.tasks.generation.SessionLocal") @patch("worker_app.tasks.generation.download_asset") def test_missing_asset_ids_raises_value_error(self, mock_download, mock_session_factory): """指定的 asset_ids 不存在时抛 ValueError.""" from worker_app.tasks.generation import _download_library_assets # DB 返回空(asset_ids 不存在,query 过滤后无结果) mock_session = _setup_mock_session([]) mock_session_factory.return_value = mock_session with pytest.raises(RuntimeError, match="未找到视频素材"): _download_library_assets( Path("/tmp"), asset_library_id="lib-1", asset_ids=["nonexistent-id"], ) @patch("worker_app.tasks.generation.SessionLocal") @patch("worker_app.tasks.generation.download_asset") def test_project_ownership_mismatch_raises(self, mock_download, mock_session_factory): """项目级模式下归属不匹配抛 ValueError.""" from worker_app.tasks.generation import _download_library_assets assets = [ _make_mock_asset( "a1", "video1.mp4", "uploads/video1.mp4", asset_library_id="", project_id="proj-2", ), ] mock_session = _setup_mock_session(assets) mock_session_factory.return_value = mock_session with pytest.raises(ValueError, match="素材不属于指定项目"): _download_library_assets( Path("/tmp"), project_id="proj-1", asset_ids=["a1"], ) mock_download.assert_not_called() @patch("worker_app.tasks.generation.SessionLocal") @patch("worker_app.tasks.generation.download_asset") def test_ownership_pass_then_download(self, mock_download, mock_session_factory): """归属校验通过后正常下载.""" from worker_app.tasks.generation import _download_library_assets assets = [ _make_mock_asset("a1", "video1.mp4", "uploads/video1.mp4", asset_library_id="lib-1"), ] mock_session = _setup_mock_session(assets) mock_session_factory.return_value = mock_session mock_download.return_value = True result = _download_library_assets( Path("/tmp"), asset_library_id="lib-1", asset_ids=["a1"], ) assert len(result) == 1 mock_download.assert_called_once() @patch("worker_app.tasks.generation.SessionLocal") @patch("worker_app.tasks.generation.download_asset") def test_single_session_used(self, mock_download, mock_session_factory): """验证只创建了一个 DB session(P3-2 核心).""" from worker_app.tasks.generation import _download_library_assets assets = [ _make_mock_asset("a1", "video1.mp4", "uploads/video1.mp4", asset_library_id="lib-1"), ] mock_session = _setup_mock_session(assets) mock_session_factory.return_value = mock_session mock_download.return_value = True _download_library_assets( Path("/tmp"), asset_library_id="lib-1", asset_ids=["a1"], ) # SessionLocal 只调用一次(合并前会调用两次:校验 + 下载) assert mock_session_factory.call_count == 1