fix(test): batch_download测试适配celery未初始化场景 #1198

Merged
xiaoxia merged 1 commits from fix/wave218-batch-download-ci-failure into develop 2026-07-30 13:23:45 +08:00
+59 -6
View File
@@ -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