测试极速成片
This commit is contained in:
@@ -5,6 +5,26 @@ import requests
|
||||
|
||||
from .volcano import VolcanoArkProvider
|
||||
|
||||
# gpt-image / New API 网关只认这三档。业务层为了「真 16:9」会传 1536x864,
|
||||
# 不在这里收成合法横图的话,立绘能成、三视图 image_edit 直接 400。
|
||||
_OPENAI_IMAGE_SIZES = {"1024x1024", "1024x1536", "1536x1024"}
|
||||
|
||||
|
||||
def openai_image_size(size: str, default: str = "1024x1536") -> str:
|
||||
"""把任意宽高收成 gpt-image 接受的三档:方 / 竖 / 横。"""
|
||||
raw = (size or "").strip()
|
||||
if raw in _OPENAI_IMAGE_SIZES:
|
||||
return raw
|
||||
try:
|
||||
width, height = (int(part) for part in raw.lower().replace("×", "x").split("x", 1))
|
||||
except (TypeError, ValueError):
|
||||
return default
|
||||
if width <= 0 or height <= 0:
|
||||
return default
|
||||
if width == height:
|
||||
return "1024x1024"
|
||||
return "1536x1024" if width > height else "1024x1536"
|
||||
|
||||
|
||||
def _downscale_ref_for_edit(data: bytes, content_type: str, max_edge: int = 2048) -> tuple[bytes, str]:
|
||||
"""图生图参考图过大时按需降采样,避免中转站(yunqi/gpt-image edits)拒收大图报 400
|
||||
@@ -91,7 +111,7 @@ class OpenAICompatibleProvider(VolcanoArkProvider):
|
||||
"""文生图(可选单图参考 base64)。多图参考请用 image_edit。返回体含 url 或 b64_json。"""
|
||||
if not self.api_key:
|
||||
raise ValueError("中转站 api_key 未配置")
|
||||
body: dict[str, Any] = {"model": model, "prompt": prompt, "size": size, "n": 1}
|
||||
body: dict[str, Any] = {"model": model, "prompt": prompt, "size": openai_image_size(size), "n": 1}
|
||||
if image:
|
||||
body["image"] = image
|
||||
# 实测中转站生图延迟可达 75s+,超时给到 300s
|
||||
@@ -135,7 +155,7 @@ class OpenAICompatibleProvider(VolcanoArkProvider):
|
||||
files.append(("image[]", (f"ref{idx + 1}.{ext}", img_bytes, content_type)))
|
||||
if not files:
|
||||
raise ValueError("image_edit 至少需要一张参考图")
|
||||
data = {"model": model, "prompt": prompt, "size": size, "n": "1"}
|
||||
data = {"model": model, "prompt": prompt, "size": openai_image_size(size), "n": "1"}
|
||||
response = requests.post(
|
||||
self._endpoint_url(endpoint),
|
||||
headers={"Authorization": f"Bearer {self.api_key}"}, # multipart 不要手设 Content-Type
|
||||
|
||||
@@ -3,6 +3,7 @@ from typing import Any
|
||||
import requests
|
||||
from django.conf import settings
|
||||
|
||||
from .openai_compatible import openai_image_size
|
||||
from .volcano import VolcanoArkProvider
|
||||
|
||||
|
||||
@@ -32,7 +33,7 @@ class YunqiProvider(VolcanoArkProvider):
|
||||
raise ValueError("YUNQI_API_KEY is not configured")
|
||||
if image:
|
||||
raise ValueError("gpt-image-2 via YunQi only supports text-to-image; reference image is not supported")
|
||||
body: dict[str, Any] = {"model": model, "prompt": prompt, "size": size, "n": 1}
|
||||
body: dict[str, Any] = {"model": model, "prompt": prompt, "size": openai_image_size(size), "n": 1}
|
||||
# 实测该网关生图平均延迟 75s+,超时须显著高于火山的 180s
|
||||
response = requests.post(
|
||||
f"{self.base_url.rstrip('/')}/{endpoint.lstrip('/')}",
|
||||
|
||||
@@ -1411,6 +1411,7 @@ def _ratio_to_image_size(ratio: str) -> str:
|
||||
"9:16": "1024x1536", # 近似竖图(网关无精确 9:16)
|
||||
"4:3": "1536x1024",
|
||||
"16:9": "1536x864", # 真 16:9(原来误用 1536x1024 = 3:2)
|
||||
"21:9": "1536x1024", # 超宽近似横图(网关无精确 21:9)
|
||||
}
|
||||
normalized = (ratio or "").strip()
|
||||
if normalized in known:
|
||||
@@ -1429,6 +1430,56 @@ def _ratio_to_image_size(ratio: str) -> str:
|
||||
return "1024x1024"
|
||||
|
||||
|
||||
def project_output_spec(project) -> dict:
|
||||
"""专业创作 / 极速成片共用的成片规格:画幅、分辨率、视频模型。缺省 9:16 + 720p。"""
|
||||
wizard = dict((project.metadata or {}).get("wizard") or {})
|
||||
return {
|
||||
"aspect_ratio": str(wizard.get("aspect_ratio") or "9:16").strip() or "9:16",
|
||||
"resolution": str(wizard.get("resolution") or "720p").strip().lower() or "720p",
|
||||
"video_model_config_id": str(wizard.get("video_model_config_id") or "").strip() or None,
|
||||
}
|
||||
|
||||
|
||||
def _storyboard_canvas_phrase(ratio: str) -> str:
|
||||
r = (ratio or "9:16").strip() or "9:16"
|
||||
if r in {"9:16", "3:4"}:
|
||||
return f"电商竖屏 {r}"
|
||||
if r == "1:1":
|
||||
return f"电商方形 {r}"
|
||||
return f"电商横屏 {r}"
|
||||
|
||||
|
||||
def _apply_storyboard_output_ratio(text: str, project) -> str:
|
||||
phrase = _storyboard_canvas_phrase(project_output_spec(project)["aspect_ratio"])
|
||||
return (text or "").replace("电商竖屏 9:16", phrase)
|
||||
|
||||
|
||||
def _sync_timeline_output_spec(project, *, aspect_ratio: str, resolution: str) -> None:
|
||||
from apps.ai.video_pricing import get_resolution
|
||||
|
||||
try:
|
||||
width, height = get_resolution(aspect_ratio, resolution)
|
||||
except Exception: # noqa: BLE001 — 规格不合法时不挡提交,时间线保持原值
|
||||
return
|
||||
pixels = f"{width}x{height}"
|
||||
timeline, created = Timeline.objects.get_or_create(
|
||||
project=project,
|
||||
defaults={
|
||||
"name": f"{project.name} Timeline",
|
||||
"duration_seconds": 60,
|
||||
"aspect_ratio": aspect_ratio,
|
||||
"resolution": pixels,
|
||||
},
|
||||
)
|
||||
if created:
|
||||
return
|
||||
if timeline.aspect_ratio == aspect_ratio and timeline.resolution == pixels:
|
||||
return
|
||||
timeline.aspect_ratio = aspect_ratio
|
||||
timeline.resolution = pixels
|
||||
timeline.save(update_fields=["aspect_ratio", "resolution", "updated_at"])
|
||||
|
||||
|
||||
def _ratio_to_volcano_size(ratio: str) -> str:
|
||||
"""前端比例 → 火山 Seedream 尺寸(~2K 面积,各边夹在 [1024,4096] 且取 16 的倍数)。
|
||||
预设比例直接给好尺寸;自定义 W:H 按 2K 面积换算;解析不到回落 '2K'。"""
|
||||
@@ -2060,9 +2111,9 @@ def generate_base_asset(*, project, user, kind: str, prompt: str, label: str = "
|
||||
"model": model_config.name, "endpoint": model_config.endpoint, "prompt": gen_prompt,
|
||||
"kind": kind, "label": label or "", "group_id": str(group_id) if group_id else "",
|
||||
"use_edit": use_edit, "reference_image": ref_url, "model_routing_v1": True,
|
||||
# 角色立绘不再自动接力三视图;三视图只通过角色详情里的显式按钮生成。
|
||||
# 保留字段为 False,兼容旧前端/旧任务读取,但不允许再触发自动链路。
|
||||
"auto_triview": False,
|
||||
# 角色从这里生成时,立绘落库后由 worker 自动接力生成绑定它的三视图。
|
||||
# 非角色一律忽略该参数,避免商品/场景误入人物三视图链路。
|
||||
"auto_triview": bool(auto_triview and kind == BaseAssetGroup.Kind.PERSON),
|
||||
}
|
||||
task = create_ai_task(
|
||||
project=project,
|
||||
@@ -2189,6 +2240,21 @@ def run_base_asset_task(*, task_id: str) -> None:
|
||||
from apps.assets.review import submit_asset_for_review
|
||||
|
||||
transaction.on_commit(lambda a=asset: submit_asset_for_review(a))
|
||||
# 三视图是专业创作可选的增强资产。只有调用方明确要求时才接力,
|
||||
# 极速成片只需人物立绘即可进入故事板,不能为三视图额外等待或失败。
|
||||
if kind == BaseAssetGroup.Kind.PERSON and payload.get("auto_triview"):
|
||||
def _kickoff_person_triview(portrait=asset):
|
||||
try:
|
||||
generate_person_triview(
|
||||
project=project, user=user, portrait_asset=portrait
|
||||
)
|
||||
except Exception:
|
||||
logger.exception(
|
||||
"auto person triview kickoff failed for portrait %s",
|
||||
getattr(portrait, "id", ""),
|
||||
)
|
||||
|
||||
transaction.on_commit(_kickoff_person_triview)
|
||||
except Exception as exc: # noqa: BLE001 — 失败要退费并把错误记进 AITask 供前端轮询读取;不向上抛(避免 celery 重试二次扣费)
|
||||
task.status = AITask.Status.FAILED
|
||||
task.error_message = str(exc)
|
||||
@@ -2201,6 +2267,16 @@ def run_base_asset_task(*, task_id: str) -> None:
|
||||
)
|
||||
|
||||
|
||||
def _triview_reference_url(*, asset_id: str, fallback: str = "") -> str:
|
||||
"""三视图 worker 现取立绘可访问 URL。提交时写进 payload 的 TOS 签名链接会过期,
|
||||
立绘落库瞬间也可能还没签出 URL;执行时再取一次,避免「立绘成了、三视图没参考图」。"""
|
||||
if not asset_id:
|
||||
return str(fallback or "")
|
||||
portrait = Asset.objects.filter(id=asset_id, is_deleted=False).first()
|
||||
live = _asset_preview_url(portrait)
|
||||
return live or str(fallback or "")
|
||||
|
||||
|
||||
def generate_person_triview(*, project, user, portrait_asset) -> AITask:
|
||||
"""流程步骤4 · 据「某一版立绘资产」生成它配套的三视图(**异步**:image_edit 慢,交给 worker)。
|
||||
Web 请求只建 RESERVED 任务 + 预留额度后秒回;worker 内跑 image_edit 并把三视图归组(run_triview_task)。
|
||||
@@ -2214,7 +2290,7 @@ def generate_person_triview(*, project, user, portrait_asset) -> AITask:
|
||||
model_config = get_default_model(ModelConfig.Capability.IMAGE)
|
||||
if model_config is None:
|
||||
raise ValueError("no active image model configured")
|
||||
ref_url = _asset_preview_url(portrait_asset)
|
||||
ref_url = _triview_reference_url(asset_id=asset_key)
|
||||
# 人物三视图提示词:正文可在 admin「提示词」页改(无占位符)
|
||||
tri_prompt = render_prompt("person_triview", THREE_VIEW_PROMPT)
|
||||
portrait_label = ""
|
||||
@@ -2222,7 +2298,11 @@ def generate_person_triview(*, project, user, portrait_asset) -> AITask:
|
||||
meta = group.metadata or {}
|
||||
if meta.get("triview_of"):
|
||||
continue
|
||||
candidates = [str(value) for value in (group.candidate_assets or [])]
|
||||
# candidate_assets 是 Django 的 ManyRelatedManager,不能直接遍历;
|
||||
# 资产生成完成后这里会立即触发人物三视图,直接遍历会抛
|
||||
# “ManyRelatedManager object is not iterable”,进而让极速成片
|
||||
# 在所有基础资产已成功时被错误终止。
|
||||
candidates = [str(value) for value in group.candidate_assets.all()]
|
||||
if str(group.adopted_asset_id or "") == asset_key or asset_key in candidates:
|
||||
portrait_label = str(meta.get("label") or "")
|
||||
break
|
||||
@@ -2343,7 +2423,10 @@ def run_model_triview_task(*, task_id: str) -> None:
|
||||
payload = task.request_payload or {}
|
||||
model_id = str(payload.get("model_id") or "")
|
||||
portrait_asset_id = str(payload.get("portrait_asset_id") or "")
|
||||
ref_url = str(payload.get("reference_image") or "")
|
||||
ref_url = _triview_reference_url(
|
||||
asset_id=portrait_asset_id,
|
||||
fallback=str(payload.get("reference_image") or ""),
|
||||
)
|
||||
prompt = str(payload.get("prompt") or "")
|
||||
use_model_routing = bool(payload.get("model_routing_v1"))
|
||||
provider = None if use_model_routing else get_image_provider(task.model_config)
|
||||
@@ -2465,13 +2548,19 @@ def run_triview_task(*, task_id: str) -> None:
|
||||
user = task.created_by
|
||||
payload = task.request_payload or {}
|
||||
asset_key = str(payload.get("triview_of") or "")
|
||||
ref_url = str(payload.get("reference_image") or "")
|
||||
ref_url = _triview_reference_url(
|
||||
asset_id=asset_key,
|
||||
fallback=str(payload.get("reference_image") or ""),
|
||||
)
|
||||
prompt = str(payload.get("prompt") or THREE_VIEW_PROMPT)
|
||||
model_config = task.model_config
|
||||
use_model_routing = bool(payload.get("model_routing_v1"))
|
||||
provider = None if use_model_routing else get_image_provider(model_config)
|
||||
reservation = task.credit_reservation
|
||||
try:
|
||||
if not asset_key or not ref_url:
|
||||
raise ValueError("立绘参考图不可用,无法生成三视图")
|
||||
# 构图意图仍是 16:9;gpt-image 不认 1536x864,OpenAICompatibleProvider 会收成 1536x1024。
|
||||
tri_size = prompt_ratio_size("person_triview", "1536x864")
|
||||
|
||||
def _make():
|
||||
@@ -2603,7 +2692,8 @@ def build_storyboard_frame_prompt(project, segment, extra_prompt: str = "") -> s
|
||||
脚本=(_segment_script_text(segment) or f"第 {segment.sort_order + 1} 镜"),
|
||||
补充=(("\n" + extra_prompt.strip()) if extra_prompt else ""),
|
||||
)
|
||||
return "\n".join(line for line in rendered.split("\n") if line.strip())
|
||||
cleaned = "\n".join(line for line in rendered.split("\n") if line.strip())
|
||||
return _apply_storyboard_output_ratio(cleaned, project)
|
||||
|
||||
|
||||
def build_video_segment_prompt(project, video_segment, scene, refs, user_prompt: str = "") -> str:
|
||||
@@ -2813,7 +2903,8 @@ def build_storyboard_frame_prompt_refs(project, segment, refs: list[dict], extra
|
||||
脚本=(_segment_script_text(segment) or f"第 {segment.sort_order + 1} 镜"),
|
||||
补充=(("\n" + extra_prompt.strip()) if extra_prompt else ""),
|
||||
)
|
||||
return "\n".join(line for line in rendered.split("\n") if line.strip())
|
||||
cleaned = "\n".join(line for line in rendered.split("\n") if line.strip())
|
||||
return _apply_storyboard_output_ratio(cleaned, project)
|
||||
|
||||
|
||||
def _is_transient_error(exc: Exception) -> bool:
|
||||
@@ -2968,6 +3059,9 @@ def _storyboard_shot_worker(task_id, shot_id, user_id) -> None:
|
||||
model_config = task.model_config
|
||||
reservation = task.credit_reservation
|
||||
extra_prompt = (project.metadata or {}).get("storyboard_prompt", "") or ""
|
||||
spec = project_output_spec(project)
|
||||
frame_ratio = spec["aspect_ratio"]
|
||||
frame_size = _ratio_to_image_size(frame_ratio)
|
||||
task.status = AITask.Status.SUBMITTED
|
||||
task.save(update_fields=["status", "updated_at"])
|
||||
try:
|
||||
@@ -2990,9 +3084,9 @@ def _storyboard_shot_worker(task_id, shot_id, user_id) -> None:
|
||||
primary_model=model_config,
|
||||
prompt=frame_prompt,
|
||||
reference_images=ref_urls,
|
||||
aspect_ratio="9:16",
|
||||
edit_size="1024x1536",
|
||||
direct_size="1024x1536",
|
||||
aspect_ratio=frame_ratio,
|
||||
edit_size=frame_size,
|
||||
direct_size=frame_size,
|
||||
request_summary={
|
||||
"storyboard_shot": str(shot.id),
|
||||
"storyboard_sort_order": shot.sort_order,
|
||||
@@ -3006,7 +3100,7 @@ def _storyboard_shot_worker(task_id, shot_id, user_id) -> None:
|
||||
model=model_config.name,
|
||||
prompt=frame_prompt,
|
||||
images=ref_urls,
|
||||
size="1024x1536",
|
||||
size=frame_size,
|
||||
)
|
||||
)
|
||||
else:
|
||||
@@ -3270,9 +3364,14 @@ def submit_video_segment(
|
||||
user,
|
||||
prompt: str,
|
||||
model_config_id=None,
|
||||
aspect_ratio: str = "9:16",
|
||||
resolution: str = "720p",
|
||||
aspect_ratio: str | None = None,
|
||||
resolution: str | None = None,
|
||||
) -> VideoSegmentVersion | None:
|
||||
spec = project_output_spec(video_segment.project)
|
||||
aspect_ratio = str(aspect_ratio or spec["aspect_ratio"] or "9:16")
|
||||
resolution = str(resolution or spec["resolution"] or "720p").lower()
|
||||
if not model_config_id:
|
||||
model_config_id = spec["video_model_config_id"]
|
||||
model_config = None
|
||||
if model_config_id:
|
||||
model_config = (
|
||||
@@ -3392,6 +3491,7 @@ def submit_video_segment(
|
||||
)
|
||||
video_segment.status = VideoSegment.Status.RUNNING
|
||||
video_segment.save(update_fields=["status", "updated_at"])
|
||||
_sync_timeline_output_spec(project, aspect_ratio=aspect_ratio, resolution=resolution)
|
||||
return None
|
||||
except Exception as exc:
|
||||
public_error = classify_generation_error(
|
||||
|
||||
@@ -0,0 +1,34 @@
|
||||
from io import BytesIO
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
from django.test import SimpleTestCase
|
||||
|
||||
from apps.ai.providers.openai_compatible import OpenAICompatibleProvider, openai_image_size
|
||||
|
||||
|
||||
class OpenAICompatibleImageSizeTests(SimpleTestCase):
|
||||
def test_coerces_true_16_9_to_supported_landscape(self):
|
||||
self.assertEqual(openai_image_size("1536x864"), "1536x1024")
|
||||
self.assertEqual(openai_image_size("1024x1536"), "1024x1536")
|
||||
self.assertEqual(openai_image_size("1536x1024"), "1536x1024")
|
||||
self.assertEqual(openai_image_size("1024x1024"), "1024x1024")
|
||||
self.assertEqual(openai_image_size("512x768"), "1024x1536")
|
||||
|
||||
@patch("apps.ai.providers.openai_compatible.requests.post")
|
||||
@patch.object(OpenAICompatibleProvider, "media_to_bytes")
|
||||
def test_image_edit_sends_gateway_supported_size(self, media_to_bytes, post):
|
||||
media_to_bytes.return_value = (BytesIO(b"img"), "image/png")
|
||||
response = Mock()
|
||||
response.ok = True
|
||||
response.json.return_value = {"data": [{"b64_json": "x"}]}
|
||||
post.return_value = response
|
||||
provider = OpenAICompatibleProvider(base_url="https://img.example/v1", api_key="secret")
|
||||
|
||||
provider.image_edit(
|
||||
model="gpt-image-2",
|
||||
prompt="三视图",
|
||||
images=["http://example.test/portrait.png"],
|
||||
size="1536x864",
|
||||
)
|
||||
|
||||
self.assertEqual(post.call_args.kwargs["data"]["size"], "1536x1024")
|
||||
@@ -150,6 +150,24 @@ class ProjectPersonTriviewRoutingTests(TestCase):
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
|
||||
|
||||
@patch("apps.ai.services.create_ai_task")
|
||||
def test_submit_reads_candidate_assets_manager_to_keep_portrait_label(self, create_task):
|
||||
"""基础资产落库后自动接力三视图时,关联集合必须可正常读取。"""
|
||||
self.model(self.provider("project-tri-label", 20), "project-tri-label")
|
||||
group = BaseAssetGroup.objects.create(
|
||||
project=self.project,
|
||||
kind=BaseAssetGroup.Kind.PERSON,
|
||||
adopted_asset=self.portrait,
|
||||
metadata={"label": "专业测评师"},
|
||||
)
|
||||
group.candidate_assets.add(self.portrait)
|
||||
create_task.return_value = Mock(id="triview-task")
|
||||
|
||||
task = self.submit()
|
||||
|
||||
self.assertEqual(create_task.call_args.kwargs["request_payload"]["label"], "专业测评师")
|
||||
self.assertEqual(task, create_task.return_value)
|
||||
|
||||
def test_retry_then_dynamic_candidate_success_charges_once(self):
|
||||
primary = self.model(
|
||||
self.provider("project-tri-fallback-primary", 100),
|
||||
@@ -232,3 +250,37 @@ class ProjectPersonTriviewRoutingTests(TestCase):
|
||||
self.assertEqual(call.kwargs["image"], ["http://example.test/portrait.png"])
|
||||
self.assertEqual(call.kwargs["size"], "1536x864")
|
||||
self.assertEqual(task.model_attempts.get().public_model_name, primary.display_name)
|
||||
|
||||
def test_worker_refreshes_stale_reference_url_from_portrait(self):
|
||||
"""立绘成功后立刻出三视图:worker 必须用当前立绘 URL,不能沿用已过期的 payload 链接。"""
|
||||
self.model(self.provider("project-tri-refresh", 20), "project-tri-refresh")
|
||||
task = self.submit()
|
||||
payload = dict(task.request_payload)
|
||||
payload["reference_image"] = "http://expired.test/old-portrait.png"
|
||||
task.request_payload = payload
|
||||
task.save(update_fields=["request_payload", "updated_at"])
|
||||
|
||||
run_triview_task(task_id=str(task.id))
|
||||
|
||||
task.refresh_from_db()
|
||||
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
|
||||
primary = next(iter(self.provider_mocks.values()))
|
||||
self.assertEqual(
|
||||
primary.image_edit.call_args.kwargs["images"],
|
||||
["http://example.test/portrait.png"],
|
||||
)
|
||||
|
||||
def test_worker_fails_clearly_when_portrait_reference_is_missing(self):
|
||||
self.model(self.provider("project-tri-missing-ref", 20), "project-tri-missing-ref")
|
||||
self.portrait.files.all().delete()
|
||||
task = self.submit()
|
||||
payload = dict(task.request_payload)
|
||||
payload["reference_image"] = ""
|
||||
task.request_payload = payload
|
||||
task.save(update_fields=["request_payload", "updated_at"])
|
||||
|
||||
run_triview_task(task_id=str(task.id))
|
||||
|
||||
task.refresh_from_db()
|
||||
self.assertEqual(task.status, AITask.Status.FAILED)
|
||||
self.assertIn("立绘参考图不可用", task.error_message)
|
||||
|
||||
@@ -24,8 +24,8 @@ class ModelRoutingPolicyTests(SimpleTestCase):
|
||||
self.assertEqual(policy.jitter_ratio, 0.20)
|
||||
self.assertEqual(policy.text.retry_delays, (1.0, 3.0))
|
||||
self.assertEqual(policy.text.request_timeout, 120.0)
|
||||
self.assertEqual(policy.text.stream_timeout, 300.0)
|
||||
self.assertEqual(policy.text.total_timeout, 480.0)
|
||||
self.assertEqual(policy.text.stream_timeout, 1740.0)
|
||||
self.assertEqual(policy.text.total_timeout, 1800.0)
|
||||
self.assertEqual(policy.image.retry_delays, (3.0,))
|
||||
self.assertEqual(policy.image.request_timeout, 300.0)
|
||||
self.assertEqual(policy.image.total_timeout, 900.0)
|
||||
|
||||
@@ -185,6 +185,16 @@ class StoryboardRoutingTests(TestCase):
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
|
||||
|
||||
def test_wizard_aspect_ratio_drives_storyboard_size(self):
|
||||
self.project.metadata = {"wizard": {"aspect_ratio": "16:9", "resolution": "480p"}}
|
||||
self.project.save(update_fields=["metadata"])
|
||||
primary = self.model(self.provider("storyboard-wide", 20), "storyboard-wide")
|
||||
task = self.enqueue()
|
||||
attempt = task.model_attempts.get()
|
||||
self.assertEqual(attempt.request_summary["aspect_ratio"], "16:9")
|
||||
call = self.provider_mocks[primary.id].image_edit.call_args
|
||||
self.assertEqual(call.kwargs["size"], "1536x864")
|
||||
|
||||
def test_retry_then_dynamic_candidate_success_charges_once(self):
|
||||
primary = self.model(
|
||||
self.provider("storyboard-fallback-primary", 100),
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
import json
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
|
||||
from django.conf import settings
|
||||
from django.test import SimpleTestCase
|
||||
@@ -1544,6 +1545,7 @@ class StandaloneCategoryTests(TestCase):
|
||||
)
|
||||
|
||||
|
||||
@override_settings(CACHES={"default": {"BACKEND": "django.core.cache.backends.locmem.LocMemCache"}})
|
||||
class TriviewModelDecouplingTests(TestCase):
|
||||
"""项目角色三视图只归项目,不自动创建或更新团队模特。"""
|
||||
|
||||
@@ -1606,9 +1608,9 @@ class TriviewModelDecouplingTests(TestCase):
|
||||
model.refresh_from_db()
|
||||
self.assertEqual(model.triview_asset_id, existing_tri.id)
|
||||
|
||||
@patch("apps.ai.services.generate_person_triview")
|
||||
def test_person_base_asset_ignores_legacy_auto_triview_flag(self, kickoff_triview):
|
||||
"""即使旧前端/旧任务传 auto_triview=True,角色立绘完成后也不能自动接力三视图。"""
|
||||
@patch("apps.ai.tasks.generate_base_asset_task.delay")
|
||||
def test_person_base_asset_marks_auto_triview_for_worker(self, _enqueue_base_asset):
|
||||
"""角色立绘任务需标记自动接力三视图,实际接力只在立绘成功落库后发生。"""
|
||||
from apps.ai.services import generate_base_asset
|
||||
from apps.projects.models import BaseAssetGroup
|
||||
|
||||
@@ -1633,8 +1635,47 @@ class TriviewModelDecouplingTests(TestCase):
|
||||
)
|
||||
|
||||
task.refresh_from_db()
|
||||
self.assertFalse(task.request_payload.get("auto_triview"))
|
||||
kickoff_triview.assert_not_called()
|
||||
self.assertTrue(task.request_payload.get("auto_triview"))
|
||||
|
||||
@patch("apps.assets.review.submit_asset_for_review")
|
||||
@patch("apps.ai.services.generate_person_triview")
|
||||
@patch("apps.ai.services._store_generated_media")
|
||||
@patch("apps.ai.services.execute_routed_image_request")
|
||||
@patch("apps.ai.tasks.generate_base_asset_task.delay")
|
||||
def test_completed_auto_person_generation_queues_triview(
|
||||
self, _enqueue_base_asset, routed, store_media, kickoff_triview, _review
|
||||
):
|
||||
"""角色立绘成功落库后,worker 必须自动创建绑定该立绘的三视图任务。"""
|
||||
from apps.ai.services import generate_base_asset, run_base_asset_task
|
||||
from apps.projects.models import BaseAssetGroup
|
||||
|
||||
ModelConfig.objects.create(
|
||||
provider=ModelProvider.objects.create(name="auto-tri-success-provider", display_name="Auto Tri Success Provider"),
|
||||
name="auto-tri-success-img",
|
||||
display_name="Auto Tri Success Img",
|
||||
capability=ModelConfig.Capability.IMAGE,
|
||||
endpoint="images/generations",
|
||||
unit_price="1.0000",
|
||||
is_default=True,
|
||||
)
|
||||
routed.return_value = SimpleNamespace(value=({"data": [{"url": "http://x/person.png"}]}, "http://x/person.png"))
|
||||
store_media.return_value = self.portrait
|
||||
task = generate_base_asset(
|
||||
project=self.project,
|
||||
user=self.user,
|
||||
kind=BaseAssetGroup.Kind.PERSON,
|
||||
prompt="项目角色立绘",
|
||||
label="女主",
|
||||
# 专业模式显式要求时才接力三视图;极速成片传 False 会跳过。
|
||||
auto_triview=True,
|
||||
)
|
||||
|
||||
with self.captureOnCommitCallbacks(execute=True):
|
||||
run_base_asset_task(task_id=str(task.id))
|
||||
|
||||
kickoff_triview.assert_called_once_with(
|
||||
project=self.project, user=self.user, portrait_asset=self.portrait
|
||||
)
|
||||
|
||||
|
||||
class ModelLibraryTriviewTaskTests(TestCase):
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
import uuid
|
||||
|
||||
from django.core.exceptions import ObjectDoesNotExist
|
||||
from rest_framework import serializers
|
||||
|
||||
from apps.assets.serializers import AssetFileSerializer
|
||||
@@ -341,6 +342,7 @@ class QuickCreateJobSerializer(serializers.ModelSerializer):
|
||||
phase_index = serializers.SerializerMethodField()
|
||||
settings = serializers.SerializerMethodField()
|
||||
result = serializers.SerializerMethodField()
|
||||
error_message = serializers.SerializerMethodField()
|
||||
|
||||
class Meta:
|
||||
model = QuickCreateJob
|
||||
@@ -365,6 +367,15 @@ class QuickCreateJobSerializer(serializers.ModelSerializer):
|
||||
]
|
||||
read_only_fields = fields
|
||||
|
||||
def get_error_message(self, obj) -> str:
|
||||
from apps.projects.services.quick_create import _public_error
|
||||
|
||||
stored = obj.error_message or ""
|
||||
if obj.status != QuickCreateJob.Status.FAILED:
|
||||
return stored
|
||||
hidden = str((obj.metadata or {}).get("internal_error") or "")
|
||||
return _public_error(" ".join(part for part in (stored, hidden) if part)) or stored
|
||||
|
||||
def get_product_images(self, obj) -> list[dict]:
|
||||
product = getattr(obj.project, "product", None)
|
||||
if product is None:
|
||||
@@ -403,7 +414,8 @@ class QuickCreateJobSerializer(serializers.ModelSerializer):
|
||||
def get_result(self, obj) -> dict | None:
|
||||
project = obj.project
|
||||
settings = self.get_settings(obj)
|
||||
video_url = _final_video_url(project)
|
||||
final_video_url = _final_video_url(project)
|
||||
video_url = final_video_url
|
||||
segments = list(project.video_segments.all())
|
||||
if not video_url:
|
||||
for segment in sorted(segments, key=lambda item: item.sort_order):
|
||||
@@ -446,6 +458,7 @@ class QuickCreateJobSerializer(serializers.ModelSerializer):
|
||||
)
|
||||
return {
|
||||
"video_url": video_url,
|
||||
"final_video_url": final_video_url,
|
||||
"poster_url": poster_url,
|
||||
"duration_seconds": duration or 15,
|
||||
"aspect_ratio": settings["aspect_ratio"],
|
||||
@@ -479,6 +492,23 @@ class ScriptVersionSerializer(serializers.ModelSerializer):
|
||||
read_only_fields = fields
|
||||
|
||||
|
||||
def _quick_create_job(obj: Project):
|
||||
try:
|
||||
return obj.quick_create_job
|
||||
except ObjectDoesNotExist:
|
||||
return None
|
||||
|
||||
|
||||
def _quick_create_status(obj: Project) -> str:
|
||||
job = _quick_create_job(obj)
|
||||
return job.status if job else ""
|
||||
|
||||
|
||||
def _quick_create_job_id(obj: Project) -> str:
|
||||
job = _quick_create_job(obj)
|
||||
return str(job.id) if job else ""
|
||||
|
||||
|
||||
class ProjectListSerializer(serializers.ModelSerializer):
|
||||
"""列表/仪表盘/侧栏用的轻量项目序列化:不嵌套 阶段/片段/故事板/时间线(那些只详情页要)。
|
||||
脚本数/镜数走 annotate 计数(见 ProjectViewSet.get_queryset),避免逐项目拉全套关联(原列表 2-3s)。"""
|
||||
@@ -490,13 +520,15 @@ class ProjectListSerializer(serializers.ModelSerializer):
|
||||
# 合成成片地址:项目列表的播放按钮据此直接播成片(没合成过为空 → 退回进流水线)
|
||||
final_video_url = serializers.SerializerMethodField()
|
||||
quick_create = serializers.SerializerMethodField()
|
||||
quick_create_status = serializers.SerializerMethodField()
|
||||
quick_create_job_id = serializers.SerializerMethodField()
|
||||
|
||||
class Meta:
|
||||
model = Project
|
||||
fields = [
|
||||
"id", "name", "product", "product_title", "cover_preview_url",
|
||||
"status", "current_stage", "script_version_count", "video_segment_count",
|
||||
"final_video_url", "quick_create",
|
||||
"final_video_url", "quick_create", "quick_create_status", "quick_create_job_id",
|
||||
"is_deleted", "purged_at", "created_at", "updated_at",
|
||||
]
|
||||
|
||||
@@ -509,6 +541,12 @@ class ProjectListSerializer(serializers.ModelSerializer):
|
||||
def get_quick_create(self, obj) -> bool:
|
||||
return bool((obj.metadata or {}).get("quick_create"))
|
||||
|
||||
def get_quick_create_status(self, obj) -> str:
|
||||
return _quick_create_status(obj)
|
||||
|
||||
def get_quick_create_job_id(self, obj) -> str:
|
||||
return _quick_create_job_id(obj)
|
||||
|
||||
|
||||
class ProjectSerializer(serializers.ModelSerializer):
|
||||
stages = ProjectStageSerializer(many=True, read_only=True)
|
||||
@@ -520,6 +558,8 @@ class ProjectSerializer(serializers.ModelSerializer):
|
||||
timeline = TimelineSerializer(read_only=True)
|
||||
# 合成成片地址(最新一次成功拼接):视频阶段的「播放成片 / 下载成片」直接用它
|
||||
final_video_url = serializers.SerializerMethodField()
|
||||
quick_create_status = serializers.SerializerMethodField()
|
||||
quick_create_job_id = serializers.SerializerMethodField()
|
||||
|
||||
class Meta:
|
||||
model = Project
|
||||
@@ -542,6 +582,8 @@ class ProjectSerializer(serializers.ModelSerializer):
|
||||
"video_segments",
|
||||
"timeline",
|
||||
"final_video_url",
|
||||
"quick_create_status",
|
||||
"quick_create_job_id",
|
||||
"created_at",
|
||||
"updated_at",
|
||||
]
|
||||
@@ -550,6 +592,12 @@ class ProjectSerializer(serializers.ModelSerializer):
|
||||
def get_final_video_url(self, obj) -> str:
|
||||
return _final_video_url(obj)
|
||||
|
||||
def get_quick_create_status(self, obj) -> str:
|
||||
return _quick_create_status(obj)
|
||||
|
||||
def get_quick_create_job_id(self, obj) -> str:
|
||||
return _quick_create_job_id(obj)
|
||||
|
||||
|
||||
class ScriptTemplateSerializer(serializers.ModelSerializer):
|
||||
"""套路模板 · 列表与详情共用。写入只开放 name(其余字段由存模板端点从脚本抽)。"""
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -47,9 +47,15 @@ def run_export_job_task(self, export_job_id: str) -> str:
|
||||
QUICK_CREATE_QUEUE = "airshelf.quick"
|
||||
|
||||
|
||||
@app.task(bind=True, max_retries=0, soft_time_limit=240, time_limit=270, queue=QUICK_CREATE_QUEUE)
|
||||
# 豆包长思考脚本允许完整跑 30 分钟;硬上限额外留 60 秒让 soft timeout 的
|
||||
# 收尾/落库完成,避免 15 分钟时仍在正常输出却被 worker 强制中断。
|
||||
@app.task(bind=True, max_retries=0, soft_time_limit=1800, time_limit=1860, queue=QUICK_CREATE_QUEUE)
|
||||
def run_quick_script_task(self, quick_job_id: str) -> str:
|
||||
"""脚本生成单独跑,避免把整条极速成片编排堵在一次 SSE 消费里。"""
|
||||
"""脚本生成单独跑,避免把整条极速成片编排堵在一次 SSE 消费里。
|
||||
|
||||
软超时必须长于豆包思考流(允许最长 30 分钟)。短于 HTTP 流超时会 SIGUSR1 掐连接,
|
||||
任务监视器就记成 stream aborted (client disconnected)。
|
||||
"""
|
||||
from celery.exceptions import SoftTimeLimitExceeded
|
||||
|
||||
from apps.projects.models import QuickCreateJob
|
||||
@@ -57,41 +63,36 @@ def run_quick_script_task(self, quick_job_id: str) -> str:
|
||||
|
||||
try:
|
||||
consume_quick_script(quick_job_id)
|
||||
except SoftTimeLimitExceeded:
|
||||
except SoftTimeLimitExceeded as exc:
|
||||
job = QuickCreateJob.objects.select_related("project").filter(id=quick_job_id).first()
|
||||
if job is not None and job.status not in {
|
||||
QuickCreateJob.Status.SUCCEEDED,
|
||||
QuickCreateJob.Status.FAILED,
|
||||
QuickCreateJob.Status.CANCELLED,
|
||||
}:
|
||||
fail_quick_create(job, "脚本生成超时,请稍后重试或进入专业模式查看")
|
||||
raise
|
||||
fail_quick_create(
|
||||
job,
|
||||
"脚本生成时间较长,系统会自动重试",
|
||||
internal_error=f"SoftTimeLimitExceeded: stream aborted (client disconnected); {exc}",
|
||||
)
|
||||
return quick_job_id
|
||||
return quick_job_id
|
||||
|
||||
|
||||
@app.task(bind=True, max_retries=0, queue=QUICK_CREATE_QUEUE)
|
||||
@app.task(bind=True, max_retries=0)
|
||||
def advance_quick_create_task(self, quick_job_id: str) -> str:
|
||||
"""一次只推进一个可重入状态,等待型阶段通过重新入队轮询,不占 worker 睡眠。"""
|
||||
from apps.projects.services.quick_create import advance_quick_create
|
||||
"""一次只推进一个可重入状态,等待型阶段通过重新入队轮询,不占 worker 睡眠。
|
||||
|
||||
编排走默认 celery 队列:worker 即使没听 airshelf.quick,也不会停在「等待开始」。
|
||||
"""
|
||||
from apps.projects.models import QuickCreateJob
|
||||
from apps.projects.services.quick_create import _claim_next_advance, _enqueue_advance, advance_quick_create
|
||||
|
||||
next_delay = advance_quick_create(quick_job_id)
|
||||
if next_delay is not None:
|
||||
from apps.projects.services.quick_create import _claim_next_advance
|
||||
|
||||
if not _claim_next_advance(quick_job_id, int(next_delay)):
|
||||
return quick_job_id
|
||||
try:
|
||||
advance_quick_create_task.apply_async(
|
||||
args=[quick_job_id],
|
||||
countdown=max(1, int(next_delay)),
|
||||
queue=QUICK_CREATE_QUEUE,
|
||||
)
|
||||
except Exception as exc: # noqa: BLE001 — 重排失败不能留下永久“生成中”
|
||||
from apps.projects.models import QuickCreateJob
|
||||
from apps.projects.services.quick_create import fail_quick_create
|
||||
|
||||
job = QuickCreateJob.objects.select_related("project").filter(id=quick_job_id).first()
|
||||
if job is not None:
|
||||
fail_quick_create(job, "生成队列暂时中断,请稍后重试", internal_error=str(exc))
|
||||
raise
|
||||
job = QuickCreateJob.objects.filter(id=quick_job_id).first()
|
||||
if job is not None:
|
||||
_enqueue_advance(job, countdown=max(1, int(next_delay)))
|
||||
return quick_job_id
|
||||
|
||||
@@ -12,12 +12,14 @@ 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.products.models import Product, ProductImage
|
||||
from apps.projects.models import Project, ProjectStage, QuickCreateJob, ScriptSegment, ScriptVersion, VideoSegment, VideoSegmentVersion
|
||||
from apps.projects.models import BaseAssetGroup, Project, ProjectStage, QuickCreateJob, ScriptSegment, ScriptVersion, VideoSegment, VideoSegmentVersion
|
||||
from apps.projects.serializers import ProjectListSerializer, QuickCreateJobSerializer
|
||||
from apps.projects.services.pipeline import initialize_project_pipeline
|
||||
from apps.projects.services.quick_create import (
|
||||
REVIEW_FAIL_MESSAGE,
|
||||
_consume_script_agent,
|
||||
_reviews_ready,
|
||||
_safe_error,
|
||||
_start_videos,
|
||||
advance_quick_create,
|
||||
cancel_quick_create,
|
||||
@@ -94,7 +96,7 @@ class QuickCreateApiTests(TestCase):
|
||||
self.assertEqual(str(job.id), response.data["id"])
|
||||
enqueue.assert_called_once()
|
||||
self.assertEqual(enqueue.call_args.kwargs["args"], [str(job.id)])
|
||||
self.assertEqual(enqueue.call_args.kwargs["queue"], "airshelf.quick")
|
||||
self.assertIn(enqueue.call_args.kwargs["queue"], {"celery", "airshelf.quick"})
|
||||
require_worker_task.assert_called_once_with("apps.projects.tasks.advance_quick_create_task")
|
||||
self.assertEqual(get_model.call_count, 2)
|
||||
get_quick_model.assert_called_once()
|
||||
@@ -290,7 +292,7 @@ class QuickCreateApiTests(TestCase):
|
||||
self.assertEqual(job.status, QuickCreateJob.Status.RUNNING)
|
||||
enqueue.assert_called_once()
|
||||
|
||||
def test_history_lists_team_jobs_including_failed(self):
|
||||
def test_history_lists_completed_jobs_only(self):
|
||||
product = Product.objects.create(team=self.team, created_by=self.user, title="历史商品")
|
||||
project = Project.objects.create(team=self.team, created_by=self.user, product=product, name="历史商品 · 极速成片")
|
||||
QuickCreateJob.objects.create(
|
||||
@@ -334,12 +336,45 @@ class QuickCreateApiTests(TestCase):
|
||||
|
||||
response = self.client.get("/api/projects/quick-create-history/")
|
||||
self.assertEqual(response.status_code, 200)
|
||||
self.assertEqual(response.data["count"], 2)
|
||||
self.assertEqual(response.data["count"], 1)
|
||||
titles = [item["title"] for item in response.data["results"]]
|
||||
self.assertIn("历史商品 · 极速成片", titles)
|
||||
self.assertIn("失败商品 · 极速成片", titles)
|
||||
self.assertNotIn("失败商品 · 极速成片", titles)
|
||||
self.assertNotIn("进行中商品 · 极速成片", titles)
|
||||
|
||||
def test_history_hides_deleted_projects(self):
|
||||
product = Product.objects.create(team=self.team, created_by=self.user, title="还在的商品")
|
||||
alive = Project.objects.create(team=self.team, created_by=self.user, product=product, name="还在的商品 · 极速成片")
|
||||
QuickCreateJob.objects.create(
|
||||
team=self.team,
|
||||
created_by=self.user,
|
||||
project=alive,
|
||||
status=QuickCreateJob.Status.SUCCEEDED,
|
||||
phase=QuickCreateJob.Phase.COMPLETE,
|
||||
)
|
||||
deleted_product = Product.objects.create(team=self.team, created_by=self.user, title="已删商品")
|
||||
deleted = Project.objects.create(
|
||||
team=self.team,
|
||||
created_by=self.user,
|
||||
product=deleted_product,
|
||||
name="已删商品 · 极速成片",
|
||||
is_deleted=True,
|
||||
)
|
||||
QuickCreateJob.objects.create(
|
||||
team=self.team,
|
||||
created_by=self.user,
|
||||
project=deleted,
|
||||
status=QuickCreateJob.Status.CANCELLED,
|
||||
phase=QuickCreateJob.Phase.SCRIPT,
|
||||
)
|
||||
|
||||
response = self.client.get("/api/projects/quick-create-history/")
|
||||
self.assertEqual(response.status_code, 200)
|
||||
titles = [item["title"] for item in response.data["results"]]
|
||||
self.assertEqual(response.data["count"], 1)
|
||||
self.assertIn("还在的商品 · 极速成片", titles)
|
||||
self.assertNotIn("已删商品 · 极速成片", titles)
|
||||
|
||||
def test_list_serializer_flags_quick_create_projects(self):
|
||||
product = Product.objects.create(team=self.team, created_by=self.user, title="列表商品")
|
||||
quick = Project.objects.create(
|
||||
@@ -357,6 +392,44 @@ class QuickCreateApiTests(TestCase):
|
||||
)
|
||||
self.assertTrue(ProjectListSerializer(quick).data["quick_create"])
|
||||
self.assertFalse(ProjectListSerializer(normal).data["quick_create"])
|
||||
self.assertEqual(ProjectListSerializer(quick).data["quick_create_status"], "")
|
||||
self.assertEqual(ProjectListSerializer(quick).data["quick_create_job_id"], "")
|
||||
|
||||
QuickCreateJob.objects.create(team=self.team, created_by=self.user, project=quick, status=QuickCreateJob.Status.RUNNING)
|
||||
self.assertEqual(ProjectListSerializer(quick).data["quick_create_status"], "running")
|
||||
self.assertTrue(ProjectListSerializer(quick).data["quick_create_job_id"])
|
||||
|
||||
def test_running_quick_create_blocks_professional_edits(self):
|
||||
product = Product.objects.create(team=self.team, created_by=self.user, title="锁单商品")
|
||||
project = Project.objects.create(
|
||||
team=self.team,
|
||||
created_by=self.user,
|
||||
product=product,
|
||||
name="锁单商品 · 极速成片",
|
||||
metadata={"quick_create": True},
|
||||
)
|
||||
job = QuickCreateJob.objects.create(
|
||||
team=self.team,
|
||||
created_by=self.user,
|
||||
project=project,
|
||||
status=QuickCreateJob.Status.RUNNING,
|
||||
)
|
||||
|
||||
blocked = self.client.patch(f"/api/projects/{project.id}/", {"name": "不该改"}, format="json")
|
||||
self.assertEqual(blocked.status_code, 409)
|
||||
self.assertIn("极速成片", str(blocked.data))
|
||||
|
||||
allowed = self.client.get(f"/api/projects/{project.id}/")
|
||||
self.assertEqual(allowed.status_code, 200)
|
||||
self.assertEqual(allowed.data["quick_create_status"], "running")
|
||||
self.assertEqual(allowed.data["quick_create_job_id"], str(job.id))
|
||||
|
||||
job.status = QuickCreateJob.Status.FAILED
|
||||
job.save(update_fields=["status", "updated_at"])
|
||||
resumed = self.client.patch(f"/api/projects/{project.id}/", {"name": "专业模式可改"}, format="json")
|
||||
self.assertEqual(resumed.status_code, 200)
|
||||
project.refresh_from_db()
|
||||
self.assertEqual(project.name, "专业模式可改")
|
||||
|
||||
|
||||
class QuickCreateCoordinatorTests(TestCase):
|
||||
@@ -410,7 +483,7 @@ class QuickCreateCoordinatorTests(TestCase):
|
||||
self.assertTrue(self.job.metadata.get("script_started"))
|
||||
start_script.assert_called_once()
|
||||
self.assertEqual(start_script.call_args.kwargs["args"], [str(self.job.id)])
|
||||
self.assertEqual(start_script.call_args.kwargs["queue"], "airshelf.quick")
|
||||
self.assertIn(start_script.call_args.kwargs["queue"], {"celery", "airshelf.quick"})
|
||||
|
||||
def test_adopted_script_moves_to_assets_without_rerunning(self):
|
||||
script = ScriptVersion.objects.create(project=self.project, title="脚本", content="口播", is_adopted=True)
|
||||
@@ -431,8 +504,8 @@ class QuickCreateCoordinatorTests(TestCase):
|
||||
self.assertIsNone(delay)
|
||||
self.assertEqual(self.job.status, QuickCreateJob.Status.CANCELLED)
|
||||
|
||||
@patch("apps.projects.services.quick_create._run_quick_script_in_thread")
|
||||
def test_recover_reruns_script_locally_when_queue_drops_it(self, run_local):
|
||||
@patch("apps.projects.services.quick_create._enqueue_script")
|
||||
def test_recover_requeues_script_when_queue_drops_it(self, enqueue_script):
|
||||
self.job.status = QuickCreateJob.Status.RUNNING
|
||||
self.job.phase = QuickCreateJob.Phase.SCRIPT
|
||||
self.job.metadata = {
|
||||
@@ -443,9 +516,8 @@ class QuickCreateCoordinatorTests(TestCase):
|
||||
|
||||
recover_quick_create(self.job)
|
||||
self.job.refresh_from_db()
|
||||
self.assertTrue(self.job.metadata.get("script_local"))
|
||||
self.assertEqual(self.job.status, QuickCreateJob.Status.RUNNING)
|
||||
run_local.assert_called_once_with(str(self.job.id))
|
||||
enqueue_script.assert_called_once_with(str(self.job.id))
|
||||
|
||||
def test_recover_starts_stale_queued_job_without_waiting_for_quick_queue(self):
|
||||
QuickCreateJob.objects.filter(id=self.job.id).update(
|
||||
@@ -455,13 +527,13 @@ class QuickCreateCoordinatorTests(TestCase):
|
||||
updated_at=timezone.now() - timedelta(seconds=30),
|
||||
)
|
||||
self.job.refresh_from_db()
|
||||
with patch("apps.projects.services.quick_create._advance_without_quick_queue") as advance:
|
||||
with patch("apps.projects.services.quick_create._enqueue_advance") as enqueue:
|
||||
recover_quick_create(self.job)
|
||||
advance.assert_called_once_with(str(self.job.id))
|
||||
enqueue.assert_called_once()
|
||||
|
||||
@patch("apps.common.celery_health.worker_consumes_queue", return_value=False)
|
||||
@patch("apps.projects.services.quick_create._run_quick_script_in_thread")
|
||||
def test_advance_script_runs_locally_when_quick_queue_has_no_consumer(self, run_local, _listens):
|
||||
@patch("apps.projects.tasks.run_quick_script_task.apply_async")
|
||||
def test_advance_script_falls_back_to_celery_when_quick_queue_has_no_consumer(self, start_script, _listens):
|
||||
self.job.status = QuickCreateJob.Status.RUNNING
|
||||
self.job.phase = QuickCreateJob.Phase.SCRIPT
|
||||
self.job.save(update_fields=["status", "phase", "updated_at"])
|
||||
@@ -470,11 +542,11 @@ class QuickCreateCoordinatorTests(TestCase):
|
||||
self.job.refresh_from_db()
|
||||
self.assertEqual(delay, 5)
|
||||
self.assertTrue(self.job.metadata.get("script_started"))
|
||||
self.assertTrue(self.job.metadata.get("script_local"))
|
||||
run_local.assert_called_once_with(str(self.job.id))
|
||||
start_script.assert_called_once()
|
||||
self.assertEqual(start_script.call_args.kwargs["queue"], "celery")
|
||||
|
||||
@patch("apps.projects.services.quick_create._run_quick_script_in_thread")
|
||||
def test_recover_does_not_steal_script_when_model_call_is_inflight(self, run_local):
|
||||
@patch("apps.projects.services.quick_create._enqueue_script")
|
||||
def test_recover_does_not_steal_script_when_model_call_is_inflight(self, enqueue_script):
|
||||
self._script_task(AITask.Status.SUBMITTED, key="script-inflight-1")
|
||||
self.job.status = QuickCreateJob.Status.RUNNING
|
||||
self.job.phase = QuickCreateJob.Phase.SCRIPT
|
||||
@@ -486,9 +558,8 @@ class QuickCreateCoordinatorTests(TestCase):
|
||||
|
||||
recover_quick_create(self.job)
|
||||
self.job.refresh_from_db()
|
||||
self.assertFalse(self.job.metadata.get("script_local"))
|
||||
self.assertEqual(self.job.status, QuickCreateJob.Status.RUNNING)
|
||||
run_local.assert_not_called()
|
||||
enqueue_script.assert_not_called()
|
||||
|
||||
def test_advance_script_ignores_old_failed_task_while_latest_is_running(self):
|
||||
self._script_task(AITask.Status.FAILED, key="script-old-failed", error_message="stream aborted (client disconnected)")
|
||||
@@ -506,6 +577,25 @@ class QuickCreateCoordinatorTests(TestCase):
|
||||
self.assertEqual(delay, 5)
|
||||
self.assertEqual(self.job.status, QuickCreateJob.Status.RUNNING)
|
||||
|
||||
@patch("apps.projects.services.quick_create._enqueue_script")
|
||||
def test_advance_script_retries_retryable_failure_without_failing_job(self, enqueue_script):
|
||||
self._script_task(AITask.Status.FAILED, key="script-aborted", error_message="stream aborted (client disconnected)")
|
||||
self.job.status = QuickCreateJob.Status.RUNNING
|
||||
self.job.phase = QuickCreateJob.Phase.SCRIPT
|
||||
self.job.metadata = {
|
||||
"script_started": True,
|
||||
"script_started_at": timezone.now().isoformat(),
|
||||
}
|
||||
self.job.save(update_fields=["status", "phase", "metadata", "updated_at"])
|
||||
|
||||
delay = advance_quick_create(str(self.job.id))
|
||||
self.job.refresh_from_db()
|
||||
self.assertEqual(delay, 5)
|
||||
self.assertEqual(self.job.status, QuickCreateJob.Status.RUNNING)
|
||||
self.assertTrue(self.job.metadata.get("script_started"))
|
||||
self.assertEqual(self.job.metadata.get("transient_retries"), 1)
|
||||
enqueue_script.assert_called_once_with(str(self.job.id))
|
||||
|
||||
@patch("apps.projects.tasks.advance_quick_create_task.apply_async")
|
||||
def test_recover_restarts_script_after_stream_abort(self, enqueue):
|
||||
self.job.status = QuickCreateJob.Status.FAILED
|
||||
@@ -525,6 +615,42 @@ class QuickCreateCoordinatorTests(TestCase):
|
||||
self.assertEqual(self.job.metadata.get("transient_retries"), 1)
|
||||
enqueue.assert_called_once()
|
||||
|
||||
@patch("apps.projects.tasks.advance_quick_create_task.apply_async")
|
||||
def test_recover_retries_soft_time_limit_as_transient(self, enqueue):
|
||||
self.job.status = QuickCreateJob.Status.FAILED
|
||||
self.job.phase = QuickCreateJob.Phase.SCRIPT
|
||||
self.job.error_message = "脚本生成时间较长,系统会自动重试"
|
||||
self.job.metadata = {
|
||||
"script_started": True,
|
||||
"script_started_at": timezone.now().isoformat(),
|
||||
"internal_error": "SoftTimeLimitExceeded: stream aborted (client disconnected)",
|
||||
}
|
||||
self.job.save(update_fields=["status", "phase", "error_message", "metadata", "updated_at"])
|
||||
|
||||
recover_quick_create(self.job)
|
||||
self.job.refresh_from_db()
|
||||
self.assertEqual(self.job.status, QuickCreateJob.Status.RUNNING)
|
||||
self.assertFalse(self.job.metadata.get("script_started"))
|
||||
enqueue.assert_called_once()
|
||||
|
||||
@patch("apps.projects.tasks.advance_quick_create_task.apply_async")
|
||||
def test_recover_retries_failed_script_task_without_internal_error(self, enqueue):
|
||||
self._script_task(
|
||||
AITask.Status.FAILED,
|
||||
key="script-soft-limit",
|
||||
error_message="SoftTimeLimitExceeded()",
|
||||
)
|
||||
self.job.status = QuickCreateJob.Status.FAILED
|
||||
self.job.phase = QuickCreateJob.Phase.SCRIPT
|
||||
self.job.error_message = "脚本生成超时,请稍后重试或进入专业模式查看"
|
||||
self.job.metadata = {"script_started": True}
|
||||
self.job.save(update_fields=["status", "phase", "error_message", "metadata", "updated_at"])
|
||||
|
||||
recover_quick_create(self.job)
|
||||
self.job.refresh_from_db()
|
||||
self.assertEqual(self.job.status, QuickCreateJob.Status.RUNNING)
|
||||
enqueue.assert_called_once()
|
||||
|
||||
@patch("apps.projects.services.quick_create.get_quick_script_model")
|
||||
@patch("apps.projects.services.quick_create.stream_script_agent")
|
||||
def test_consume_script_does_not_hit_db_on_every_sse_frame(self, stream_fn, get_model):
|
||||
@@ -581,12 +707,29 @@ class QuickCreateCoordinatorTests(TestCase):
|
||||
self.assertEqual(self.project.status, Project.Status.VIDEOING)
|
||||
self.assertEqual(self.project.failure_reason, "")
|
||||
|
||||
@patch("apps.projects.tasks.advance_quick_create_task.apply_async")
|
||||
def test_recover_resumes_production_after_orchestrator_timeout(self, enqueue):
|
||||
self.project.status = Project.Status.VIDEOING
|
||||
self.project.current_stage = ProjectStage.Stage.VIDEO
|
||||
self.project.save(update_fields=["status", "current_stage", "updated_at"])
|
||||
self.job.status = QuickCreateJob.Status.FAILED
|
||||
self.job.phase = QuickCreateJob.Phase.PRODUCTION
|
||||
self.job.error_message = "脚本已保留,后续步骤遇到网络波动。点重试会从上次进度继续"
|
||||
self.job.metadata = {"storyboard_started": True, "internal_error": "Timeout reading from socket"}
|
||||
self.job.save(update_fields=["status", "phase", "error_message", "metadata", "updated_at"])
|
||||
|
||||
recover_quick_create(self.job)
|
||||
self.job.refresh_from_db()
|
||||
self.assertEqual(self.job.status, QuickCreateJob.Status.RUNNING)
|
||||
self.assertEqual(self.job.metadata.get("transient_retries"), 1)
|
||||
enqueue.assert_called_once()
|
||||
|
||||
@patch("apps.projects.tasks.advance_quick_create_task.apply_async")
|
||||
def test_resume_failed_production_job_keeps_progress(self, enqueue):
|
||||
self.job.status = QuickCreateJob.Status.FAILED
|
||||
self.job.phase = QuickCreateJob.Phase.PRODUCTION
|
||||
self.job.error_message = "极速成片暂未完成,请稍后重试或进入专业模式查看"
|
||||
self.job.metadata = {"storyboard_started": True, "transient_retries": 8}
|
||||
self.job.metadata = {"storyboard_started": True, "transient_retries": 8, "video_fail_retries": 2}
|
||||
self.job.save(update_fields=["status", "phase", "error_message", "metadata", "updated_at"])
|
||||
|
||||
resume_quick_create(self.job)
|
||||
@@ -594,6 +737,7 @@ class QuickCreateCoordinatorTests(TestCase):
|
||||
self.assertEqual(self.job.status, QuickCreateJob.Status.RUNNING)
|
||||
self.assertEqual(self.job.error_message, "")
|
||||
self.assertIsNone(self.job.metadata.get("transient_retries"))
|
||||
self.assertIsNone(self.job.metadata.get("video_fail_retries"))
|
||||
enqueue.assert_called_once()
|
||||
|
||||
@patch("apps.projects.services.quick_create.assets_client.is_enabled", return_value=True)
|
||||
@@ -610,6 +754,40 @@ class QuickCreateCoordinatorTests(TestCase):
|
||||
self.job.refresh_from_db()
|
||||
self.assertTrue(self.job.metadata.get("reviews_skipped"))
|
||||
|
||||
def test_safe_error_maps_image_moderation_to_review_copy(self):
|
||||
self.assertEqual(
|
||||
_safe_error(ValueError("400 moderation_blocked safety_violations=[sexual]")),
|
||||
REVIEW_FAIL_MESSAGE,
|
||||
)
|
||||
self.assertEqual(
|
||||
_safe_error(RuntimeError("InputImageSensitiveContentDetected")),
|
||||
REVIEW_FAIL_MESSAGE,
|
||||
)
|
||||
|
||||
@patch("apps.projects.services.quick_create.assets_client.is_enabled", return_value=True)
|
||||
@patch("apps.projects.services.quick_create.poll_team_reviews", return_value={})
|
||||
@patch(
|
||||
"apps.projects.services.quick_create.collect_video_review_blockers",
|
||||
return_value=[{"review_status": "failed", "name": "女主立绘"}],
|
||||
)
|
||||
def test_failed_asset_review_stops_job_with_clear_message(self, _blockers, _poll, _enabled):
|
||||
self.job.status = QuickCreateJob.Status.RUNNING
|
||||
self.job.phase = QuickCreateJob.Phase.PRODUCTION
|
||||
self.job.save(update_fields=["status", "phase", "updated_at"])
|
||||
self.assertIsNone(_reviews_ready(self.job))
|
||||
self.job.refresh_from_db()
|
||||
self.assertEqual(self.job.status, QuickCreateJob.Status.FAILED)
|
||||
self.assertEqual(self.job.error_message, REVIEW_FAIL_MESSAGE)
|
||||
|
||||
def test_failed_job_serializer_rewrites_hidden_moderation_error(self):
|
||||
self.job.status = QuickCreateJob.Status.FAILED
|
||||
self.job.phase = QuickCreateJob.Phase.PRODUCTION
|
||||
self.job.error_message = "极速成片暂未完成,请稍后重试或进入专业模式查看"
|
||||
self.job.metadata = {"internal_error": "400 moderation_blocked safety_violations=[sexual]"}
|
||||
self.job.save(update_fields=["status", "phase", "error_message", "metadata", "updated_at"])
|
||||
data = QuickCreateJobSerializer(self.job).data
|
||||
self.assertEqual(data["error_message"], REVIEW_FAIL_MESSAGE)
|
||||
|
||||
def test_recover_marks_success_when_video_already_finished(self):
|
||||
self.project.video_segments.exclude(sort_order=0).delete()
|
||||
segment = self.project.video_segments.get(sort_order=0)
|
||||
@@ -635,6 +813,24 @@ class QuickCreateCoordinatorTests(TestCase):
|
||||
self.assertEqual(self.job.status, QuickCreateJob.Status.SUCCEEDED)
|
||||
self.assertEqual(self.job.phase, QuickCreateJob.Phase.COMPLETE)
|
||||
|
||||
def test_recover_marks_success_when_professional_mode_already_completed(self):
|
||||
"""专业模式完成后,极速任务的旧失败状态必须自动被回收。"""
|
||||
self.job.status = QuickCreateJob.Status.FAILED
|
||||
self.job.phase = QuickCreateJob.Phase.PRODUCTION
|
||||
self.job.error_message = "极速成片暂未完成,请稍后重试或进入专业模式查看"
|
||||
self.job.save(update_fields=["status", "phase", "error_message", "updated_at"])
|
||||
self.project.status = Project.Status.COMPLETED
|
||||
self.project.current_stage = ProjectStage.Stage.VIDEO
|
||||
self.project.failure_reason = ""
|
||||
self.project.save(update_fields=["status", "current_stage", "failure_reason", "updated_at"])
|
||||
|
||||
recover_quick_create(self.job)
|
||||
|
||||
self.job.refresh_from_db()
|
||||
self.assertEqual(self.job.status, QuickCreateJob.Status.SUCCEEDED)
|
||||
self.assertEqual(self.job.phase, QuickCreateJob.Phase.COMPLETE)
|
||||
self.assertEqual(self.job.progress, 100)
|
||||
|
||||
@patch("apps.projects.services.quick_create.generate_base_asset", side_effect=TimeoutError("Timeout reading from socket"))
|
||||
def test_asset_start_timeout_retries_without_locking_the_job(self, _generate):
|
||||
script = ScriptVersion.objects.create(project=self.project, title="脚本", content="口播", is_adopted=True)
|
||||
@@ -680,6 +876,213 @@ class QuickCreateCoordinatorTests(TestCase):
|
||||
self.assertEqual(self.job.status, QuickCreateJob.Status.FAILED)
|
||||
self.assertNotEqual(self.project.status, Project.Status.FAILED)
|
||||
|
||||
def _ready_asset_tasks(self):
|
||||
script = ScriptVersion.objects.create(project=self.project, title="脚本", content="口播", is_adopted=True)
|
||||
ScriptSegment.objects.create(script_version=script, sort_order=0, narration="开场")
|
||||
ids = []
|
||||
for kind, key in (
|
||||
(AITask.Type.PRODUCT_IMAGE, "asset-product"),
|
||||
(AITask.Type.PERSON_IMAGE, "asset-person"),
|
||||
(AITask.Type.SCENE_IMAGE, "asset-scene"),
|
||||
):
|
||||
task = self._script_task(AITask.Status.SUCCEEDED, key=key)
|
||||
task.task_type = kind
|
||||
task.save(update_fields=["task_type"])
|
||||
ids.append(str(task.id))
|
||||
return ids
|
||||
|
||||
def _portrait_group(self, *, task=None, name="推荐模特"):
|
||||
portrait = Asset.objects.create(
|
||||
team=self.team,
|
||||
created_by=self.user,
|
||||
name=name,
|
||||
asset_type=Asset.Type.IMAGE,
|
||||
source=Asset.Source.AI_GENERATED,
|
||||
category=Asset.Category.PERSON,
|
||||
)
|
||||
return BaseAssetGroup.objects.create(
|
||||
project=self.project,
|
||||
kind=BaseAssetGroup.Kind.PERSON,
|
||||
task=task,
|
||||
adopted_asset=portrait,
|
||||
metadata={"label": name},
|
||||
), portrait
|
||||
|
||||
@patch("apps.projects.services.quick_create.generate_person_triview")
|
||||
def test_assets_skip_inflight_person_triview(self, start_triview):
|
||||
base_ids = self._ready_asset_tasks()
|
||||
_group, portrait = self._portrait_group()
|
||||
inflight = self._script_task(AITask.Status.SUBMITTED, key="auto-triview")
|
||||
inflight.request_payload = {"triview_of": str(portrait.id)}
|
||||
inflight.save(update_fields=["request_payload"])
|
||||
self.job.status = QuickCreateJob.Status.RUNNING
|
||||
self.job.phase = QuickCreateJob.Phase.ASSETS
|
||||
self.job.metadata = {"base_asset_task_ids": base_ids, "assets_started": True}
|
||||
self.job.save(update_fields=["status", "phase", "metadata", "updated_at"])
|
||||
|
||||
delay = advance_quick_create(str(self.job.id))
|
||||
self.job.refresh_from_db()
|
||||
start_triview.assert_not_called()
|
||||
self.assertEqual(delay, 1)
|
||||
self.assertIsNone(self.job.metadata.get("triview_task_ids"))
|
||||
self.assertTrue(self.job.metadata.get("triview_skipped"))
|
||||
self.assertEqual(self.job.status, QuickCreateJob.Status.RUNNING)
|
||||
self.assertEqual(self.job.phase, QuickCreateJob.Phase.PRODUCTION)
|
||||
|
||||
@patch("apps.projects.services.quick_create.generate_person_triview")
|
||||
def test_assets_never_starts_triview(self, start_triview):
|
||||
new_id = uuid.uuid4()
|
||||
start_triview.return_value = SimpleNamespace(id=new_id)
|
||||
base_ids = self._ready_asset_tasks()
|
||||
group, portrait = self._portrait_group()
|
||||
self.job.status = QuickCreateJob.Status.RUNNING
|
||||
self.job.phase = QuickCreateJob.Phase.ASSETS
|
||||
self.job.metadata = {"base_asset_task_ids": base_ids, "assets_started": True}
|
||||
self.job.save(update_fields=["status", "phase", "metadata", "updated_at"])
|
||||
|
||||
delay = advance_quick_create(str(self.job.id))
|
||||
self.job.refresh_from_db()
|
||||
start_triview.assert_not_called()
|
||||
self.assertEqual(delay, 1)
|
||||
self.assertIsNone(self.job.metadata.get("triview_task_ids"))
|
||||
self.assertTrue(self.job.metadata.get("triview_skipped"))
|
||||
self.assertEqual(self.job.phase, QuickCreateJob.Phase.PRODUCTION)
|
||||
self.assertIsNone(group.task_id)
|
||||
|
||||
def test_failed_triview_does_not_block_storyboard(self):
|
||||
base_ids = self._ready_asset_tasks()
|
||||
_group, portrait = self._portrait_group()
|
||||
failed = self._script_task(AITask.Status.FAILED, key="triview-failed", error_message="image_edit timeout")
|
||||
failed.request_payload = {"triview_of": str(portrait.id)}
|
||||
failed.save(update_fields=["request_payload"])
|
||||
self.job.status = QuickCreateJob.Status.RUNNING
|
||||
self.job.phase = QuickCreateJob.Phase.ASSETS
|
||||
self.job.metadata = {
|
||||
"base_asset_task_ids": base_ids,
|
||||
"assets_started": True,
|
||||
"triview_task_ids": [str(failed.id)],
|
||||
"triview_fail_retries": 2,
|
||||
}
|
||||
self.job.save(update_fields=["status", "phase", "metadata", "updated_at"])
|
||||
|
||||
delay = advance_quick_create(str(self.job.id))
|
||||
self.job.refresh_from_db()
|
||||
self.assertEqual(delay, 1)
|
||||
self.assertEqual(self.job.phase, QuickCreateJob.Phase.PRODUCTION)
|
||||
self.assertTrue(self.job.metadata.get("triview_skipped"))
|
||||
self.assertEqual(self.job.status, QuickCreateJob.Status.RUNNING)
|
||||
|
||||
@patch("apps.projects.services.quick_create.generate_base_asset")
|
||||
def test_assets_start_all_characters_and_scenes(self, generate):
|
||||
generate.side_effect = lambda **kwargs: SimpleNamespace(id=uuid.uuid4())
|
||||
ScriptVersion.objects.create(project=self.project, title="脚本", content="口播", is_adopted=True)
|
||||
self.project.metadata = {
|
||||
"script_entities": [
|
||||
{"id": "c1", "type": "character", "name": "女主", "visual_prompt": "都市女性"},
|
||||
{"id": "c2", "type": "character", "name": "闺蜜", "visual_prompt": "活泼女生"},
|
||||
{"id": "s1", "type": "scene", "name": "客厅", "visual_prompt": "暖光客厅"},
|
||||
{"id": "s2", "type": "scene", "name": "咖啡馆", "visual_prompt": "街边咖啡馆"},
|
||||
]
|
||||
}
|
||||
self.project.save(update_fields=["metadata", "updated_at"])
|
||||
self.job.status = QuickCreateJob.Status.RUNNING
|
||||
self.job.phase = QuickCreateJob.Phase.ASSETS
|
||||
self.job.save(update_fields=["status", "phase", "updated_at"])
|
||||
|
||||
delay = advance_quick_create(str(self.job.id))
|
||||
self.job.refresh_from_db()
|
||||
self.assertEqual(delay, 10)
|
||||
self.assertEqual(generate.call_count, 5)
|
||||
labels = [(call.kwargs["kind"], call.kwargs["label"]) for call in generate.call_args_list]
|
||||
self.assertEqual(
|
||||
labels,
|
||||
[
|
||||
(BaseAssetGroup.Kind.PRODUCT, "测试精华"),
|
||||
(BaseAssetGroup.Kind.PERSON, "女主"),
|
||||
(BaseAssetGroup.Kind.PERSON, "闺蜜"),
|
||||
(BaseAssetGroup.Kind.SCENE, "客厅"),
|
||||
(BaseAssetGroup.Kind.SCENE, "咖啡馆"),
|
||||
],
|
||||
)
|
||||
for call in generate.call_args_list:
|
||||
if call.kwargs["kind"] == BaseAssetGroup.Kind.PERSON:
|
||||
self.assertFalse(call.kwargs["auto_triview"])
|
||||
|
||||
@patch("apps.projects.services.quick_create.generate_person_triview")
|
||||
def test_assets_skip_triview_for_every_character(self, start_triview):
|
||||
start_triview.side_effect = lambda **kwargs: SimpleNamespace(id=uuid.uuid4())
|
||||
base_ids = self._ready_asset_tasks()
|
||||
_g1, portrait_a = self._portrait_group(name="女主")
|
||||
_g2, portrait_b = self._portrait_group(name="闺蜜")
|
||||
self.job.status = QuickCreateJob.Status.RUNNING
|
||||
self.job.phase = QuickCreateJob.Phase.ASSETS
|
||||
self.job.metadata = {"base_asset_task_ids": base_ids, "assets_started": True}
|
||||
self.job.save(update_fields=["status", "phase", "metadata", "updated_at"])
|
||||
|
||||
delay = advance_quick_create(str(self.job.id))
|
||||
self.job.refresh_from_db()
|
||||
self.assertEqual(delay, 1)
|
||||
start_triview.assert_not_called()
|
||||
self.assertIsNone(self.job.metadata.get("triview_task_ids"))
|
||||
self.assertTrue(self.job.metadata.get("triview_skipped"))
|
||||
self.assertEqual(self.job.phase, QuickCreateJob.Phase.PRODUCTION)
|
||||
self.assertEqual(self.job.status, QuickCreateJob.Status.RUNNING)
|
||||
|
||||
@patch("apps.projects.services.quick_create.generate_person_triview", side_effect=ValueError("no active image model configured"))
|
||||
def test_triview_submit_error_does_not_fail_job(self, _start_triview):
|
||||
base_ids = self._ready_asset_tasks()
|
||||
self._portrait_group()
|
||||
self.job.status = QuickCreateJob.Status.RUNNING
|
||||
self.job.phase = QuickCreateJob.Phase.ASSETS
|
||||
self.job.metadata = {"base_asset_task_ids": base_ids, "assets_started": True}
|
||||
self.job.save(update_fields=["status", "phase", "metadata", "updated_at"])
|
||||
|
||||
delay = advance_quick_create(str(self.job.id))
|
||||
self.job.refresh_from_db()
|
||||
self.assertEqual(delay, 1)
|
||||
self.assertEqual(self.job.status, QuickCreateJob.Status.RUNNING)
|
||||
self.assertEqual(self.job.phase, QuickCreateJob.Phase.PRODUCTION)
|
||||
self.assertTrue(self.job.metadata.get("triview_skipped"))
|
||||
|
||||
def test_failed_triview_task_in_base_ids_does_not_fail_assets(self):
|
||||
base_ids = self._ready_asset_tasks()
|
||||
_group, portrait = self._portrait_group()
|
||||
failed = self._script_task(AITask.Status.FAILED, key="triview-in-base", error_message="image_edit timeout")
|
||||
failed.request_payload = {"kind": "person", "label": "推荐模特", "triview_of": str(portrait.id)}
|
||||
failed.save(update_fields=["request_payload"])
|
||||
self.job.status = QuickCreateJob.Status.RUNNING
|
||||
self.job.phase = QuickCreateJob.Phase.ASSETS
|
||||
self.job.metadata = {"base_asset_task_ids": [*base_ids, str(failed.id)], "assets_started": True}
|
||||
self.job.save(update_fields=["status", "phase", "metadata", "updated_at"])
|
||||
|
||||
delay = advance_quick_create(str(self.job.id))
|
||||
self.job.refresh_from_db()
|
||||
self.assertEqual(delay, 1)
|
||||
self.assertEqual(self.job.status, QuickCreateJob.Status.RUNNING)
|
||||
self.assertEqual(self.job.phase, QuickCreateJob.Phase.PRODUCTION)
|
||||
|
||||
@patch("apps.projects.tasks.advance_quick_create_task.apply_async")
|
||||
def test_resume_failed_assets_continues_into_storyboard(self, enqueue):
|
||||
base_ids = self._ready_asset_tasks()
|
||||
self.job.status = QuickCreateJob.Status.FAILED
|
||||
self.job.phase = QuickCreateJob.Phase.ASSETS
|
||||
self.job.error_message = "极速成片暂未完成,请稍后重试或进入专业模式查看"
|
||||
self.job.metadata = {
|
||||
"base_asset_task_ids": base_ids,
|
||||
"assets_started": True,
|
||||
"triview_task_ids": [],
|
||||
"triview_ready": True,
|
||||
"triview_skipped": True,
|
||||
}
|
||||
self.job.save(update_fields=["status", "phase", "error_message", "metadata", "updated_at"])
|
||||
|
||||
resume_quick_create(self.job)
|
||||
self.job.refresh_from_db()
|
||||
self.assertEqual(self.job.status, QuickCreateJob.Status.RUNNING)
|
||||
self.assertEqual(self.job.phase, QuickCreateJob.Phase.PRODUCTION)
|
||||
self.assertEqual(self.job.error_message, "")
|
||||
enqueue.assert_called_once()
|
||||
|
||||
@patch("apps.projects.services.quick_create._reviews_ready", return_value=True)
|
||||
def test_production_counts_ready_videos_without_adding_version_ids(self, _reviews):
|
||||
self.project.video_segments.exclude(sort_order=0).delete()
|
||||
@@ -739,6 +1142,7 @@ class QuickCreateCoordinatorTests(TestCase):
|
||||
data = QuickCreateJobSerializer(self.job).data
|
||||
self.assertEqual(data["phase_index"], 3)
|
||||
self.assertEqual(data["result"]["video_url"], "https://cdn.example/quick.mp4")
|
||||
self.assertEqual(data["result"]["final_video_url"], "")
|
||||
self.assertEqual(data["result"]["duration_seconds"], 15)
|
||||
|
||||
def test_phase_index_matches_four_step_ui(self):
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
from django.test import TestCase
|
||||
from django.test import TestCase, override_settings
|
||||
from unittest.mock import patch
|
||||
|
||||
from rest_framework.test import APIClient
|
||||
@@ -351,9 +351,11 @@ class ProjectApiTests(TestCase):
|
||||
group = BaseAssetGroup.objects.get(project=project, kind=BaseAssetGroup.Kind.PERSON)
|
||||
self.assertEqual(group.metadata.get("label"), "女主")
|
||||
|
||||
@override_settings(CACHES={"default": {"BACKEND": "django.core.cache.backends.locmem.LocMemCache"}})
|
||||
@patch("apps.ai.tasks.generate_base_asset_task.delay")
|
||||
@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):
|
||||
def test_generate_person_base_asset_always_enables_auto_triview(self, get_provider, store_media, _enqueue_base_asset):
|
||||
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",
|
||||
@@ -370,13 +372,13 @@ class ProjectApiTests(TestCase):
|
||||
|
||||
response = self.client.post(
|
||||
f"/api/projects/{project.id}/generate-base-asset/",
|
||||
{"kind": "person", "prompt": "portrait", "label": "hero", "auto_triview": True},
|
||||
{"kind": "person", "prompt": "portrait", "label": "hero"},
|
||||
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"))
|
||||
self.assertTrue(task.request_payload.get("auto_triview"))
|
||||
|
||||
@patch("apps.ai.services._store_generated_media")
|
||||
@patch("apps.ai.services.get_image_provider")
|
||||
@@ -1348,6 +1350,13 @@ class VideoSegmentTrueUpTests(TestCase):
|
||||
self.assertEqual(task.credit_reservation.amount, video_reserve_amount(quote.points))
|
||||
self.assertEqual(task.request_payload["estimated_tokens"], tokens)
|
||||
|
||||
def test_submit_reads_wizard_output_spec(self):
|
||||
self.project.metadata = {"wizard": {"aspect_ratio": "16:9", "resolution": "480p"}}
|
||||
self.project.save(update_fields=["metadata"])
|
||||
task = self._submit()
|
||||
self.assertEqual(task.request_payload["ratio"], "16:9")
|
||||
self.assertEqual(task.request_payload["resolution"], "480p")
|
||||
|
||||
@patch("apps.ai.services._store_generated_media")
|
||||
def test_poll_settles_by_actual_usage_tokens(self, store):
|
||||
from decimal import Decimal
|
||||
|
||||
@@ -8,7 +8,7 @@ from django.http import HttpResponse, JsonResponse, StreamingHttpResponse
|
||||
from django.utils import timezone
|
||||
from rest_framework import status
|
||||
from rest_framework.decorators import action
|
||||
from rest_framework.exceptions import ValidationError
|
||||
from rest_framework.exceptions import APIException, ValidationError
|
||||
from rest_framework.parsers import FormParser, MultiPartParser
|
||||
from rest_framework.renderers import BaseRenderer
|
||||
from rest_framework.response import Response
|
||||
@@ -95,6 +95,12 @@ from .tasks import poll_video_segment_task
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class QuickCreateInProgress(APIException):
|
||||
status_code = status.HTTP_409_CONFLICT
|
||||
default_detail = "该项目正在极速成片中,请到极速成片页查看进度"
|
||||
default_code = "quick_create_running"
|
||||
|
||||
|
||||
class ServerSentEventRenderer(BaseRenderer):
|
||||
"""让 DRF 内容协商接受 Accept: text/event-stream(否则流式端点直接 406)。
|
||||
实际响应由视图返回 StreamingHttpResponse 直接下发,这个 renderer 只用于通过协商。"""
|
||||
@@ -235,7 +241,7 @@ def settle_video_completion(project: Project) -> bool:
|
||||
|
||||
|
||||
class ProjectViewSet(TeamScopedViewSetMixin, ModelViewSet):
|
||||
queryset = Project.objects.select_related("product", "timeline").prefetch_related(
|
||||
queryset = Project.objects.select_related("product", "timeline", "quick_create_job").prefetch_related(
|
||||
"stages",
|
||||
"video_segments",
|
||||
"video_segments__adopted_version__asset__files",
|
||||
@@ -339,7 +345,7 @@ class ProjectViewSet(TeamScopedViewSetMixin, ModelViewSet):
|
||||
# ——原列表把每个项目的 阶段/片段/故事板/时间线/资产文件全拉出,20 个项目实测 ~2s。
|
||||
if self.action == "list":
|
||||
qs = (
|
||||
Project.objects.select_related("product", "product__cover_asset", "timeline")
|
||||
Project.objects.select_related("product", "product__cover_asset", "timeline", "quick_create_job")
|
||||
.prefetch_related(
|
||||
"product__cover_asset__files",
|
||||
# 成片地址(final_video_url)只需要「成功的导出任务」,预取到位后列表不再逐项目查库
|
||||
@@ -363,6 +369,21 @@ class ProjectViewSet(TeamScopedViewSetMixin, ModelViewSet):
|
||||
).order_by("-updated_at")
|
||||
return super().get_queryset().filter(is_deleted=False, purged_at__isnull=True)
|
||||
|
||||
def initial(self, request, *args, **kwargs):
|
||||
super().initial(request, *args, **kwargs)
|
||||
if request.method in ("GET", "HEAD", "OPTIONS"):
|
||||
return
|
||||
if self.action in {"create", "destroy"}:
|
||||
return
|
||||
pk = kwargs.get("pk")
|
||||
if not pk:
|
||||
return
|
||||
if QuickCreateJob.objects.filter(
|
||||
project_id=pk,
|
||||
status__in=[QuickCreateJob.Status.QUEUED, QuickCreateJob.Status.RUNNING],
|
||||
).exists():
|
||||
raise QuickCreateInProgress()
|
||||
|
||||
def perform_destroy(self, instance):
|
||||
instance.is_deleted = True
|
||||
instance.save(update_fields=["is_deleted", "updated_at"])
|
||||
@@ -730,30 +751,19 @@ class ProjectViewSet(TeamScopedViewSetMixin, ModelViewSet):
|
||||
job = self._quick_job_queryset().get(id=job.id)
|
||||
return Response(QuickCreateJobSerializer(job).data)
|
||||
|
||||
@action(detail=False, methods=["post"], url_path=r"quick-create-retry/(?P<job_id>[^/.]+)")
|
||||
def quick_create_retry(self, request, job_id=None):
|
||||
job = self._quick_job_queryset().filter(id=job_id).first()
|
||||
if job is None:
|
||||
return Response({"detail": "极速成片任务不存在"}, status=status.HTTP_404_NOT_FOUND)
|
||||
if job.status == QuickCreateJob.Status.SUCCEEDED:
|
||||
return Response({"detail": "任务已经完成"}, status=status.HTTP_400_BAD_REQUEST)
|
||||
if job.status == QuickCreateJob.Status.CANCELLED:
|
||||
return Response({"detail": "已取消的任务请重新开始"}, status=status.HTTP_400_BAD_REQUEST)
|
||||
from .services.quick_create import resume_quick_create
|
||||
|
||||
resume_quick_create(job)
|
||||
job = self._quick_job_queryset().get(id=job.id)
|
||||
return Response(QuickCreateJobSerializer(job).data)
|
||||
|
||||
@action(detail=False, methods=["get"], url_path="quick-create-history")
|
||||
def quick_create_history(self, request):
|
||||
from .services.quick_create import restore_false_failed_quick_creates
|
||||
|
||||
restore_false_failed_quick_creates(self.get_team())
|
||||
# 进行中的任务看上方状态卡;列表要能找回失败后去专业模式继续的项目。
|
||||
# 进行中 / 未完成的任务回填上方表单;过往列表只放已完成成片。
|
||||
jobs = (
|
||||
self._quick_job_queryset()
|
||||
.exclude(status__in=[QuickCreateJob.Status.QUEUED, QuickCreateJob.Status.RUNNING])
|
||||
.filter(
|
||||
project__is_deleted=False,
|
||||
project__purged_at__isnull=True,
|
||||
status=QuickCreateJob.Status.SUCCEEDED,
|
||||
)
|
||||
.order_by("-created_at")
|
||||
)
|
||||
return Response({
|
||||
@@ -975,8 +985,9 @@ class ProjectViewSet(TeamScopedViewSetMixin, ModelViewSet):
|
||||
prompt=request.data.get("prompt", ""),
|
||||
label=request.data.get("label", ""),
|
||||
reference_asset_id=request.data.get("reference_asset_id") or None,
|
||||
# 角色立绘不再自动接力三视图;三视图只由角色详情里的显式按钮生成。
|
||||
auto_triview=False,
|
||||
# 用户点击角色 AI 生成 = 生成立绘并在完成后自动接力三视图。
|
||||
# 三视图任务由 worker 创建,页面刷新或离开也不会漏掉。
|
||||
auto_triview=kind == BaseAssetGroup.Kind.PERSON,
|
||||
)
|
||||
except ValueError as exc: # 无可用模型 / 余额不足等,立即反馈
|
||||
internal_kind = "user_credit_insufficient" if str(exc).strip().lower() == "insufficient credit" else ""
|
||||
@@ -1453,7 +1464,14 @@ class ProjectViewSet(TeamScopedViewSetMixin, ModelViewSet):
|
||||
# 火山/中转报错(如人脸需走素材库的 InputImageSensitiveContentDetected)→ 返回真实报错的 JSON,
|
||||
# 任务仍保留原始错误供排障,但普通用户仅收到安全错误对象,不让 500 HTML 导致前端白屏。
|
||||
try:
|
||||
submit_video_segment(video_segment=segment, user=request.user, prompt=request.data.get("prompt", ""))
|
||||
submit_video_segment(
|
||||
video_segment=segment,
|
||||
user=request.user,
|
||||
prompt=request.data.get("prompt", ""),
|
||||
model_config_id=request.data.get("model_config_id") or None,
|
||||
aspect_ratio=request.data.get("aspect_ratio") or None,
|
||||
resolution=request.data.get("resolution") or None,
|
||||
)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
public_error = classify_generation_error(exc, operation="video_generate")
|
||||
return Response(
|
||||
|
||||
Reference in New Issue
Block a user