diff --git a/apps/api/app/api/routes/viral_video.py b/apps/api/app/api/routes/viral_video.py index e9936203b..2ebae48ee 100644 --- a/apps/api/app/api/routes/viral_video.py +++ b/apps/api/app/api/routes/viral_video.py @@ -16,6 +16,7 @@ from __future__ import annotations import logging from app.auth import AuthenticatedUser, get_current_user +from app.core.celery_app import celery_app from app.dependencies import get_db_session from app.schemas.viral_video import ( AnalyzeStyleRequest, @@ -122,9 +123,7 @@ def create_viral_video( # 入队 Celery 任务 try: - from worker_app.tasks.viral_video import run_viral_video_pipeline - - run_viral_video_pipeline.delay(job.id) + celery_app.send_task("worker.run_viral_video_pipeline", args=[job.id]) logger.info("[爆款视频] 任务已入队: job_id=%s user_id=%s", job.id, job.user_id) except Exception as e: logger.error("[爆款视频] 入队失败: %s", e, exc_info=True) @@ -210,9 +209,7 @@ def retry_viral_video_job( # 重新入队 try: - from worker_app.tasks.viral_video import run_viral_video_pipeline - - run_viral_video_pipeline.delay(job.id) + celery_app.send_task("worker.run_viral_video_pipeline", args=[job.id]) logger.info("[爆款视频] 重试入队: job_id=%s retry_count=%d", job.id, job.retry_count) except Exception as e: logger.error("[爆款视频] 重试入队失败: %s", e, exc_info=True) @@ -249,9 +246,7 @@ def confirm_intent( # 从断点恢复 Celery 任务 try: - from worker_app.tasks.viral_video import resume_viral_video_pipeline - - resume_viral_video_pipeline.delay(job.id) + celery_app.send_task("worker.resume_viral_video_pipeline", args=[job.id]) logger.info("[爆款视频] 意图确认,恢复流水线: job_id=%s", job.id) except Exception as e: logger.error("[爆款视频] 恢复流水线失败: %s", e, exc_info=True) @@ -284,9 +279,7 @@ def analyze_style( # 入队风格分析任务 try: - from worker_app.tasks.viral_video import run_video_style_analysis - - run_video_style_analysis.delay(job.id) + celery_app.send_task("worker.run_video_style_analysis", args=[job.id]) logger.info("[爆款视频] 风格分析入队: job_id=%s", job.id) except Exception as e: logger.error("[爆款视频] 风格分析入队失败: %s", e, exc_info=True) diff --git a/tests/unit/test_viral_video_routes.py b/tests/unit/test_viral_video_routes.py new file mode 100644 index 000000000..02e986220 --- /dev/null +++ b/tests/unit/test_viral_video_routes.py @@ -0,0 +1,176 @@ +"""viral_video.py HTTP 端点单元测试(celery send_task 分支覆盖)。 + +直接调用路由函数(不启动 TestClient),通过 patch 注入 repo/session/user, +覆盖 4 个 celery_app.send_task(...) 调用点: + +- create_viral_video (generate) -> worker.run_viral_video_pipeline +- retry_viral_video_job (retry) -> worker.run_viral_video_pipeline +- confirm_intent -> worker.resume_viral_video_pipeline +- analyze_style -> worker.run_video_style_analysis +""" + +from __future__ import annotations + +from types import SimpleNamespace +from unittest.mock import MagicMock, patch + + +def _auth_user(uid: str = "u1"): + return SimpleNamespace(user=SimpleNamespace(id=uid)) + + +def _make_job(job_id: str = "job-1", user_id: str = "u1", status: str = "pending", **kwargs): + from packages.domain.viral_video import ViralVideoStatus + + job = MagicMock() + job.id = job_id + job.user_id = user_id + job.status = ViralVideoStatus(status) if isinstance(status, str) else status + job.images = kwargs.pop("images", ["img-1"]) + job.industry = kwargs.pop("industry", "电商") + job.target_customer = kwargs.pop("target_customer", "年轻人") + for k, v in { + "persona_id": "", + "viral_structure": "", + "marketing_purpose": "", + "bgm_preference": "", + "duration": 30, + "user_copy_text": "", + "fusion_level": "ai_polish", + "reference_audio_path": "", + "reference_video_url": "", + "style_strength": "medium", + "style_template_id": "", + "retry_count": 0, + "error_msg": "", + "result_video_url": "", + "style_guide": None, + "created_at": None, + "started_at": None, + "completed_at": None, + "stage": "", + "progress": 0.0, + "intent_result": None, + "updated_at": None, + }.items(): + setattr(job, k, kwargs.pop(k, v)) + return job + + +# ── generate ──────────────────────────────────────────────────────────── + + +class TestCreateViralVideo: + def _req(self, **kw): + from app.schemas.viral_video import CreateViralVideoRequest + + d = {"images": ["https://x.com/a.jpg"], "industry": "电商", "target_customer": "年轻人"} + d.update(kw) + return CreateViralVideoRequest(**d) + + def test_generate_dispatches_celery_task(self): + from app.api.routes import viral_video as vv_mod + + req = self._req() + user = _auth_user("u1") + session = MagicMock() + saved_job = _make_job(job_id="job-new", user_id="u1", status="pending") + repo = MagicMock() + + def fake_save(job): + job.id = saved_job.id + + repo.save.side_effect = fake_save + + with ( + patch.object(vv_mod, "_get_job_repo", return_value=repo), + patch.object(vv_mod.celery_app, "send_task") as mock_send, + ): + resp = vv_mod.create_viral_video(req, authenticated_user=user, session=session) + + mock_send.assert_called_once_with("worker.run_viral_video_pipeline", args=[saved_job.id]) + assert resp.id == saved_job.id + + +# ── retry ─────────────────────────────────────────────────────────────── + + +class TestRetryViralVideo: + def test_retry_dispatches_celery_task(self): + from app.api.routes import viral_video as vv_mod + + from packages.domain.viral_video import ViralVideoStatus + + user = _auth_user("u1") + session = MagicMock() + job = _make_job(job_id="job-retry", user_id="u1", status=ViralVideoStatus.FAILED, retry_count=1) + repo = MagicMock() + repo.get.return_value = job + + with ( + patch.object(vv_mod, "_get_job_repo", return_value=repo), + patch.object(vv_mod.celery_app, "send_task") as mock_send, + ): + resp = vv_mod.retry_viral_video_job("job-retry", authenticated_user=user, session=session) + + assert job.status == ViralVideoStatus.PENDING + assert job.retry_count == 2 + mock_send.assert_called_once_with("worker.run_viral_video_pipeline", args=["job-retry"]) + assert resp.id == "job-retry" + + +# ── confirm-intent ────────────────────────────────────────────────────── + + +class TestConfirmIntent: + def test_confirm_intent_dispatches_resume_task(self): + from app.api.routes import viral_video as vv_mod + from app.schemas.viral_video import ConfirmIntentRequest + + from packages.domain.viral_video import ViralVideoStatus + + user = _auth_user("u1") + session = MagicMock() + job = _make_job(job_id="job-cfm", user_id="u1", status=ViralVideoStatus.WAIT_USER_CONFIRM) + repo = MagicMock() + repo.get.return_value = job + req = ConfirmIntentRequest(confirmed_copy="确认后的文案") + + with ( + patch.object(vv_mod, "_get_job_repo", return_value=repo), + patch.object(vv_mod.celery_app, "send_task") as mock_send, + ): + resp = vv_mod.confirm_intent("job-cfm", req, authenticated_user=user, session=session) + + assert job.user_copy_text == "确认后的文案" + job.resume_from_confirm.assert_called_once() + mock_send.assert_called_once_with("worker.resume_viral_video_pipeline", args=["job-cfm"]) + assert resp.id == "job-cfm" + + +# ── analyze-style ─────────────────────────────────────────────────────── + + +class TestAnalyzeStyle: + def test_analyze_style_dispatches_analysis_task(self): + from app.api.routes import viral_video as vv_mod + from app.schemas.viral_video import AnalyzeStyleRequest + + user = _auth_user("u1") + session = MagicMock() + job = _make_job(job_id="job-sty", user_id="u1", status="pending") + repo = MagicMock() + repo.get.return_value = job + req = AnalyzeStyleRequest(reference_video_url="https://x.com/ref.mp4", style_template_id="tpl-1") + + with ( + patch.object(vv_mod, "_get_job_repo", return_value=repo), + patch.object(vv_mod.celery_app, "send_task") as mock_send, + ): + resp = vv_mod.analyze_style("job-sty", req, authenticated_user=user, session=session) + + assert job.reference_video_url == "https://x.com/ref.mp4" + assert job.style_template_id == "tpl-1" + mock_send.assert_called_once_with("worker.run_video_style_analysis", args=["job-sty"]) + assert resp.job_id == "job-sty" + assert resp.status == "analyzing"