fix(api): 爆款视频 API 用 celery_app.send_task() 替换直接 import worker #2105

Merged
xiaoxia merged 4 commits from fix/2105-viral-video-celery-import into develop 2026-09-30 16:22:59 +08:00
2 changed files with 181 additions and 12 deletions
+5 -12
View File
@@ -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)
+176
View File
@@ -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"