diff --git a/apps/api/app/api/routes/viral_video.py b/apps/api/app/api/routes/viral_video.py index 93468c46b..bf2bba725 100644 --- a/apps/api/app/api/routes/viral_video.py +++ b/apps/api/app/api/routes/viral_video.py @@ -280,11 +280,18 @@ def generate_copy( raise HTTPException(status_code=404, detail="任务不存在") if job.user_id != authenticated_user.user.id: raise HTTPException(status_code=403, detail="无权操作此任务") - if job.status not in (ViralVideoStatus.IMAGE_ANALYZED, ViralVideoStatus.PENDING, ViralVideoStatus.FAILED): + # 允许首次进入(IMAGE_ANALYZED/PENDING)、失败重试(FAILED)、文案重新生成(COPY_GENERATED/COMPLETED) + if job.status not in ( + ViralVideoStatus.IMAGE_ANALYZED, + ViralVideoStatus.PENDING, + ViralVideoStatus.FAILED, + ViralVideoStatus.COPY_GENERATED, + ViralVideoStatus.COMPLETED, + ): raise HTTPException(status_code=409, detail=f"任务当前状态 {job.status} 不能生成文案") - # 允许失败任务重试:重置 - if job.status == ViralVideoStatus.FAILED: + # 失败重试 / 重新生成:retry_count 自增 + if job.status in (ViralVideoStatus.FAILED, ViralVideoStatus.COPY_GENERATED, ViralVideoStatus.COMPLETED): job.retry_count += 1 job.error_msg = "" diff --git a/packages/domain/viral_video.py b/packages/domain/viral_video.py index 09119a00a..760864539 100755 --- a/packages/domain/viral_video.py +++ b/packages/domain/viral_video.py @@ -191,11 +191,40 @@ class ViralVideoJob: self.updated_at = datetime.now(timezone.utc) def resume_from_image_analyzed(self, **kwargs) -> None: - if self.status not in (ViralVideoStatus.IMAGE_ANALYZED, ViralVideoStatus.PENDING): + """阶段2入口:允许从 IMAGE_ANALYZED/PENDING 首次进入,也允许从 COPY_GENERATED/COMPLETED/FAILED 重新生成文案。 + + 重新生成时清空上一轮文案产物(copy_result/intent_result/storyboard/generated_copy_text), + 并重置 completed_at/result_video_url/error_msg,确保前端轮询能看到新的阶段2进度。 + """ + _allowed = ( + ViralVideoStatus.IMAGE_ANALYZED, + ViralVideoStatus.PENDING, + ViralVideoStatus.COPY_GENERATED, + ViralVideoStatus.COMPLETED, + ViralVideoStatus.FAILED, + ) + if self.status not in _allowed: raise ValueError(f"Cannot resume from {self.status} to copy-gen") + _is_regen = self.status in ( + ViralVideoStatus.COPY_GENERATED, + ViralVideoStatus.COMPLETED, + ViralVideoStatus.FAILED, + ) for k, v in kwargs.items(): if hasattr(self, k) and v not in (None, "", []): setattr(self, k, v) + if _is_regen: + # 清空上一轮文案/视频产物,避免前端拿到旧数据 + self.intent_result = None + self.copy_result = None + self.storyboard = None + self.generated_copy_text = "" + self.result_video_url = "" + self.current_stage = "" + self.phase_message = "" + self.error_msg = "" + self.completed_at = None + self.heartbeat_at = None self.status = ViralVideoStatus.RUNNING self.updated_at = datetime.now(timezone.utc) diff --git a/tests/unit/test_domain_entities_extended.py b/tests/unit/test_domain_entities_extended.py index d638320fa..c7b97db19 100755 --- a/tests/unit/test_domain_entities_extended.py +++ b/tests/unit/test_domain_entities_extended.py @@ -521,3 +521,81 @@ class TestIngestJob: storage_key="k", ) assert job.error_message == "" + + +class TestViralVideoResumeForRegenerate: + """#2222: resume_from_image_analyzed 应支持 COPY_GENERATED/COMPLETED/FAILED 重新生成文案。""" + + def test_regen_from_copy_generated_clears_old_copy(self): + from datetime import datetime, timezone + + from packages.domain.viral_video import ViralVideoJob, ViralVideoStatus + + job = ViralVideoJob(user_id="u1", images=["img1"]) + # 模拟已经生成过文案和视频 + job.status = ViralVideoStatus.COPY_GENERATED + job.copy_result = {"shots": [{"x": 1}], "voiceover_script": "旧文案"} + job.intent_result = {"intent": "旧意图"} + job.storyboard = [{"x": 1}] + job.generated_copy_text = "旧文案" + job.result_video_url = "http://old.mp4" + job.completed_at = datetime(2026, 10, 6, tzinfo=timezone.utc) + job.error_msg = "" + job.current_stage = "tts_generation" + job.phase_message = "TTS完成" + + # 重新生成 + job.resume_from_image_analyzed() + + assert job.status == ViralVideoStatus.RUNNING + assert job.copy_result is None + assert job.intent_result is None + assert job.storyboard is None + assert job.generated_copy_text == "" + assert job.result_video_url == "" + assert job.completed_at is None + assert job.error_msg == "" + assert job.current_stage == "" + assert job.phase_message == "" + + def test_regen_from_completed_clears_old_copy(self): + from packages.domain.viral_video import ViralVideoJob, ViralVideoStatus + + job = ViralVideoJob(user_id="u1", images=["img1"]) + job.status = ViralVideoStatus.COMPLETED + job.copy_result = {"shots": [], "voiceover_script": "xx"} + job.intent_result = {"intent": "x"} + job.result_video_url = "http://v.mp4" + + job.resume_from_image_analyzed() + + assert job.status == ViralVideoStatus.RUNNING + assert job.copy_result is None + assert job.intent_result is None + assert job.result_video_url == "" + + def test_first_call_from_image_analyzed_keeps_fields(self): + """首次进入(IMAGE_ANALYZED)不应清空任何已有的字段。""" + from packages.domain.viral_video import ViralVideoJob, ViralVideoStatus + + job = ViralVideoJob(user_id="u1", images=["img1"]) + job.status = ViralVideoStatus.IMAGE_ANALYZED + job.image_analysis = {"products": []} + job.industry = "美妆" + + job.resume_from_image_analyzed() + + assert job.status == ViralVideoStatus.RUNNING + assert job.image_analysis == {"products": []} + assert job.industry == "美妆" + + def test_wait_user_confirm_rejected(self): + """wait_user_confirm 中间状态应被拒绝(前端正在编辑/确认文案)。""" + import pytest + + from packages.domain.viral_video import ViralVideoJob, ViralVideoStatus + + job = ViralVideoJob(user_id="u1", images=["img1"]) + job.status = ViralVideoStatus.WAIT_USER_CONFIRM + with pytest.raises(ValueError, match="Cannot resume"): + job.resume_from_image_analyzed() diff --git a/tests/unit/test_viral_video_routes.py b/tests/unit/test_viral_video_routes.py index d81c5e841..e586a642e 100644 --- a/tests/unit/test_viral_video_routes.py +++ b/tests/unit/test_viral_video_routes.py @@ -431,7 +431,7 @@ class TestGenerateCopy: assert resp.id == "job-gc" def test_generate_copy_rejects_wrong_status(self): - """任务在 copy_generated/completed 时不能再 generate-copy(状态保护)。""" + """wait_user_confirm 等中间状态不允许调用 generate-copy(状态保护)。""" import pytest from app.api.routes import viral_video as vv_mod from app.schemas.viral_video import GenerateCopyRequest @@ -441,7 +441,8 @@ class TestGenerateCopy: user = _auth_user("u1") session = MagicMock() - job = _make_job(job_id="job-gc2", user_id="u1", status=ViralVideoStatus.COPY_GENERATED) + # wait_user_confirm 属于前端在编辑/确认文案的中间状态,应拒绝重新触发生成 + job = _make_job(job_id="job-gc2", user_id="u1", status=ViralVideoStatus.WAIT_USER_CONFIRM) repo = MagicMock() repo.get.return_value = job @@ -450,6 +451,31 @@ class TestGenerateCopy: vv_mod.generate_copy("job-gc2", GenerateCopyRequest(), authenticated_user=user, session=session) assert exc.value.status_code == 409 + def test_generate_copy_allows_regenerate_from_copy_generated(self): + """#2222: COPY_GENERATED/COMPLETED 状态下点「重新生成文案」应放行入队,不返回 409。""" + from unittest.mock import patch + + from app.api.routes import viral_video as vv_mod + from app.schemas.viral_video import GenerateCopyRequest + + from packages.domain.viral_video import ViralVideoStatus + + user = _auth_user("u1") + session = MagicMock() + for regen_status in (ViralVideoStatus.COPY_GENERATED, ViralVideoStatus.COMPLETED): + job = _make_job(job_id=f"job-regen-{regen_status}", user_id="u1", status=regen_status) + 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.generate_copy(f"job-regen-{regen_status}", GenerateCopyRequest(), authenticated_user=user, session=session) + mock_send.assert_called_once() + job.resume_from_image_analyzed.assert_called() + assert job.retry_count >= 1 + assert resp.id == f"job-regen-{regen_status}" + def test_generate_copy_persists_voice_and_ratio(self): """generate-copy 应把 voice_id/voice_source/video_ratio 写入 job。""" from unittest.mock import patch