diff --git a/tests/unit/test_batch_download.py b/tests/unit/test_batch_download.py index a5f10dd62..58c669bce 100755 --- a/tests/unit/test_batch_download.py +++ b/tests/unit/test_batch_download.py @@ -36,6 +36,42 @@ class _FakeGeneratedVideoRepository: # module, not at the batch_download module. +def _make_bound_task(): + """构建 mock 的 celery task self 对象(bind=True 场景)。""" + task = MagicMock() + task.name = "worker.batch_download_videos" + task.max_retries = 1 + task.retry = MagicMock() + return task + + +def _get_raw_task_fn(fn): + """从 celery Task 对象或 PromiseProxy 中提取原始函数(带 self 参数). + + 用于单元测试:绕过 celery 的 self 注入,手动传入 mock task 对象。 + """ + # 解开 PromiseProxy + if hasattr(fn, "_get_current_object"): + fn = fn._get_current_object() + # 从 celery Task 中提取原始函数(__wrapped__ 是 bound method,__func__ 才是裸函数) + if hasattr(fn, "__wrapped__"): + wrapped = fn.__wrapped__ + if hasattr(wrapped, "__func__"): + return wrapped.__func__ + return wrapped + # 已经是裸函数 + return fn + + +def _call_task(fn, bound_task, video_ids, user_id=""): + """调用 celery task 函数,自动适配 Task 对象和裸函数两种情况. + + 统一手动传 mock self,不依赖 celery 运行时注入。 + """ + raw_fn = _get_raw_task_fn(fn) + return raw_fn(bound_task, video_ids, user_id) + + def _run_with_fakes( videos, user_id="user_1", @@ -78,6 +114,8 @@ def _run_with_fakes( captured["upload_calls"].append((local_path, storage_key)) return upload_fn(local_path, storage_key) + bound_task = _make_bound_task() + with patch( "packages.adapters.sqlalchemy_impl.generated_video_repository.SQLAlchemyGeneratedVideoRepository", return_value=repo, @@ -94,7 +132,7 @@ def _run_with_fakes( "apps.worker.worker_app.tasks.batch_download._download_video_to_file", download_fn, ): - result = batch_download_videos([v.id for v in videos], user_id) + result = _call_task(batch_download_videos, bound_task, [v.id for v in videos], user_id) captured["result"] = result return captured @@ -124,6 +162,7 @@ def test_batch_download_no_videos_raises(): from apps.worker.worker_app.tasks.batch_download import batch_download_videos repo = _FakeGeneratedVideoRepository([]) + bound_task = _make_bound_task() with patch( "packages.adapters.sqlalchemy_impl.generated_video_repository.SQLAlchemyGeneratedVideoRepository", @@ -131,7 +170,7 @@ def test_batch_download_no_videos_raises(): ): with patch("worker_app.db.SessionLocal", MagicMock()): with pytest.raises(ValueError, match="No videos found"): - batch_download_videos(["nonexistent"], "user_1") + _call_task(batch_download_videos, bound_task, ["nonexistent"], "user_1") def test_batch_download_all_downloads_fail_raises(): @@ -226,6 +265,7 @@ def test_batch_download_closes_session_on_error(): session = MagicMock() session_maker = MagicMock(return_value=session) + bound_task = _make_bound_task() with patch( "packages.adapters.sqlalchemy_impl.generated_video_repository.SQLAlchemyGeneratedVideoRepository", @@ -233,7 +273,7 @@ def test_batch_download_closes_session_on_error(): ): with patch("worker_app.db.SessionLocal", session_maker): with pytest.raises(RuntimeError, match="db down"): - batch_download_videos(["v1"], "u") + _call_task(batch_download_videos, bound_task, ["v1"], "u") session.close.assert_called_once() @@ -378,8 +418,21 @@ def test_download_http_propagates_error(): def test_batch_download_task_name(): - """Task has correct name and retry settings.""" + """Task has correct name and retry settings. + + 通过源码断言装饰器参数来验证元数据,避免依赖 celery 运行时状态 + (CI 环境中 celery 可能未完整初始化)。 + """ + import inspect + from apps.worker.worker_app.tasks.batch_download import batch_download_videos - assert batch_download_videos.name == "worker.batch_download_videos" - assert batch_download_videos.max_retries == 1 + # 优先用 Task 对象的属性(本地/完整环境) + if hasattr(batch_download_videos, "name"): + assert batch_download_videos.name == "worker.batch_download_videos" + assert batch_download_videos.max_retries == 1 + else: + # 降级:检查源码中装饰器参数 + source = inspect.getsource(batch_download_videos) + assert 'name="worker.batch_download_videos"' in source or "name = 'worker.batch_download_videos'" in source + assert "max_retries=1" in source or "max_retries = 1" in source