1397 lines
76 KiB
Python
1397 lines
76 KiB
Python
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 AITask, ModelConfig, ModelProvider
|
|
from apps.assets.models import Asset, AssetFile
|
|
from apps.billing.models import CreditAccount, CreditLedger
|
|
from apps.products.models import Product
|
|
from apps.projects.models import (
|
|
ExportJob,
|
|
Project,
|
|
ProjectStage,
|
|
ScriptSegment,
|
|
ScriptVersion,
|
|
StoryboardFrame,
|
|
StoryboardShot,
|
|
StoryboardShotVersion,
|
|
StoryboardVersion,
|
|
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")
|
|
# 迁移 0023(自由创作视频种子)已在测试库预建 volcengine provider,这里 get_or_create 避免撞唯一键
|
|
self.provider, _ = ModelProvider.objects.get_or_create(
|
|
name="volcengine",
|
|
defaults={
|
|
"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())
|
|
|
|
def test_delete_restore_and_purge_project_trash(self):
|
|
project = Project.objects.create(team=self.team, created_by=self.user, product=self.product, name="Trash Project")
|
|
segment = VideoSegment.objects.create(project=project, sort_order=0, target_duration_seconds=15)
|
|
asset = Asset.objects.create(
|
|
team=self.team,
|
|
created_by=self.user,
|
|
name="项目视频片段",
|
|
asset_type=Asset.Type.VIDEO,
|
|
source=Asset.Source.AI_GENERATED,
|
|
category=Asset.Category.VIDEO_CLIP,
|
|
)
|
|
asset_file = AssetFile.objects.create(
|
|
asset=asset,
|
|
object_key="projects/clip.mp4",
|
|
bucket="test",
|
|
content_type="video/mp4",
|
|
is_primary=True,
|
|
)
|
|
version = VideoSegmentVersion.objects.create(video_segment=segment, asset=asset, is_adopted=True)
|
|
segment.adopted_version = version
|
|
segment.save(update_fields=["adopted_version"])
|
|
|
|
delete = self.client.delete(f"/api/projects/{project.id}/")
|
|
self.assertEqual(delete.status_code, 204)
|
|
project.refresh_from_db()
|
|
self.assertTrue(project.is_deleted)
|
|
self.assertIsNone(project.purged_at)
|
|
list_ids = [p["id"] for p in self.client.get("/api/projects/").json()["results"]]
|
|
self.assertNotIn(str(project.id), list_ids)
|
|
trash_ids = [p["id"] for p in self.client.get("/api/projects/trash/").json()["results"]]
|
|
self.assertIn(str(project.id), trash_ids)
|
|
asset.refresh_from_db()
|
|
self.assertFalse(asset.is_deleted)
|
|
self.assertIsNone(asset.purged_at)
|
|
|
|
restore = self.client.post(f"/api/projects/{project.id}/restore/")
|
|
self.assertEqual(restore.status_code, 200)
|
|
project.refresh_from_db()
|
|
self.assertFalse(project.is_deleted)
|
|
self.assertIsNone(project.purged_at)
|
|
list_ids = [p["id"] for p in self.client.get("/api/projects/").json()["results"]]
|
|
self.assertIn(str(project.id), list_ids)
|
|
self.assertTrue(VideoSegmentVersion.objects.filter(id=version.id).exists())
|
|
|
|
self.client.delete(f"/api/projects/{project.id}/")
|
|
purge = self.client.delete(f"/api/projects/{project.id}/purge/")
|
|
self.assertEqual(purge.status_code, 204)
|
|
project.refresh_from_db()
|
|
self.assertTrue(project.is_deleted)
|
|
self.assertIsNotNone(project.purged_at)
|
|
self.assertTrue(Project.objects.filter(id=project.id).exists())
|
|
self.assertTrue(VideoSegment.objects.filter(project=project).exists())
|
|
self.assertTrue(VideoSegmentVersion.objects.filter(id=version.id).exists())
|
|
asset.refresh_from_db()
|
|
self.assertFalse(asset.is_deleted)
|
|
self.assertIsNone(asset.purged_at)
|
|
self.assertTrue(AssetFile.objects.filter(id=asset_file.id).exists())
|
|
trash_ids = [p["id"] for p in self.client.get("/api/projects/trash/").json()["results"]]
|
|
self.assertNotIn(str(project.id), trash_ids)
|
|
|
|
@patch("apps.ai.services.VolcanoArkProvider")
|
|
def test_rerun_script_segment_via_agent_rewrites_one_and_charges_once(self, provider_cls):
|
|
"""单镜重跑改走 agent:读结构化全脚本,只重写目标镜、保留其余镜,落新版本,计费一次。"""
|
|
import json as _json
|
|
base_draft = {
|
|
"hook": "钩子", "tone": "种草", "aspect_ratio": "9:16", "total_duration": 30, "segment_count": 2,
|
|
"entities": [],
|
|
"segments": [
|
|
{"index": 0, "duration": 15, "role": "钩子", "narration": "旧0", "visual": "画0", "speaker": None, "product_exposure": "", "entity_refs": [], "dialogue": []},
|
|
{"index": 1, "duration": 15, "role": "CTA", "narration": "旧1", "visual": "画1", "speaker": None, "product_exposure": "", "entity_refs": [], "dialogue": []},
|
|
],
|
|
}
|
|
new_draft = _json.loads(_json.dumps(base_draft))
|
|
new_draft["segments"][1]["narration"] = "全新口播文案"
|
|
new_draft["segments"][1]["visual"] = "全新画面描述"
|
|
provider = provider_cls.return_value
|
|
provider.chat_completion.return_value = {"choices": [{"message": {"content": "x"}}]}
|
|
provider.extract_text.return_value = "```json\n" + _json.dumps(new_draft, ensure_ascii=False) + "\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=_json.dumps(base_draft, ensure_ascii=False),
|
|
is_adopted=True, metadata={"hook": "钩子", "entities": []},
|
|
)
|
|
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)
|
|
# agent 落新 ScriptVersion:目标镜改了,其余镜保留
|
|
new_version = ScriptVersion.objects.filter(project=project).order_by("-created_at").first()
|
|
segs = list(new_version.segments.order_by("sort_order"))
|
|
self.assertEqual(segs[1].narration, "全新口播文案")
|
|
self.assertEqual(segs[1].visual_prompt, "全新画面描述")
|
|
self.assertEqual(segs[0].narration, "旧0")
|
|
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_extract_entities_streams_and_persists(self, provider_cls):
|
|
"""提取走流式通道(豆包 2.0 Pro 思考模型):思考事件丢弃、只收正文 JSON →
|
|
落库 cast/scenes + 每镜 entity_refs + entities_extracted 标记,计费一次。"""
|
|
import json as _json
|
|
extracted = {
|
|
"entities": [
|
|
{"id": "c1", "type": "character", "name": "女主", "visual_prompt": "26岁都市女性,利落短发"},
|
|
{"id": "s1", "type": "scene", "name": "出租屋客厅", "visual_prompt": "ins风小客厅,暖白自然光"},
|
|
],
|
|
"segments": [
|
|
{"index": 0, "entity_refs": ["c1", "s1"]},
|
|
{"index": 1, "entity_refs": ["c1", "s1"]},
|
|
],
|
|
}
|
|
provider = provider_cls.return_value
|
|
provider.chat_completion_stream.return_value = iter([
|
|
{"type": "reasoning", "text": "先通读分镜"}, # 思考事件:必须被丢弃,不进正文
|
|
{"type": "delta", "text": "```json\n" + _json.dumps(extracted, ensure_ascii=False)},
|
|
{"type": "delta", "text": "\n```"},
|
|
{"type": "done"},
|
|
])
|
|
|
|
project = Project.objects.create(team=self.team, created_by=self.user, product=self.product, name="P")
|
|
script = ScriptVersion.objects.create(
|
|
project=project, title="脚本", content="x", is_adopted=True, metadata={"entities": []},
|
|
)
|
|
ScriptSegment.objects.create(script_version=script, sort_order=0, narration="旧0", visual_prompt="画0")
|
|
ScriptSegment.objects.create(script_version=script, sort_order=1, narration="旧1", visual_prompt="画1")
|
|
|
|
response = self.client.post(f"/api/projects/{project.id}/extract-entities/", {}, format="json")
|
|
self.assertEqual(response.status_code, 200)
|
|
# 锁定豆包 2.0 Pro,且走的是流式通道(而非非流式 chat_completion)
|
|
self.assertEqual(provider.chat_completion_stream.call_args.kwargs["model"], "doubao-seed-2-0-pro-260215")
|
|
provider.chat_completion.assert_not_called()
|
|
# 实体落库:cast/scenes + entities_extracted 标记 + 每镜 entity_refs 回填
|
|
project.refresh_from_db()
|
|
self.assertIn("女主", project.metadata.get("cast", []))
|
|
self.assertIn("出租屋客厅", project.metadata.get("scenes", []))
|
|
self.assertTrue(project.metadata.get("entities_extracted"))
|
|
segs = list(script.segments.order_by("sort_order"))
|
|
self.assertEqual(segs[0].entity_refs, ["c1", "s1"])
|
|
self.assertEqual(CreditLedger.objects.filter(team=self.team, ledger_type=CreditLedger.Type.CHARGE).count(), 1)
|
|
|
|
@patch("apps.ai.services.VolcanoArkProvider")
|
|
def test_extract_entities_recovers_when_content_empty_uses_reasoning(self, provider_cls):
|
|
"""极端兜底:思考模型整轮只发了 reasoning、正文 content 为空 → 回退用 reasoning 里的 JSON,
|
|
不再像旧非流式那样拿到空串报「未返回有效 JSON」。"""
|
|
import json as _json
|
|
extracted = {
|
|
"entities": [{"id": "c1", "type": "character", "name": "男主", "visual_prompt": "28岁居家男性"}],
|
|
"segments": [{"index": 0, "entity_refs": ["c1"]}],
|
|
}
|
|
provider = provider_cls.return_value
|
|
provider.chat_completion_stream.return_value = iter([
|
|
{"type": "reasoning", "text": _json.dumps(extracted, ensure_ascii=False)},
|
|
{"type": "done"},
|
|
])
|
|
|
|
project = Project.objects.create(team=self.team, created_by=self.user, product=self.product, name="P2")
|
|
script = ScriptVersion.objects.create(
|
|
project=project, title="脚本", content="x", is_adopted=True, metadata={"entities": []},
|
|
)
|
|
ScriptSegment.objects.create(script_version=script, sort_order=0, narration="旧0", visual_prompt="画0")
|
|
|
|
response = self.client.post(f"/api/projects/{project.id}/extract-entities/", {}, format="json")
|
|
self.assertEqual(response.status_code, 200)
|
|
project.refresh_from_db()
|
|
self.assertIn("男主", project.metadata.get("cast", []))
|
|
|
|
@patch("apps.ai.services.VolcanoArkProvider")
|
|
def test_extract_entities_inflight_guard_reuses_task(self, provider_cls):
|
|
"""已有在途提取任务时再次提交 → 复用同一任务,不新建、不二次预扣、不重跑模型
|
|
(用户刷新后手痒重点 / 并发点击都不会重复扣费)。"""
|
|
from apps.ai.services import submit_extract_entities
|
|
|
|
project = Project.objects.create(team=self.team, created_by=self.user, product=self.product, name="P3")
|
|
script = ScriptVersion.objects.create(
|
|
project=project, title="脚本", content="x", is_adopted=True, metadata={"entities": []},
|
|
)
|
|
ScriptSegment.objects.create(script_version=script, sort_order=0, narration="旧0", visual_prompt="画0")
|
|
inflight = AITask.objects.create(
|
|
team=self.team, created_by=self.user, project=project, model_config=self.model,
|
|
task_type=AITask.Type.ENTITY_EXTRACTION, status=AITask.Status.SUBMITTED,
|
|
idempotency_key="entity_extraction:inflight:1",
|
|
)
|
|
before = AITask.objects.filter(project=project, task_type=AITask.Type.ENTITY_EXTRACTION).count()
|
|
|
|
task = submit_extract_entities(project=project, user=self.user)
|
|
self.assertEqual(str(task.id), str(inflight.id)) # 复用在途任务
|
|
self.assertEqual( # 没新建任务
|
|
AITask.objects.filter(project=project, task_type=AITask.Type.ENTITY_EXTRACTION).count(), before
|
|
)
|
|
provider_cls.return_value.chat_completion_stream.assert_not_called() # 没重跑模型
|
|
# 没二次预扣额度(复用直接 return,根本没走到 create_ai_task)
|
|
self.assertEqual(CreditLedger.objects.filter(team=self.team, ledger_type=CreditLedger.Type.RESERVE).count(), 0)
|
|
|
|
def test_extract_status_reports_running_then_safe_failure(self):
|
|
"""extract-status:在途→running=True;失败→安全错误对象,原始任务错误不透给前端。"""
|
|
project = Project.objects.create(team=self.team, created_by=self.user, product=self.product, name="P4")
|
|
t = AITask.objects.create(
|
|
team=self.team, created_by=self.user, project=project, model_config=self.model,
|
|
task_type=AITask.Type.ENTITY_EXTRACTION, status=AITask.Status.SUBMITTED,
|
|
idempotency_key="entity_extraction:status:1",
|
|
)
|
|
res = self.client.get(f"/api/projects/{project.id}/extract-status/")
|
|
self.assertEqual(res.status_code, 200)
|
|
self.assertTrue(res.data["running"])
|
|
|
|
t.status = AITask.Status.FAILED
|
|
t.error_message = "提取结果解析失败(模型未返回有效 JSON),请重试"
|
|
t.save(update_fields=["status", "error_message"])
|
|
res = self.client.get(f"/api/projects/{project.id}/extract-status/")
|
|
self.assertFalse(res.data["running"])
|
|
self.assertEqual(res.data["status"], "failed")
|
|
self.assertEqual(res.data["error"]["code"], "unknown")
|
|
self.assertNotIn("有效 JSON", res.data["error_message"])
|
|
|
|
@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_generate_base_asset_ignores_auto_triview_request(self, get_provider, store_media):
|
|
ModelConfig.objects.create(
|
|
provider=self.provider, name="img-model-auto-tri", display_name="Img Auto Tri",
|
|
capability=ModelConfig.Capability.IMAGE, endpoint="images/generations", unit_price="1.0000",
|
|
)
|
|
asset = Asset.objects.create(
|
|
team=self.team, created_by=self.user, name="person",
|
|
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-auto-tri")
|
|
|
|
response = self.client.post(
|
|
f"/api/projects/{project.id}/generate-base-asset/",
|
|
{"kind": "person", "prompt": "portrait", "label": "hero", "auto_triview": True},
|
|
format="json",
|
|
)
|
|
|
|
self.assertEqual(response.status_code, 202)
|
|
task = AITask.objects.get(id=response.data["task"]["id"])
|
|
self.assertFalse(task.request_payload.get("auto_triview"))
|
|
|
|
@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_script_syncs_video_segments_to_shot_count(self):
|
|
"""采用脚本时把视频片段数对齐到分镜数:用户在「采用前」删掉一镜(4→3),
|
|
采用后视频步骤的片段数必须跟着变 3,而不是停在项目创建固定铺的 4 段。"""
|
|
# 项目创建 → 固定 4 段视频片段
|
|
project = Project.objects.create(team=self.team, created_by=self.user, product=self.product, name="P")
|
|
for i in range(4):
|
|
VideoSegment.objects.create(project=project, sort_order=i, target_duration_seconds=15)
|
|
# 草稿脚本只剩 3 镜(模拟用户建了 4 镜又删 1,且尚未采用)
|
|
script = ScriptVersion.objects.create(project=project, is_adopted=False)
|
|
for i in range(3):
|
|
ScriptSegment.objects.create(script_version=script, sort_order=i, narration=f"n{i}", visual_prompt=f"v{i}")
|
|
self.assertEqual(project.video_segments.count(), 4)
|
|
|
|
response = self.client.post(
|
|
f"/api/projects/{project.id}/adopt-script/",
|
|
{"script_version_id": str(script.id)},
|
|
format="json",
|
|
)
|
|
|
|
self.assertEqual(response.status_code, 200)
|
|
# 采用后:多出的「从未生成过」尾段被裁掉,片段数 = 分镜数
|
|
self.assertEqual(project.video_segments.count(), 3)
|
|
|
|
def test_retrieve_self_heals_stale_video_segment_count(self):
|
|
"""历史脏项目(已在视频阶段、片段数比采用版分镜多)在详情加载时自愈到分镜数。"""
|
|
project = Project.objects.create(
|
|
team=self.team, created_by=self.user, product=self.product, name="P",
|
|
current_stage=ProjectStage.Stage.VIDEO,
|
|
)
|
|
for i in range(4):
|
|
VideoSegment.objects.create(project=project, sort_order=i, target_duration_seconds=15)
|
|
script = ScriptVersion.objects.create(project=project, is_adopted=True)
|
|
for i in range(3):
|
|
ScriptSegment.objects.create(script_version=script, sort_order=i, narration=f"n{i}", visual_prompt=f"v{i}")
|
|
|
|
response = self.client.get(f"/api/projects/{project.id}/")
|
|
|
|
self.assertEqual(response.status_code, 200)
|
|
self.assertEqual(project.video_segments.count(), 3)
|
|
self.assertEqual(len(response.data["video_segments"]), 3)
|
|
|
|
def test_adopt_script_keeps_generated_segment_even_if_excess(self):
|
|
"""收口绝不动已生成的段:即便采用版分镜少于已出片的片段数,已采用版本的尾段保留。"""
|
|
project = Project.objects.create(team=self.team, created_by=self.user, product=self.product, name="P")
|
|
segs = [VideoSegment.objects.create(project=project, sort_order=i, target_duration_seconds=15) for i in range(2)]
|
|
# 第 2 段(sort_order=1)已出片(有 adopted_version)→ 不可裁
|
|
asset = Asset.objects.create(
|
|
team=self.team, created_by=self.user, name="clip",
|
|
asset_type=Asset.Type.VIDEO, source=Asset.Source.AI_GENERATED, category=Asset.Category.VIDEO_CLIP,
|
|
)
|
|
ver = VideoSegmentVersion.objects.create(video_segment=segs[1], asset=asset, is_adopted=True)
|
|
segs[1].adopted_version = ver
|
|
segs[1].status = VideoSegment.Status.SUCCEEDED
|
|
segs[1].save(update_fields=["adopted_version", "status"])
|
|
# 采用版只有 1 镜
|
|
script = ScriptVersion.objects.create(project=project, is_adopted=False)
|
|
ScriptSegment.objects.create(script_version=script, sort_order=0, narration="n", visual_prompt="v")
|
|
|
|
response = self.client.post(
|
|
f"/api/projects/{project.id}/adopt-script/",
|
|
{"script_version_id": str(script.id)},
|
|
format="json",
|
|
)
|
|
|
|
self.assertEqual(response.status_code, 200)
|
|
# 已出片的尾段保留 → 不裁到 1,停在 2
|
|
self.assertEqual(project.video_segments.count(), 2)
|
|
|
|
def _mk_image_model(self):
|
|
return ModelConfig.objects.create(
|
|
provider=self.provider, name="gpt-image-2", display_name="img",
|
|
capability=ModelConfig.Capability.IMAGE, endpoint="images/generations", unit_price="1.0000",
|
|
)
|
|
|
|
def _mk_shot(self, project, script, order, *, review_status="active"):
|
|
"""建一个已出片的 StoryboardShot(+ 已采用版本,资产审核态可控),供过审/采用测试。"""
|
|
seg = script.segments.filter(sort_order=order).first() or ScriptSegment.objects.create(
|
|
script_version=script, sort_order=order, narration=f"n{order}", visual_prompt=f"v{order}")
|
|
a = Asset.objects.create(team=self.team, created_by=self.user, name=f"sb{order}", asset_type="image",
|
|
source="ai_generated", category=Asset.Category.STORYBOARD, review_status=review_status)
|
|
shot = StoryboardShot.objects.create(project=project, script_segment=seg, sort_order=order,
|
|
status=StoryboardShot.Status.SUCCEEDED)
|
|
ver = StoryboardShotVersion.objects.create(shot=shot, asset=a, is_adopted=True)
|
|
shot.adopted_version = ver
|
|
shot.save(update_fields=["adopted_version"])
|
|
return shot, ver, a
|
|
|
|
def test_submit_storyboard_creates_one_shot_per_segment_all_queued(self):
|
|
"""开始生成故事板:确保每镜一个 shot、全部置 QUEUED(真正出图交给 poll)。"""
|
|
from apps.ai.services import submit_storyboard
|
|
self._mk_image_model()
|
|
project = Project.objects.create(team=self.team, created_by=self.user, product=self.product, name="P")
|
|
script = ScriptVersion.objects.create(project=project, is_adopted=True)
|
|
for i in range(3):
|
|
ScriptSegment.objects.create(script_version=script, sort_order=i, narration=f"n{i}", visual_prompt=f"v{i}")
|
|
|
|
targets = submit_storyboard(project=project, user=self.user, prompt="风格")
|
|
|
|
self.assertEqual(project.storyboard_shots.count(), 3)
|
|
self.assertEqual(len(targets), 3)
|
|
self.assertTrue(all(s.status == StoryboardShot.Status.QUEUED for s in project.storyboard_shots.all()))
|
|
project.refresh_from_db()
|
|
self.assertEqual((project.metadata or {}).get("storyboard_prompt"), "风格") # 整张风格存项目级
|
|
|
|
def test_submit_storyboard_single_shot_only_targets_that_shot(self):
|
|
"""单场重跑:只把指定 shot 置 QUEUED,其余已出片的场不动。"""
|
|
from apps.ai.services import submit_storyboard
|
|
self._mk_image_model()
|
|
project = Project.objects.create(team=self.team, created_by=self.user, product=self.product, name="P")
|
|
script = ScriptVersion.objects.create(project=project, is_adopted=True)
|
|
for i in range(2):
|
|
ScriptSegment.objects.create(script_version=script, sort_order=i, narration=f"n{i}", visual_prompt=f"v{i}")
|
|
shot0, _, _ = self._mk_shot(project, script, 0)
|
|
shot1, _, _ = self._mk_shot(project, script, 1)
|
|
|
|
submit_storyboard(project=project, user=self.user, shot_ids=[str(shot1.id)])
|
|
|
|
shot0.refresh_from_db(); shot1.refresh_from_db()
|
|
self.assertEqual(shot0.status, StoryboardShot.Status.SUCCEEDED) # 没动
|
|
self.assertEqual(shot1.status, StoryboardShot.Status.QUEUED) # 只重跑这场
|
|
|
|
def test_adopt_storyboard_shot_version_switches_adopted(self):
|
|
"""采用某场的历史版本:切换该场采用的分镜图。"""
|
|
self._mk_image_model()
|
|
project = Project.objects.create(team=self.team, created_by=self.user, product=self.product, name="P")
|
|
script = ScriptVersion.objects.create(project=project, is_adopted=True)
|
|
ScriptSegment.objects.create(script_version=script, sort_order=0, narration="n", visual_prompt="v")
|
|
shot, v1, _ = self._mk_shot(project, script, 0)
|
|
a2 = Asset.objects.create(team=self.team, created_by=self.user, name="sb-v2", asset_type="image",
|
|
source="ai_generated", category=Asset.Category.STORYBOARD)
|
|
v2 = StoryboardShotVersion.objects.create(shot=shot, asset=a2, is_adopted=False)
|
|
|
|
res = self.client.post(
|
|
f"/api/projects/{project.id}/adopt-storyboard-shot-version/",
|
|
{"shot_id": str(shot.id), "version_id": str(v2.id)}, format="json",
|
|
)
|
|
self.assertEqual(res.status_code, 200)
|
|
shot.refresh_from_db(); v1.refresh_from_db(); v2.refresh_from_db()
|
|
self.assertEqual(shot.adopted_version_id, v2.id)
|
|
self.assertTrue(v2.is_adopted)
|
|
self.assertFalse(v1.is_adopted)
|
|
|
|
def test_video_review_precheck_flags_unreviewed_storyboard_then_clears(self):
|
|
"""生成视频前的过审闸:某镜故事板分镜未过审(非 active)→ precheck 列为 blocker;
|
|
过审(active)后 → blockers 清空,可放行生成。"""
|
|
project = Project.objects.create(
|
|
team=self.team, created_by=self.user, product=self.product, name="P",
|
|
current_stage=ProjectStage.Stage.VIDEO,
|
|
)
|
|
VideoSegment.objects.create(project=project, sort_order=0, target_duration_seconds=15)
|
|
script = ScriptVersion.objects.create(project=project, is_adopted=True)
|
|
ScriptSegment.objects.create(script_version=script, sort_order=0, narration="n", visual_prompt="v")
|
|
_, _, frame_asset = self._mk_shot(project, script, 0, review_status="")
|
|
|
|
res = self.client.post(f"/api/projects/{project.id}/video-review-precheck/", format="json")
|
|
self.assertEqual(res.status_code, 200)
|
|
blockers = res.json()["blockers"]
|
|
self.assertTrue(any(b["kind"] == "storyboard" and b["scene_no"] == 1 for b in blockers))
|
|
|
|
# 过审后再查 → 不再拦
|
|
frame_asset.review_status = "active"
|
|
frame_asset.save(update_fields=["review_status"])
|
|
res2 = self.client.post(f"/api/projects/{project.id}/video-review-precheck/", format="json")
|
|
self.assertEqual(res2.json()["blockers"], [])
|
|
|
|
def test_video_review_precheck_skips_succeeded_segments_in_batch(self):
|
|
"""整批校验(不传 segment_id)只看未出片的段:已 succeeded 的段即便分镜未过审也不拦。"""
|
|
project = Project.objects.create(
|
|
team=self.team, created_by=self.user, product=self.product, name="P",
|
|
current_stage=ProjectStage.Stage.VIDEO,
|
|
)
|
|
seg = VideoSegment.objects.create(project=project, sort_order=0, status=VideoSegment.Status.SUCCEEDED)
|
|
script = ScriptVersion.objects.create(project=project, is_adopted=True)
|
|
ScriptSegment.objects.create(script_version=script, sort_order=0, narration="n", visual_prompt="v")
|
|
self._mk_shot(project, script, 0, review_status="")
|
|
|
|
# 整批:succeeded 段跳过 → 不拦
|
|
res = self.client.post(f"/api/projects/{project.id}/video-review-precheck/", format="json")
|
|
self.assertEqual(res.json()["blockers"], [])
|
|
# 指定该段(单段重跑):无论 succeeded 都校验 → 拦
|
|
res2 = self.client.post(
|
|
f"/api/projects/{project.id}/video-review-precheck/",
|
|
{"video_segment_id": str(seg.id)}, format="json",
|
|
)
|
|
self.assertTrue(len(res2.json()["blockers"]) >= 1)
|
|
|
|
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 test_video_done_settles_even_if_stage_stuck_at_storyboard(self):
|
|
"""ZWQ#15:项目 current_stage 还停在 storyboard,但视频段都出片并采用 → 必须收口成 completed,
|
|
而不是永远卡在「3/4 故事板」。修复前 settle_video_completion 因 current_stage != VIDEO 直接返回 → 卡死。"""
|
|
project = Project.objects.create(
|
|
team=self.team, created_by=self.user, product=self.product, name="P",
|
|
current_stage=ProjectStage.Stage.STORYBOARD, status=Project.Status.STORYBOARDING,
|
|
)
|
|
seg = VideoSegment.objects.create(project=project, sort_order=0, target_duration_seconds=15)
|
|
asset = Asset.objects.create(
|
|
team=self.team, created_by=self.user, name="clip",
|
|
asset_type=Asset.Type.VIDEO, source=Asset.Source.AI_GENERATED, category=Asset.Category.VIDEO_CLIP,
|
|
)
|
|
ver = VideoSegmentVersion.objects.create(video_segment=seg, asset=asset)
|
|
|
|
res = self.client.post(
|
|
f"/api/projects/{project.id}/adopt-video-version/",
|
|
{"video_segment_id": str(seg.id), "version_id": str(ver.id)},
|
|
format="json",
|
|
)
|
|
self.assertEqual(res.status_code, 200)
|
|
project.refresh_from_db()
|
|
self.assertEqual(project.current_stage, ProjectStage.Stage.VIDEO)
|
|
self.assertEqual(project.status, Project.Status.COMPLETED)
|
|
|
|
def test_retrieve_recovers_stuck_project_to_completed(self):
|
|
"""ZWQ#15(存量项目):已经卡在 storyboard 阶段、视频段都出片采用的历史项目,
|
|
进入详情(GET)时应自动补推阶段并收口成 completed,不必用户再做任何操作。"""
|
|
project = Project.objects.create(
|
|
team=self.team, created_by=self.user, product=self.product, name="P",
|
|
current_stage=ProjectStage.Stage.STORYBOARD, status=Project.Status.STORYBOARDING,
|
|
)
|
|
for i in range(2):
|
|
seg = VideoSegment.objects.create(project=project, sort_order=i, target_duration_seconds=15)
|
|
asset = Asset.objects.create(
|
|
team=self.team, created_by=self.user, name=f"clip{i}",
|
|
asset_type=Asset.Type.VIDEO, source=Asset.Source.AI_GENERATED, category=Asset.Category.VIDEO_CLIP,
|
|
)
|
|
ver = VideoSegmentVersion.objects.create(video_segment=seg, asset=asset, is_adopted=True)
|
|
seg.adopted_version = ver
|
|
seg.status = VideoSegment.Status.SUCCEEDED
|
|
seg.save(update_fields=["adopted_version", "status"])
|
|
|
|
res = self.client.get(f"/api/projects/{project.id}/")
|
|
self.assertEqual(res.status_code, 200)
|
|
project.refresh_from_db()
|
|
self.assertEqual(project.current_stage, ProjectStage.Stage.VIDEO)
|
|
self.assertEqual(project.status, Project.Status.COMPLETED)
|
|
|
|
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.assertEqual(response.json()["error"]["code"], "provider_config_error")
|
|
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_export_command_uses_selected_output_size(self):
|
|
"""极速成片选横屏/清晰度后,多段合成必须沿用该尺寸,不能回退成固定竖屏。"""
|
|
from apps.projects.services.export import _build_export_command
|
|
|
|
cmd = _build_export_command(
|
|
n=1,
|
|
specs=[{"ts": 0.0, "te": 15.0, "dur": 15.0}],
|
|
starts=[0.0],
|
|
total=15.0,
|
|
transition="none",
|
|
sub_overlays=[],
|
|
bgm_name=None,
|
|
bgm_volume=1.0,
|
|
output_width=1920,
|
|
output_height=1080,
|
|
)
|
|
graph = cmd[cmd.index("-filter_complex") + 1]
|
|
self.assertIn("scale=1920:1080", graph)
|
|
self.assertIn("pad=1920:1080", 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
|
|
|
|
|
|
class ExportClipsTests(TestCase):
|
|
"""导出全部:把项目已采用视频片段打成一个 zip 下载。"""
|
|
|
|
def setUp(self):
|
|
self.user = User.objects.create_user(username="exp", password="p")
|
|
self.team = Team.objects.create(name="ExpT", owner=self.user)
|
|
TeamMember.objects.create(team=self.team, user=self.user, role=TeamMember.Role.OWNER)
|
|
self.product = Product.objects.create(team=self.team, created_by=self.user, title="P")
|
|
self.project = Project.objects.create(team=self.team, name="项目甲", product=self.product, created_by=self.user)
|
|
self.client = APIClient()
|
|
self.client.force_authenticate(self.user)
|
|
|
|
def _add_clip(self, order):
|
|
from apps.assets.models import AssetFile
|
|
|
|
asset = Asset.objects.create(team=self.team, name=f"clip{order}", asset_type="video", source="ai_generated", category="video_clip")
|
|
AssetFile.objects.create(asset=asset, object_key=f"k{order}.mp4", bucket="b", content_type="video/mp4", is_primary=True)
|
|
seg = VideoSegment.objects.create(project=self.project, sort_order=order)
|
|
ver = VideoSegmentVersion.objects.create(video_segment=seg, asset=asset, is_adopted=True)
|
|
seg.adopted_version = ver
|
|
seg.save(update_fields=["adopted_version"])
|
|
|
|
def test_export_zip_packs_all_clips(self):
|
|
self._add_clip(0)
|
|
self._add_clip(1)
|
|
with patch("apps.projects.views.TosStorage") as Tos, patch("requests.get") as rget:
|
|
Tos.return_value.presigned_get_url.return_value = "http://x/c.mp4"
|
|
rget.return_value.content = b"MP4DATA"
|
|
res = self.client.get(f"/api/projects/{self.project.id}/export-clips/")
|
|
self.assertEqual(res.status_code, 200)
|
|
self.assertEqual(res["Content-Type"], "application/zip")
|
|
import io
|
|
import zipfile
|
|
|
|
zf = zipfile.ZipFile(io.BytesIO(res.content))
|
|
self.assertEqual(len(zf.namelist()), 2)
|
|
|
|
def test_export_empty_returns_400(self):
|
|
res = self.client.get(f"/api/projects/{self.project.id}/export-clips/")
|
|
self.assertEqual(res.status_code, 400)
|
|
|
|
|
|
class MergeFinalVideoTests(TestCase):
|
|
"""合成成片(submit-export):按当前采用片段重建时间线 · 在跑的任务复用 · 成片地址下发。"""
|
|
|
|
def setUp(self):
|
|
self.user = User.objects.create_user(username="mrg", password="p")
|
|
self.team = Team.objects.create(name="MrgT", owner=self.user)
|
|
TeamMember.objects.create(team=self.team, user=self.user, role=TeamMember.Role.OWNER)
|
|
self.product = Product.objects.create(team=self.team, created_by=self.user, title="P")
|
|
self.project = Project.objects.create(
|
|
team=self.team, name="项目甲", product=self.product, created_by=self.user,
|
|
current_stage=ProjectStage.Stage.VIDEO, status=Project.Status.COMPLETED,
|
|
)
|
|
self.client = APIClient()
|
|
self.client.force_authenticate(self.user)
|
|
|
|
def _video_asset(self, tag):
|
|
asset = Asset.objects.create(
|
|
team=self.team, name=f"clip-{tag}", asset_type="video", source="ai_generated", category="video_clip",
|
|
)
|
|
AssetFile.objects.create(
|
|
asset=asset, object_key=f"{tag}.mp4", bucket="b", content_type="video/mp4",
|
|
is_primary=True, preview_url=f"http://x/{tag}.mp4",
|
|
)
|
|
return asset
|
|
|
|
def _segment(self, order, tag):
|
|
asset = self._video_asset(tag)
|
|
seg = VideoSegment.objects.create(
|
|
project=self.project, sort_order=order, target_duration_seconds=10,
|
|
status=VideoSegment.Status.SUCCEEDED,
|
|
)
|
|
version = VideoSegmentVersion.objects.create(video_segment=seg, asset=asset, is_adopted=True)
|
|
seg.adopted_version = version
|
|
seg.save(update_fields=["adopted_version"])
|
|
return seg, asset
|
|
|
|
def _submit(self):
|
|
with patch("apps.projects.views.run_export_job_in_thread") as runner:
|
|
with self.captureOnCommitCallbacks(execute=True):
|
|
res = self.client.post(f"/api/projects/{self.project.id}/submit-export/")
|
|
return res, runner
|
|
|
|
def test_submit_builds_timeline_and_keeps_video_stage(self):
|
|
self._segment(0, "a")
|
|
self._segment(1, "b")
|
|
res, runner = self._submit()
|
|
self.assertEqual(res.status_code, 202)
|
|
runner.assert_called_once()
|
|
timeline = Timeline.objects.get(project=self.project)
|
|
clips = list(timeline.clips.order_by("sort_order"))
|
|
self.assertEqual(len(clips), 2)
|
|
self.assertEqual([c.start_ms for c in clips], [0, 10000])
|
|
# V1 视频是末阶段:合成不能把项目推到被雪藏的第 5 阶段(否则流水线页落到看不见的页面)
|
|
self.project.refresh_from_db()
|
|
self.assertEqual(self.project.current_stage, ProjectStage.Stage.VIDEO)
|
|
|
|
def test_submit_reuses_running_job(self):
|
|
self._segment(0, "a")
|
|
first, _ = self._submit()
|
|
job_id = first.data["id"]
|
|
ExportJob.objects.filter(id=job_id).update(status=ExportJob.Status.RUNNING)
|
|
second, runner = self._submit()
|
|
self.assertEqual(second.status_code, 202)
|
|
self.assertEqual(str(second.data["id"]), str(job_id))
|
|
self.assertEqual(ExportJob.objects.filter(timeline__project=self.project).count(), 1)
|
|
runner.assert_not_called()
|
|
|
|
def test_resubmit_rebuilds_clips_after_segment_rerun(self):
|
|
seg, _old = self._segment(0, "old")
|
|
self._submit()
|
|
# 上一次合成已收工(否则「重新合成」会复用在跑的任务,见 test_submit_reuses_running_job)
|
|
ExportJob.objects.filter(timeline__project=self.project).update(status=ExportJob.Status.SUCCEEDED)
|
|
new_asset = self._video_asset("new")
|
|
version = VideoSegmentVersion.objects.create(video_segment=seg, asset=new_asset, is_adopted=True)
|
|
seg.adopted_version = version
|
|
seg.save(update_fields=["adopted_version"])
|
|
self._submit()
|
|
clips = list(Timeline.objects.get(project=self.project).clips.all())
|
|
self.assertEqual([c.asset_id for c in clips], [new_asset.id])
|
|
|
|
def test_final_video_url_exposed_after_success(self):
|
|
self._segment(0, "a")
|
|
self._submit()
|
|
timeline = Timeline.objects.get(project=self.project)
|
|
final = Asset.objects.create(
|
|
team=self.team, name="final.mp4", asset_type="video", source="exported", category="final_video",
|
|
)
|
|
AssetFile.objects.create(
|
|
asset=final, object_key="final.mp4", bucket="b", content_type="video/mp4",
|
|
is_primary=True, preview_url="http://x/final.mp4",
|
|
)
|
|
job = timeline.export_jobs.order_by("-created_at").first()
|
|
job.status = "succeeded"
|
|
job.output_asset = final
|
|
job.progress = 100
|
|
job.save(update_fields=["status", "output_asset", "progress"])
|
|
|
|
detail = self.client.get(f"/api/projects/{self.project.id}/")
|
|
self.assertEqual(detail.data["final_video_url"], "http://x/final.mp4")
|
|
listed = self.client.get("/api/projects/")
|
|
row = next(p for p in listed.data["results"] if str(p["id"]) == str(self.project.id))
|
|
self.assertEqual(row["final_video_url"], "http://x/final.mp4")
|
|
|
|
def test_final_video_url_empty_before_merge(self):
|
|
self._segment(0, "a")
|
|
detail = self.client.get(f"/api/projects/{self.project.id}/")
|
|
self.assertEqual(detail.data["final_video_url"], "")
|
|
|
|
def test_merge_drops_bgm_without_audio_stream(self):
|
|
"""BGM 文件里没有音频流(上传串了张图 / 演示种子数据)→ 丢掉 BGM 照常拼片,
|
|
不能让滤镜里的 [n:a] 匹配不到流,把整条合成炸掉。"""
|
|
from pathlib import Path
|
|
|
|
from apps.projects.models import BgmTrack
|
|
from apps.projects.services import export as export_mod
|
|
|
|
seg, _asset = self._segment(0, "a")
|
|
timeline = Timeline.objects.create(project=self.project, name="T", duration_seconds=10)
|
|
TimelineClip.objects.create(
|
|
timeline=timeline, asset=seg.adopted_version.asset, sort_order=0, start_ms=0, duration_ms=10000,
|
|
)
|
|
fake_bgm = Asset.objects.create(team=self.team, name="BGM", asset_type="audio", source="upload")
|
|
AssetFile.objects.create(
|
|
asset=fake_bgm, object_key="cover.png", bucket="b", content_type="image/png", is_primary=True,
|
|
)
|
|
BgmTrack.objects.create(timeline=timeline, asset=fake_bgm, volume=60, start_ms=0)
|
|
job = ExportJob.objects.create(timeline=timeline, status=ExportJob.Status.QUEUED)
|
|
|
|
commands = []
|
|
|
|
def fake_run(cmd, **kwargs):
|
|
commands.append(cmd)
|
|
# ffprobe:一律回「无音频流 / 无帧率」,模拟素材探测不出音轨
|
|
if cmd[0] == "ffprobe":
|
|
return type("P", (), {"returncode": 0, "stdout": b"", "stderr": b""})()
|
|
(Path(kwargs["cwd"]) / "output.mp4").write_bytes(b"MP4")
|
|
return type("P", (), {"returncode": 0, "stdout": b"", "stderr": b""})()
|
|
|
|
with patch.object(export_mod, "_download_asset_primary_file", lambda asset, path: path.write_bytes(b"X")), \
|
|
patch.object(export_mod.subprocess, "run", side_effect=fake_run), \
|
|
patch.object(export_mod, "TosStorage") as Tos:
|
|
Tos.return_value.upload_fileobj.return_value = type(
|
|
"S", (), {"object_key": "exports/x.mp4", "bucket": "b", "content_type": "video/mp4", "size_bytes": 3},
|
|
)()
|
|
export_mod.run_export_job(str(job.id))
|
|
|
|
job.refresh_from_db()
|
|
self.assertEqual(job.status, ExportJob.Status.SUCCEEDED)
|
|
ffmpeg_cmd = next(c for c in commands if c[0] == "ffmpeg")
|
|
self.assertNotIn("-stream_loop", ffmpeg_cmd)
|
|
self.assertNotIn("[abgm]", ffmpeg_cmd[ffmpeg_cmd.index("-filter_complex") + 1])
|
|
|
|
def test_final_video_url_skips_non_video_output(self):
|
|
"""演示种子数据里有「导出成功但挂了张 PNG 海报」的成片:不能下发,否则播放键打开放不出的图。"""
|
|
self._segment(0, "a")
|
|
self._submit()
|
|
poster = Asset.objects.create(
|
|
team=self.team, name="成片海报", asset_type="video", source="exported", category="final_video",
|
|
)
|
|
AssetFile.objects.create(
|
|
asset=poster, object_key="poster.png", bucket="b", content_type="image/png",
|
|
is_primary=True, preview_url="http://x/poster.png",
|
|
)
|
|
ExportJob.objects.filter(timeline__project=self.project).update(
|
|
status=ExportJob.Status.SUCCEEDED, output_asset=poster,
|
|
)
|
|
detail = self.client.get(f"/api/projects/{self.project.id}/")
|
|
self.assertEqual(detail.data["final_video_url"], "")
|
|
|
|
|
|
class BaseAssetAdoptDeleteTests(TestCase):
|
|
"""显性采用/未采用 + 删除角色组。"""
|
|
|
|
def setUp(self):
|
|
from apps.projects.models import BaseAssetGroup
|
|
|
|
self.user = User.objects.create_user(username="bad", password="x")
|
|
self.team = Team.objects.create(name="BAD", owner=self.user)
|
|
TeamMember.objects.create(team=self.team, user=self.user, role=TeamMember.Role.OWNER)
|
|
self.product = Product.objects.create(team=self.team, created_by=self.user, title="P")
|
|
self.proj = Project.objects.create(team=self.team, name="项目甲", product=self.product, created_by=self.user)
|
|
self.group = BaseAssetGroup.objects.create(project=self.proj, kind=BaseAssetGroup.Kind.PERSON, metadata={"label": "小明"})
|
|
self.client = APIClient()
|
|
self.client.force_authenticate(self.user)
|
|
|
|
def test_set_adopt_state(self):
|
|
res = self.client.post(f"/api/projects/{self.proj.id}/set-base-asset-adopt/", {"group_id": str(self.group.id), "state": "unadopted"}, format="json")
|
|
self.assertEqual(res.status_code, 200)
|
|
self.group.refresh_from_db()
|
|
self.assertEqual(self.group.metadata.get("adopt"), "unadopted")
|
|
|
|
def test_set_adopt_invalid_state(self):
|
|
res = self.client.post(f"/api/projects/{self.proj.id}/set-base-asset-adopt/", {"group_id": str(self.group.id), "state": "x"}, format="json")
|
|
self.assertEqual(res.status_code, 400)
|
|
|
|
def test_delete_base_asset(self):
|
|
from apps.projects.models import BaseAssetGroup
|
|
|
|
res = self.client.post(f"/api/projects/{self.proj.id}/delete-base-asset/", {"group_id": str(self.group.id)}, format="json")
|
|
self.assertEqual(res.status_code, 200)
|
|
self.assertFalse(BaseAssetGroup.objects.filter(id=self.group.id).exists())
|
|
|
|
|
|
class AttachOfficialModelTests(TestCase):
|
|
"""ZWQ#4/#5 回归:从演员库挑「官方模特」(跨团队)新增/替换角色,旧版因 team 隔离报 asset not found
|
|
且替换失败、还留下空角色组。修复后:官方模特形象图自动克隆进本团队并真正挂上,不再 404。"""
|
|
|
|
def setUp(self):
|
|
from apps.assets.models import Asset, AssetFile, Model
|
|
from apps.projects.models import BaseAssetGroup
|
|
|
|
self.user = User.objects.create_user(username="biz", password="x")
|
|
self.team = Team.objects.create(name="商家团队", owner=self.user)
|
|
TeamMember.objects.create(team=self.team, user=self.user, role=TeamMember.Role.OWNER)
|
|
self.product = Product.objects.create(team=self.team, created_by=self.user, title="口红")
|
|
self.proj = Project.objects.create(team=self.team, name="带货视频", product=self.product, created_by=self.user)
|
|
|
|
# 官方模特属于「官方团队」(跨团队可见,但 portrait 资产不属本商家团队)——正是触发 bug 的根因。
|
|
self.official_owner = User.objects.create_user(username="ops", password="x")
|
|
self.official_team = Team.objects.create(name="官方", owner=self.official_owner)
|
|
self.portrait = Asset.objects.create(
|
|
team=self.official_team, created_by=self.official_owner, name="都市白领 林晚",
|
|
asset_type=Asset.Type.IMAGE, source=Asset.Source.AI_GENERATED, category=Asset.Category.PERSON,
|
|
metadata={"kind": "model"},
|
|
)
|
|
AssetFile.objects.create(asset=self.portrait, object_key="m/linwan.png", bucket="b", content_type="image/png", preview_url="http://x/linwan.png", is_primary=True)
|
|
self.model = Model.objects.create(
|
|
team=self.official_team, created_by=self.official_owner, name="都市白领 林晚",
|
|
is_official=True, portrait_asset=self.portrait,
|
|
)
|
|
self.BaseAssetGroup = BaseAssetGroup
|
|
self.client = APIClient()
|
|
self.client.force_authenticate(self.user)
|
|
|
|
def test_add_actor_from_official_model_clones_and_attaches(self):
|
|
"""新增角色:无 group,传 kind+label,选官方模特 → 200,克隆进本团队并采用,不再 asset not found。"""
|
|
from apps.assets.models import Asset
|
|
|
|
res = self.client.post(
|
|
f"/api/projects/{self.proj.id}/attach-base-asset/",
|
|
{"kind": "person", "label": "林晚", "asset_id": str(self.portrait.id)},
|
|
format="json",
|
|
)
|
|
self.assertEqual(res.status_code, 200, res.content)
|
|
group = self.BaseAssetGroup.objects.get(id=res.json()["id"])
|
|
self.assertIsNotNone(group.adopted_asset_id)
|
|
clone = Asset.objects.get(id=group.adopted_asset_id)
|
|
# 克隆体真正落进本商家团队(不是官方团队),并共享同一份 TOS 对象。
|
|
self.assertEqual(clone.team_id, self.team.id)
|
|
self.assertNotEqual(clone.id, self.portrait.id)
|
|
self.assertEqual(clone.files.first().object_key, "m/linwan.png")
|
|
self.assertEqual(clone.metadata.get("cloned_from_asset"), str(self.portrait.id))
|
|
|
|
def test_replace_existing_group_with_official_model(self):
|
|
"""替换:已有角色组,传 group_id + 官方模特 asset_id → 200,采用克隆体。"""
|
|
from apps.assets.models import Asset
|
|
|
|
group = self.BaseAssetGroup.objects.create(project=self.proj, kind=self.BaseAssetGroup.Kind.PERSON, metadata={"label": "主播"})
|
|
res = self.client.post(
|
|
f"/api/projects/{self.proj.id}/attach-base-asset/",
|
|
{"group_id": str(group.id), "asset_id": str(self.portrait.id)},
|
|
format="json",
|
|
)
|
|
self.assertEqual(res.status_code, 200, res.content)
|
|
group.refresh_from_db()
|
|
clone = Asset.objects.get(id=group.adopted_asset_id)
|
|
self.assertEqual(clone.team_id, self.team.id)
|
|
|
|
def test_unknown_asset_still_404_without_ghost_group(self):
|
|
"""彻底不存在的 asset_id → 仍 404,且不留下空角色组(不再在 404 前提交建组)。"""
|
|
import uuid as _uuid
|
|
|
|
before = self.BaseAssetGroup.objects.filter(project=self.proj).count()
|
|
res = self.client.post(
|
|
f"/api/projects/{self.proj.id}/attach-base-asset/",
|
|
{"kind": "person", "label": "幽灵", "asset_id": str(_uuid.uuid4())},
|
|
format="json",
|
|
)
|
|
self.assertEqual(res.status_code, 404)
|
|
self.assertEqual(self.BaseAssetGroup.objects.filter(project=self.proj).count(), before)
|
|
|
|
def test_attach_own_actor_auto_adopts_and_submits_review(self):
|
|
"""ZWQ#7:在演员库选用「自己创建的角色」(本团队 person 资产)后,组应自动标 adopt=adopted(已采用),
|
|
且自动送审(已过审),不必用户再手动点采用 / 灰盾。"""
|
|
from unittest.mock import patch
|
|
|
|
from apps.assets.models import Asset, AssetFile
|
|
|
|
mine = 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=mine, object_key="me/host.png", bucket="b", content_type="image/png", preview_url="http://x/host.png", is_primary=True)
|
|
group = self.BaseAssetGroup.objects.create(project=self.proj, kind=self.BaseAssetGroup.Kind.PERSON, metadata={"label": "主播"})
|
|
|
|
with patch("apps.assets.review.submit_asset_for_review") as sub:
|
|
with self.captureOnCommitCallbacks(execute=True):
|
|
res = self.client.post(
|
|
f"/api/projects/{self.proj.id}/attach-base-asset/",
|
|
{"group_id": str(group.id), "asset_id": str(mine.id)},
|
|
format="json",
|
|
)
|
|
self.assertEqual(res.status_code, 200, res.content)
|
|
group.refresh_from_db()
|
|
self.assertEqual(group.adopted_asset_id, mine.id)
|
|
self.assertEqual((group.metadata or {}).get("adopt"), "adopted") # 自动采用
|
|
sub.assert_called_once() # 自动送审(on_commit 在 TestCase 事务提交点触发)
|
|
self.assertEqual(sub.call_args.args[0].id, mine.id)
|
|
|
|
|
|
class VideoSegmentTrueUpTests(TestCase):
|
|
"""pipeline 视频段积分计价 + 按火山真实 usage.total_tokens 结算(true-up)。
|
|
终结「视频 ¥1/段、成本 ¥15」的倒贴定价:预留=预估积分×buffer,终态按实结算、差额自动 RELEASE。"""
|
|
|
|
def setUp(self):
|
|
from decimal import Decimal
|
|
|
|
self.user = User.objects.create_user(username="tu-owner", password="p")
|
|
self.team = Team.objects.create(name="TU", owner=self.user)
|
|
TeamMember.objects.create(team=self.team, user=self.user, role=TeamMember.Role.OWNER)
|
|
CreditAccount.objects.create(team=self.team, balance=Decimal("10000.0000"))
|
|
self.product = Product.objects.create(team=self.team, created_by=self.user, title="P")
|
|
self.project = Project.objects.create(team=self.team, created_by=self.user, product=self.product, name="TU-P")
|
|
self.segment = VideoSegment.objects.create(project=self.project, sort_order=0, target_duration_seconds=15)
|
|
self.provider = patch("apps.ai.services.build_provider").start().return_value
|
|
self.provider.create_video_task.return_value = {"id": "ark-tu-1", "status": "queued"}
|
|
self.addCleanup(patch.stopall)
|
|
|
|
def _submit(self):
|
|
from apps.ai.services import submit_video_segment
|
|
|
|
submit_video_segment(video_segment=self.segment, user=self.user, prompt="测试")
|
|
return AITask.objects.get(team=self.team, task_type=AITask.Type.VIDEO_SEGMENT)
|
|
|
|
def test_submit_prices_by_tokens_with_buffer(self):
|
|
from apps.billing.pricing import quote_video_estimate, video_reserve_amount
|
|
|
|
task = self._submit()
|
|
model = task.model_config
|
|
tokens, quote = quote_video_estimate(model, aspect_ratio="9:16", resolution="720p", duration=15, references=[])
|
|
self.assertEqual(task.estimated_cost, quote.points)
|
|
self.assertEqual(task.base_cost, quote.base_cost_yuan)
|
|
self.assertGreater(quote.points, 100) # 告别 ¥1/段:15s 720p 竖屏应是三位数积分
|
|
self.assertEqual(task.credit_reservation.amount, video_reserve_amount(quote.points))
|
|
self.assertEqual(task.request_payload["estimated_tokens"], tokens)
|
|
|
|
@patch("apps.ai.services._store_generated_media")
|
|
def test_poll_settles_by_actual_usage_tokens(self, store):
|
|
from decimal import Decimal
|
|
|
|
from apps.ai.services import poll_video_segment
|
|
from apps.billing.pricing import quote_video_actual
|
|
|
|
task = self._submit()
|
|
store.return_value = Asset.objects.create(
|
|
team=self.team, created_by=self.user, name="clip",
|
|
asset_type=Asset.Type.VIDEO, source=Asset.Source.AI_GENERATED, category=Asset.Category.VIDEO_CLIP,
|
|
)
|
|
self.provider.poll_video_task.return_value = {
|
|
"status": "succeeded",
|
|
"usage": {"total_tokens": 300000},
|
|
"content": {"video_url": "http://x/v.mp4"},
|
|
}
|
|
self.provider.extract_first_media_url.return_value = "http://x/v.mp4"
|
|
poll_video_segment(video_segment=self.segment, user=self.user)
|
|
task.refresh_from_db()
|
|
settle = quote_video_actual(task.model_config, tokens=300000, with_video_ref=False, resolution="720p")
|
|
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
|
|
self.assertEqual(task.actual_cost, settle.points)
|
|
self.assertEqual(task.base_cost, settle.base_cost_yuan)
|
|
# 真实 tokens(30万) < 预估(32.4万):实扣 < 预留,charge 自动 RELEASE 差额,冻结清零
|
|
account = CreditAccount.objects.get(team=self.team)
|
|
self.assertEqual(account.reserved_balance, Decimal("0"))
|
|
self.assertEqual(account.balance, Decimal("10000.0000") - settle.points)
|
|
charges = CreditLedger.objects.filter(task=task, ledger_type=CreditLedger.Type.CHARGE)
|
|
self.assertEqual(charges.count(), 1)
|
|
self.assertEqual(charges.first().amount, settle.points)
|
|
|
|
@patch("apps.ai.services._store_generated_media")
|
|
def test_poll_falls_back_to_estimate_when_usage_missing(self, store):
|
|
from apps.ai.services import poll_video_segment
|
|
|
|
task = self._submit()
|
|
store.return_value = Asset.objects.create(
|
|
team=self.team, created_by=self.user, name="clip2",
|
|
asset_type=Asset.Type.VIDEO, source=Asset.Source.AI_GENERATED, category=Asset.Category.VIDEO_CLIP,
|
|
)
|
|
self.provider.poll_video_task.return_value = {"status": "succeeded", "content": {"video_url": "http://x/v.mp4"}}
|
|
self.provider.extract_first_media_url.return_value = "http://x/v.mp4"
|
|
poll_video_segment(video_segment=self.segment, user=self.user)
|
|
task.refresh_from_db()
|
|
self.assertEqual(task.actual_cost, task.estimated_cost) # usage 缺失回落预估,不阻断出片
|