from django.test import TestCase from unittest.mock import patch from rest_framework.test import APIClient from apps.accounts.models import Team, TeamMember, User from apps.ai.models import ModelConfig, ModelProvider from apps.assets.models import Asset from apps.billing.models import CreditAccount, CreditLedger from apps.products.models import Product from apps.projects.models import ( Project, ProjectStage, ScriptSegment, ScriptVersion, SubtitleTrack, Timeline, TimelineClip, VideoSegment, VideoSegmentVersion, ) class ProjectApiTests(TestCase): def setUp(self): self.user = User.objects.create_user(username="owner", password="pass") self.team = Team.objects.create(name="E2E Team", owner=self.user) TeamMember.objects.create(team=self.team, user=self.user, role=TeamMember.Role.OWNER) CreditAccount.objects.create(team=self.team, balance="100.0000") self.product = Product.objects.create(team=self.team, created_by=self.user, title="Test Product") self.provider = ModelProvider.objects.create( name="volcengine", display_name="Volcano", base_url="https://ark.cn-beijing.volces.com/api/v3", ) self.model = ModelConfig.objects.create( provider=self.provider, name="doubao-seed-2-0-pro-260215", display_name="Doubao", capability=ModelConfig.Capability.TEXT, endpoint="chat/completions", unit_price="2.0000", ) # 数据迁移(0006)会 seed 中转站文本模型,created_at 更早 → get_default_model 会选到它, # 从而路由到未被 patch 的 OpenAICompatibleProvider 打真网络。这里只保留自建 volcengine 模型为 # active,确保命中可 mock 的 VolcanoArkProvider(生产有 bootstrap 豆包在前,默认仍是豆包,不受影响)。 ModelConfig.objects.filter(capability=ModelConfig.Capability.TEXT).exclude(id=self.model.id).update(status="disabled") self.client = APIClient() self.client.force_authenticate(self.user) def test_create_project_initializes_pipeline(self): response = self.client.post( "/api/projects/", {"name": "Launch Video", "product": str(self.product.id)}, format="json", ) self.assertEqual(response.status_code, 201) project = Project.objects.get(id=response.data["id"]) self.assertEqual(project.team, self.team) self.assertEqual(project.created_by, self.user) self.assertEqual(project.stages.count(), 5) self.assertEqual(project.video_segments.count(), 4) self.assertEqual( list(project.stages.values_list("stage", flat=True)), [ ProjectStage.Stage.SCRIPT, ProjectStage.Stage.BASE_ASSETS, ProjectStage.Stage.STORYBOARD, ProjectStage.Stage.VIDEO, ProjectStage.Stage.EXPORT, ], ) self.assertEqual( list(project.video_segments.values_list("target_duration_seconds", flat=True)), [15, 15, 15, 15], ) self.assertEqual( list(project.video_segments.values_list("sort_order", flat=True)), [0, 1, 2, 3], ) self.assertTrue(VideoSegment.objects.filter(project=project).exists()) @patch("apps.ai.services.VolcanoArkProvider") def test_generate_script_creates_script_segments_and_charges_credit(self, provider_cls): provider = provider_cls.return_value provider.chat_completion.return_value = { "choices": [ { "message": { "content": "1. 开场吸引\n2. 展示卖点\n3. 使用场景\n4. 促单转化", } } ] } provider.extract_text.return_value = "1. 开场吸引\n2. 展示卖点\n3. 使用场景\n4. 促单转化" project = Project.objects.create(team=self.team, created_by=self.user, product=self.product, name="Launch Video") for stage in [ ProjectStage.Stage.SCRIPT, ProjectStage.Stage.BASE_ASSETS, ProjectStage.Stage.STORYBOARD, ProjectStage.Stage.VIDEO, ProjectStage.Stage.EXPORT, ]: ProjectStage.objects.create(project=project, stage=stage) response = self.client.post( f"/api/projects/{project.id}/generate-script/", {"prompt": "突出高转化"}, format="json", ) self.assertEqual(response.status_code, 201) script = ScriptVersion.objects.get(project=project) self.assertEqual(script.segments.count(), 4) # 出稿(1)+ 人物/场景提取(1)各计费一次;此处 mock 提取返回非 JSON,提取仍按调用成功计费 self.assertEqual(CreditLedger.objects.filter(team=self.team, ledger_type=CreditLedger.Type.CHARGE).count(), 2) @patch("apps.ai.services.VolcanoArkProvider") def test_rerun_script_segment_updates_one_segment_and_charges_once(self, provider_cls): provider = provider_cls.return_value provider.chat_completion.return_value = {"choices": [{"message": {"content": "x"}}]} provider.extract_text.return_value = "旁白:全新口播文案\n画面:全新画面描述" project = Project.objects.create(team=self.team, created_by=self.user, product=self.product, name="P") script = ScriptVersion.objects.create(project=project, title="脚本", content="...", is_adopted=True) seg0 = ScriptSegment.objects.create(script_version=script, sort_order=0, narration="旧0", visual_prompt="画0") seg1 = ScriptSegment.objects.create(script_version=script, sort_order=1, narration="旧1", visual_prompt="画1") response = self.client.post( f"/api/projects/{project.id}/rerun-script-segment/", {"segment_id": str(seg1.id), "instruction": "更俏皮"}, format="json", ) self.assertEqual(response.status_code, 200) seg0.refresh_from_db() seg1.refresh_from_db() self.assertEqual(seg1.narration, "全新口播文案") self.assertEqual(seg1.visual_prompt, "全新画面描述") # 其它分镜保持不动 self.assertEqual(seg0.narration, "旧0") self.assertEqual(seg0.visual_prompt, "画0") # 返回的是该镜所属 ScriptVersion self.assertEqual(str(response.data["id"]), str(script.id)) self.assertEqual(CreditLedger.objects.filter(team=self.team, ledger_type=CreditLedger.Type.CHARGE).count(), 1) def test_rerun_script_segment_rejects_foreign_segment(self): project = Project.objects.create(team=self.team, created_by=self.user, product=self.product, name="P") other = Project.objects.create(team=self.team, created_by=self.user, product=self.product, name="O") other_script = ScriptVersion.objects.create(project=other, title="脚本", content="...") other_seg = ScriptSegment.objects.create(script_version=other_script, sort_order=0, narration="x") response = self.client.post( f"/api/projects/{project.id}/rerun-script-segment/", {"segment_id": str(other_seg.id)}, format="json", ) self.assertEqual(response.status_code, 404) # ── 流程文档 步骤③④:出稿后自动提取人物/场景(含建议提示词)并计费 ── @patch("apps.ai.services.VolcanoArkProvider") def test_generate_script_extracts_cast_scenes_and_bills(self, provider_cls): provider = provider_cls.return_value script_text = "镜头1\n旁白:开场\n画面:地铁里女主掏出面膜\n\n镜头2\n旁白:卖点\n画面:同事围观" # 提取调用让模型回 JSON(允许外面裹解释文字,验证容错) extract_text = '这是结果:{"cast":[{"name":"女主","prompt":"26岁都市女性,自然妆"},{"name":"同事","prompt":"职场女性"}],"scenes":[{"name":"地铁","prompt":"早高峰地铁车厢"}]}' provider.chat_completion.side_effect = [{"r": 1}, {"r": 2}] provider.extract_text.side_effect = [script_text, extract_text] project = Project.objects.create(team=self.team, created_by=self.user, product=self.product, name="P") ProjectStage.objects.create(project=project, stage=ProjectStage.Stage.SCRIPT) response = self.client.post(f"/api/projects/{project.id}/generate-script/", {"prompt": "x"}, format="json") self.assertEqual(response.status_code, 201) project.refresh_from_db() self.assertEqual(project.metadata.get("cast"), ["女主", "同事"]) self.assertEqual(project.metadata.get("scenes"), ["地铁"]) self.assertEqual(project.metadata.get("cast_prompts", {}).get("女主"), "26岁都市女性,自然妆") self.assertEqual(project.metadata.get("scene_prompts", {}).get("地铁"), "早高峰地铁车厢") # 出稿 + 提取,各计费一次 self.assertEqual(CreditLedger.objects.filter(team=self.team, ledger_type=CreditLedger.Type.CHARGE).count(), 2) @patch("apps.ai.services.VolcanoArkProvider") def test_extract_failure_releases_credit_and_keeps_script(self, provider_cls): """提取调用抛错:必须退掉提取的预扣(不计费),且不影响脚本已出稿(出稿那笔照扣)。""" provider = provider_cls.return_value provider.chat_completion.side_effect = [{"r": 1}, RuntimeError("ark down")] provider.extract_text.side_effect = ["镜头1\n旁白:x\n画面:y", AssertionError("不应走到")] project = Project.objects.create(team=self.team, created_by=self.user, product=self.product, name="P") ProjectStage.objects.create(project=project, stage=ProjectStage.Stage.SCRIPT) response = self.client.post(f"/api/projects/{project.id}/generate-script/", {"prompt": "x"}, format="json") self.assertEqual(response.status_code, 201) self.assertEqual(ScriptVersion.objects.filter(project=project).count(), 1) # 出稿计费 1 笔;提取失败 → release,不产生第 2 笔 CHARGE self.assertEqual(CreditLedger.objects.filter(team=self.team, ledger_type=CreditLedger.Type.CHARGE).count(), 1) @patch("apps.ai.services._store_generated_media") @patch("apps.ai.services.get_image_provider") def test_generate_base_asset_stores_tag_label(self, get_provider, store_media): """流程步骤④:按脚本标签生成基础资产时,label 落进 group.metadata,供前端把结果归回标签卡。""" ModelConfig.objects.create( provider=self.provider, name="img-model", display_name="Img", capability=ModelConfig.Capability.IMAGE, endpoint="images/generations", unit_price="1.0000", ) asset = Asset.objects.create( team=self.team, created_by=self.user, name="女主", asset_type=Asset.Type.IMAGE, source=Asset.Source.AI_GENERATED, category=Asset.Category.PERSON, ) store_media.return_value = asset provider = get_provider.return_value provider.image_generation.return_value = {"data": [{"url": "http://x/img.png"}]} provider.extract_first_media_url.return_value = "http://x/img.png" project = Project.objects.create(team=self.team, created_by=self.user, product=self.product, name="P") response = self.client.post( f"/api/projects/{project.id}/generate-base-asset/", {"kind": "person", "prompt": "26岁都市女性", "label": "女主"}, format="json", ) # 异步:秒回 202 + 任务;EAGER 下 worker 任务就地执行,组已落库,label 进 metadata self.assertEqual(response.status_code, 202) self.assertIn("task", response.data) from apps.projects.models import BaseAssetGroup group = BaseAssetGroup.objects.get(project=project, kind=BaseAssetGroup.Kind.PERSON) self.assertEqual(group.metadata.get("label"), "女主") @patch("apps.ai.services._store_generated_media") @patch("apps.ai.services.get_image_provider") def test_product_base_asset_uses_cover_image_as_reference(self, get_provider, store_media): """商品三视图:商品有主图时走 image_edit 以主图为参考(锁包装一致),而非纯文生图。""" from apps.assets.models import AssetFile ModelConfig.objects.create( provider=self.provider, name="img-model", display_name="Img", capability=ModelConfig.Capability.IMAGE, endpoint="images/generations", unit_price="1.0000", ) # 商品主图资产(带可访问 preview_url) cover = Asset.objects.create( team=self.team, created_by=self.user, name="主图", asset_type=Asset.Type.IMAGE, source=Asset.Source.UPLOAD, category=Asset.Category.PRODUCT_IMAGE, ) AssetFile.objects.create(asset=cover, object_key="k.png", bucket="b", content_type="image/png", preview_url="http://x/cover.png", is_primary=True) self.product.cover_asset = cover self.product.save(update_fields=["cover_asset"]) out = Asset.objects.create( team=self.team, created_by=self.user, name="商品三视图", asset_type=Asset.Type.IMAGE, source=Asset.Source.AI_GENERATED, category=Asset.Category.PRODUCT_IMAGE, ) store_media.return_value = out provider = get_provider.return_value provider.image_edit.return_value = {"data": [{"url": "http://x/tri.png"}]} provider.extract_first_media_url.return_value = "http://x/tri.png" project = Project.objects.create(team=self.team, created_by=self.user, product=self.product, name="P") response = self.client.post( f"/api/projects/{project.id}/generate-base-asset/", {"kind": "product", "prompt": "商品三视图"}, format="json", ) self.assertEqual(response.status_code, 202) # 异步:秒回任务,EAGER 下 worker 就地出图 # 走 image_edit 且把商品主图 URL 作为参考图传入,未走纯文生图 provider.image_edit.assert_called_once() self.assertEqual(provider.image_edit.call_args.kwargs["images"], ["http://x/cover.png"]) provider.image_generation.assert_not_called() @patch("apps.ai.services._store_generated_media") @patch("apps.ai.services.get_image_provider") def test_product_base_asset_falls_back_without_cover(self, get_provider, store_media): """商品无主图时回落纯文生图(image_generation),不报错。""" ModelConfig.objects.create( provider=self.provider, name="img-model", display_name="Img", capability=ModelConfig.Capability.IMAGE, endpoint="images/generations", unit_price="1.0000", ) out = Asset.objects.create( team=self.team, created_by=self.user, name="商品三视图", asset_type=Asset.Type.IMAGE, source=Asset.Source.AI_GENERATED, category=Asset.Category.PRODUCT_IMAGE, ) store_media.return_value = out provider = get_provider.return_value provider.image_generation.return_value = {"data": [{"url": "http://x/tri.png"}]} provider.extract_first_media_url.return_value = "http://x/tri.png" project = Project.objects.create(team=self.team, created_by=self.user, product=self.product, name="P") response = self.client.post( f"/api/projects/{project.id}/generate-base-asset/", {"kind": "product", "prompt": "商品三视图"}, format="json", ) self.assertEqual(response.status_code, 202) # 异步:秒回任务,EAGER 下 worker 就地出图 provider.image_generation.assert_called_once() provider.image_edit.assert_not_called() @patch("apps.ai.services._store_generated_media") @patch("apps.ai.services.get_image_provider") def test_triview_async_binds_group_to_portrait(self, get_provider, store_media): """三视图异步:提交秒回 202;EAGER 下 worker 就地以立绘为参考 image_edit,生成的组 triview_of 绑定该立绘。""" from apps.assets.models import AssetFile from apps.projects.models import BaseAssetGroup ModelConfig.objects.create( provider=self.provider, name="img-model", display_name="Img", capability=ModelConfig.Capability.IMAGE, endpoint="images/generations", unit_price="1.0000", ) portrait = Asset.objects.create( team=self.team, created_by=self.user, name="立绘", asset_type=Asset.Type.IMAGE, source=Asset.Source.AI_GENERATED, category=Asset.Category.PERSON, ) AssetFile.objects.create(asset=portrait, object_key="p.png", bucket="b", content_type="image/png", preview_url="http://x/portrait.png", is_primary=True) tri = Asset.objects.create( team=self.team, created_by=self.user, name="三视图", asset_type=Asset.Type.IMAGE, source=Asset.Source.AI_GENERATED, category=Asset.Category.PERSON, ) store_media.return_value = tri provider = get_provider.return_value provider.image_edit.return_value = {"data": [{"url": "http://x/tri.png"}]} provider.extract_first_media_url.return_value = "http://x/tri.png" project = Project.objects.create(team=self.team, created_by=self.user, product=self.product, name="P") response = self.client.post( f"/api/projects/{project.id}/generate-triview/", {"portrait_asset_id": str(portrait.id)}, format="json", ) self.assertEqual(response.status_code, 202) provider.image_edit.assert_called_once() self.assertEqual(provider.image_edit.call_args.kwargs["images"], ["http://x/portrait.png"]) group = BaseAssetGroup.objects.get(project=project, kind=BaseAssetGroup.Kind.PERSON) self.assertEqual(group.metadata.get("triview_of"), str(portrait.id)) self.assertEqual(group.adopted_asset_id, tri.id) def test_extract_cast_and_scenes_parsing_is_robust(self): """纯函数:_coerce_tag_entries 去重/容错;无可用模型时 extract 返回空且不抛。""" from apps.ai.services import _coerce_tag_entries, extract_cast_and_scenes names, prompts = _coerce_tag_entries( [{"name": "女主", "prompt": "p1"}, {"name": "女主", "prompt": "dup"}, {"name": "无"}, "同事"] ) self.assertEqual(names, ["女主", "同事"]) # 去重 + 过滤占位「无」+ 容忍纯字符串项 self.assertEqual(prompts["女主"], "p1") ModelConfig.objects.filter(capability=ModelConfig.Capability.TEXT).update(status=ModelConfig.Status.DISABLED) project = Project.objects.create(team=self.team, created_by=self.user, product=self.product, name="P") self.assertEqual( extract_cast_and_scenes(project=project, user=self.user, content="脚本"), {"cast": [], "scenes": [], "cast_prompts": {}, "scene_prompts": {}}, ) def test_adopt_video_version_remaps_timeline_draft(self): """切换采用版本后,时间线草稿里引用本场旧版本资产的片段必须跟随换成新资产(剪辑台/导出都读草稿)。""" project = Project.objects.create(team=self.team, created_by=self.user, product=self.product, name="P") make_asset = lambda name: Asset.objects.create( team=self.team, created_by=self.user, name=name, asset_type=Asset.Type.VIDEO, source=Asset.Source.AI_GENERATED, category=Asset.Category.VIDEO_CLIP, ) old_asset, new_asset = make_asset("v1"), make_asset("v2") segment = VideoSegment.objects.create(project=project, sort_order=0) old_ver = VideoSegmentVersion.objects.create(video_segment=segment, asset=old_asset, is_adopted=True) new_ver = VideoSegmentVersion.objects.create(video_segment=segment, asset=new_asset) segment.adopted_version = old_ver segment.save(update_fields=["adopted_version"]) timeline = Timeline.objects.create(project=project) clip = TimelineClip.objects.create( timeline=timeline, asset=old_asset, sort_order=0, duration_ms=15000, trim_start_ms=1000, trim_end_ms=9000, ) response = self.client.post( f"/api/projects/{project.id}/adopt-video-version/", {"video_segment_id": str(segment.id), "version_id": str(new_ver.id)}, format="json", ) self.assertEqual(response.status_code, 200) clip.refresh_from_db() self.assertEqual(clip.asset_id, new_asset.id) # 旧素材上的裁剪点对新素材无意义,必须复位 self.assertEqual(clip.trim_start_ms, 0) self.assertIsNone(clip.trim_end_ms) def _make_audio_model(self): return ModelConfig.objects.create( provider=self.provider, name="doubao-voice-bigtts", display_name="豆包语音合成", capability="audio", endpoint="api/v1/tts", unit_price="2.0000", ) @patch("apps.ai.services.TosStorage") @patch("apps.ai.services.VolcanoTtsProvider") def test_generate_voiceover_creates_assets_and_charges_once(self, provider_cls, storage_cls): """每镜旁白 TTS:产出音频资产、timeline.metadata 落映射、一次调用只计一次费。""" self._make_audio_model() provider = provider_cls.return_value provider.configured = True provider.synthesize.return_value = (b"fake-mp3-bytes", 4200) stored = storage_cls.return_value.upload_fileobj.return_value stored.object_key, stored.bucket, stored.content_type, stored.size_bytes = "k.mp3", "b", "audio/mpeg", 14 project = Project.objects.create(team=self.team, created_by=self.user, product=self.product, name="VO") response = self.client.post( f"/api/projects/{project.id}/generate-voiceover/", {"items": [{"index": 0, "text": "第一镜旁白"}, {"index": 1, "text": "第二镜旁白"}]}, format="json", ) self.assertEqual(response.status_code, 201) voiceover = response.json()["voiceover"] self.assertTrue(voiceover["enabled"]) self.assertEqual(len(voiceover["items"]), 2) self.assertEqual(provider.synthesize.call_count, 2) from apps.assets.models import Asset as AssetModel self.assertEqual(AssetModel.objects.filter(team=self.team, asset_type="audio").count(), 2) self.assertEqual( CreditLedger.objects.filter(team=self.team, ledger_type=CreditLedger.Type.CHARGE).count(), 1 ) project.refresh_from_db() self.assertEqual(len(project.timeline.metadata["voiceover"]["items"]), 2) @patch("apps.ai.services.VolcanoTtsProvider") def test_generate_voiceover_unconfigured_returns_400_and_no_charge(self, provider_cls): """语音凭证未配置:400 + 人话提示,不建任务不扣费。""" self._make_audio_model() provider_cls.return_value.configured = False project = Project.objects.create(team=self.team, created_by=self.user, product=self.product, name="VO2") response = self.client.post( f"/api/projects/{project.id}/generate-voiceover/", {"items": [{"index": 0, "text": "旁白"}]}, format="json", ) self.assertEqual(response.status_code, 400) self.assertIn("语音合成未配置", response.json()["detail"]) self.assertEqual(CreditLedger.objects.filter(team=self.team, ledger_type=CreditLedger.Type.CHARGE).count(), 0) def test_export_command_mixes_voiceover_above_bgm(self): """导出命令:人声按片段起点 adelay 延迟、先与原声 amix、再与 BGM amix。""" from apps.projects.services.export import _build_export_command specs = [{"ts": 0.0, "te": 15.0, "dur": 15.0}, {"ts": 0.0, "te": 15.0, "dur": 15.0}] cmd = _build_export_command( n=2, specs=specs, starts=[0.0, 15.0], total=30.0, transition="none", sub_overlays=[], bgm_name="bgm.mp3", bgm_volume=0.6, has_audio=[False, False], voice_overlays=[("vo0.mp3", 0.0, 15.0), ("vo1.mp3", 15.0, 15.0)], ) self.assertIn("vo0.mp3", cmd) self.assertIn("vo1.mp3", cmd) graph = cmd[cmd.index("-filter_complex") + 1] self.assertIn("adelay=15000|15000", graph) self.assertIn("[avoice][vo0][vo1]amix=inputs=3", graph) self.assertIn("[anarr][abgm]amix=inputs=2", graph) def test_save_timeline_updates_voiceover_offsets(self): """拖动字幕块后保存:按 asset 回写句内起点 offset_ms,未拖动的句不受影响。""" project = Project.objects.create(team=self.team, created_by=self.user, product=self.product, name="VOFF") Timeline.objects.create(project=project, metadata={"voiceover": {"enabled": True, "items": [ {"index": 0, "cue": 0, "asset": "a1", "text": "x", "duration_ms": 2000, "offset_ms": 0}, {"index": 0, "cue": 1, "asset": "a2", "text": "y", "duration_ms": 2000, "offset_ms": 2000}, ]}}) response = self.client.post( f"/api/projects/{project.id}/save-timeline/", {"voiceover": {"items": [{"asset": "a2", "offset_ms": 9000}]}}, format="json", ) self.assertEqual(response.status_code, 200) project.timeline.refresh_from_db() items = project.timeline.metadata["voiceover"]["items"] self.assertEqual(items[0]["offset_ms"], 0) self.assertEqual(items[1]["offset_ms"], 9000) def test_subtitle_cues_use_end_ms_and_follow_voiceover(self): """字幕烧入时序:cue 自带 end_ms 时按它收尾(不再拖到片段尾)——这是配音对齐的关键; 无保存字幕的回退路径,有配音时逐句窗口也要压进语音真实时长。""" from apps.projects.services.export import _subtitle_cues project = Project.objects.create(team=self.team, created_by=self.user, product=self.product, name="SUB") timeline = Timeline.objects.create( project=project, metadata={"voiceover": {"enabled": True, "items": [{"index": 0, "asset": "x", "text": "t", "duration_ms": 4000}]}}, ) SubtitleTrack.objects.create(timeline=timeline, enabled=True, content=[ {"start_ms": 0, "end_ms": 2000, "text": "第一句"}, {"start_ms": 2000, "end_ms": 4000, "text": "第二句"}, ]) specs = [{"ts": 0.0, "te": 15.0, "dur": 15.0}] cues = _subtitle_cues(timeline, project, specs, [0.0], 15.0) self.assertEqual(len(cues), 2) self.assertEqual((cues[0][0], cues[0][1]), (0.0, 2.0)) self.assertEqual(cues[1][1], 4.0) # 最后一句在语音念完处收尾,而非拖到 15s def test_save_timeline_voiceover_toggle_and_clear(self): """保存草稿可关/开配音、可移除映射(不动已生成的音频资产)。""" project = Project.objects.create(team=self.team, created_by=self.user, product=self.product, name="VO3") Timeline.objects.create( project=project, metadata={"voiceover": {"enabled": True, "voice_type": "v", "items": [{"index": 0, "asset": "a", "text": "t"}]}}, ) response = self.client.post( f"/api/projects/{project.id}/save-timeline/", {"voiceover": {"enabled": False}}, format="json" ) self.assertEqual(response.status_code, 200) project.timeline.refresh_from_db() self.assertFalse(project.timeline.metadata["voiceover"]["enabled"]) response = self.client.post( f"/api/projects/{project.id}/save-timeline/", {"voiceover": {"clear": True}}, format="json" ) self.assertEqual(response.status_code, 200) project.timeline.refresh_from_db() self.assertNotIn("voiceover", project.timeline.metadata) class WorkerGateTests(TestCase): """无 Celery worker 时,图片/视频生成入口必须 503 拒绝(防结果无人取回=数据丢失)。""" def setUp(self): self.user = User.objects.create_user(username="gate", password="pass") self.team = Team.objects.create(name="Gate Team", owner=self.user) TeamMember.objects.create(team=self.team, user=self.user, role=TeamMember.Role.OWNER) CreditAccount.objects.create(team=self.team, balance="100.0000") self.product = Product.objects.create(team=self.team, created_by=self.user, title="Gate Product") self.project = Project.objects.create(team=self.team, created_by=self.user, name="Gate", product=self.product) self.client = APIClient() self.client.force_authenticate(self.user) def _blocked(self): return patch("apps.common.celery_health.celery_worker_available", return_value=False) def test_generation_endpoints_blocked_without_worker(self): cases = [ (f"/api/projects/{self.project.id}/generate-base-asset/", {"kind": "person"}), (f"/api/projects/{self.project.id}/generate-triview/", {"portrait_asset_id": "x"}), (f"/api/projects/{self.project.id}/generate-storyboard/", {}), (f"/api/projects/{self.project.id}/submit-video-segment/", {"video_segment_id": "x"}), ("/api/ai/generate-image/", {"prompt": "测试"}), ] with self._blocked(): for url, payload in cases: response = self.client.post(url, payload, format="json") self.assertEqual(response.status_code, 503, msg=f"{url} 应被 worker 闸拦截") self.assertIn("worker", response.json()["detail"].lower() if isinstance(response.json().get("detail"), str) else "") def test_eager_mode_bypasses_gate(self): from apps.common.celery_health import celery_worker_available # 测试配置 CELERY_TASK_ALWAYS_EAGER=True:任务就地执行,无需 worker,闸放行 self.assertTrue(celery_worker_available()) def test_ping_failure_means_unavailable(self): from apps.common import celery_health celery_health._cache["expires"] = 0.0 with self.settings(CELERY_TASK_ALWAYS_EAGER=False): with patch.object(celery_health, "_ping_workers", side_effect=ConnectionError("broker down")): self.assertFalse(celery_health.celery_worker_available()) celery_health._cache["expires"] = 0.0