完善视频复刻

This commit is contained in:
Azmat@qq.com
2026-08-27 15:56:44 +08:00
parent 3fbc2f2dab
commit d2b786a3cf
10 changed files with 722 additions and 164 deletions
+140 -10
View File
@@ -47,6 +47,7 @@ RATIOS = {"21:9", "16:9", "4:3", "1:1", "3:4", "9:16"}
RESOLUTIONS = {"480p", "720p", "1080p", "4k"} RESOLUTIONS = {"480p", "720p", "1080p", "4k"}
MODES = {"universal", "keyframe"} MODES = {"universal", "keyframe"}
IN_FLIGHT_STATUSES = ( IN_FLIGHT_STATUSES = (
AITask.Status.CREATED, # 视频复刻审核中也占并发,避免连点刷出一堆待审任务
AITask.Status.RESERVED, AITask.Status.RESERVED,
AITask.Status.SUBMITTED, AITask.Status.SUBMITTED,
AITask.Status.POLLING, AITask.Status.POLLING,
@@ -296,15 +297,16 @@ def build_content_items(*, team, prompt: str, mode: str, references: list) -> di
label_to_placeholder[label] = _placeholder_for(asset_type) label_to_placeholder[label] = _placeholder_for(asset_type)
continue continue
# 直传素材(已上传 TOS 的直链) # 直传素材(已上传 TOS 的直链)。resolved_url 可覆盖为 asset://(官方模特跨团队等)
push_url = str(ref.get("resolved_url") or url)
if ref_type == "image": if ref_type == "image":
# 参考图模式下所有图 role 必须 reference_image;keyframe 用 first_frame/last_frame # 参考图模式下所有图 role 必须 reference_image;keyframe 用 first_frame/last_frame
effective_role = "reference_image" if mode == "universal" else (role or "first_frame") effective_role = "reference_image" if mode == "universal" else (role or "first_frame")
asset_type = _push("image", url, effective_role) asset_type = _push("image", push_url, effective_role)
elif ref_type == "video": elif ref_type == "video":
asset_type = _push("video", url, role or "reference_video", duration) asset_type = _push("video", push_url, role or "reference_video", duration)
elif ref_type == "audio": elif ref_type == "audio":
asset_type = _push("audio", url, role or "reference_audio", duration) asset_type = _push("audio", push_url, role or "reference_audio", duration)
else: else:
logger.warning("unknown ref_type=%s url=%s label=%s, skipped", ref_type, url, label) logger.warning("unknown ref_type=%s url=%s label=%s, skipped", ref_type, url, label)
continue continue
@@ -366,6 +368,11 @@ def _reap_stale_free_video_tasks(*, team) -> None:
video_policy = load_model_routing_policy().video video_policy = load_model_routing_policy().video
buckets = [ buckets = [
(
[AITask.Status.CREATED],
{"updated_at__lt": now - timedelta(minutes=16)},
"素材审核超时(自动回收)",
),
( (
[AITask.Status.RESERVED], [AITask.Status.RESERVED],
{"updated_at__lt": now - timedelta(seconds=video_policy.submit_total_timeout)}, {"updated_at__lt": now - timedelta(seconds=video_policy.submit_total_timeout)},
@@ -538,7 +545,133 @@ def submit_free_video(*, team, user, params: dict) -> AITask:
task.status = AITask.Status.RESERVED task.status = AITask.Status.RESERVED
task.save(update_fields=["status", "updated_at"]) task.save(update_fields=["status", "updated_at"])
# 火山调用在事务外(不持锁调外网) return _dispatch_free_video_provider(
task=task,
built=built,
model_config=model_config,
aspect_ratio=aspect_ratio,
duration=duration,
resolution=resolution,
generate_audio=generate_audio,
seed=seed,
search_mode=search_mode,
feature=feature,
mode=mode,
poll_countdown=30,
)
def start_pending_free_video(task: AITask) -> AITask:
"""CREATED 任务审核已过后:预留积分 + 调火山。并发 poll 用行锁认领,失败不扣费。"""
if task.status != AITask.Status.CREATED:
return task
payload = task.request_payload or {}
prompt = str(payload.get("prompt") or "").strip()
mode = str(payload.get("mode") or "universal")
aspect_ratio = str(payload.get("aspect_ratio") or "16:9")
resolution = str(payload.get("resolution") or "720p")
generate_audio = bool(payload.get("generate_audio", True))
search_mode = str(payload.get("search_mode") or "off")
feature = str(payload.get("feature") or "free_video")
try:
duration = int(payload.get("duration") or 5)
except (TypeError, ValueError):
duration = 5
try:
seed = int(payload.get("seed") if payload.get("seed") is not None else -1)
except (TypeError, ValueError):
seed = -1
references = payload.get("references") or []
try:
built = build_content_items(team=task.team, prompt=prompt, mode=mode, references=references)
except ValueError as exc:
message = str(exc)
if "正在审核" in message or "已提交审核" in message or "尚未完成合规审核" in message:
return task
return _fail_pending_free_video(task, message)
tokens, quote = quote_video_estimate(
task.model_config,
aspect_ratio=aspect_ratio,
resolution=resolution,
duration=duration,
references=built["snapshots"],
team=task.team,
)
reserve_amount = video_reserve_amount(quote.points)
with transaction.atomic():
locked = (
AITask.objects.select_for_update()
.select_related("model_config", "model_config__provider", "team", "created_by")
.get(id=task.id)
)
if locked.status != AITask.Status.CREATED:
return locked
try:
reserve_credit(team=locked.team, user=locked.created_by, task=locked, amount=reserve_amount)
except ValueError as exc:
message = "团队余额不足,请充值后重试" if "insufficient credit" in str(exc) else str(exc)
locked.status = AITask.Status.FAILED
locked.error_code = "user_credit_insufficient" if "余额不足" in message else "invalid_input"
locked.error_message = message[:2000]
locked.completed_at = timezone.now()
locked.save(update_fields=["status", "error_code", "error_message", "completed_at", "updated_at"])
return locked
next_payload = dict(locked.request_payload or {})
next_payload["api_prompt"] = built["api_prompt"]
next_payload["references"] = built["snapshots"]
next_payload["estimated_tokens"] = tokens
next_payload["review_pending"] = False
locked.request_payload = next_payload
locked.estimated_cost = quote.points
locked.status = AITask.Status.RESERVED
locked.save(update_fields=["request_payload", "estimated_cost", "status", "updated_at"])
return _dispatch_free_video_provider(
task=locked,
built=built,
model_config=locked.model_config,
aspect_ratio=aspect_ratio,
duration=duration,
resolution=resolution,
generate_audio=generate_audio,
seed=seed,
search_mode=search_mode,
feature=feature,
mode=mode,
poll_countdown=30,
)
def _fail_pending_free_video(task: AITask, message: str) -> AITask:
if task.status not in (AITask.Status.CREATED, AITask.Status.RESERVED):
return task
task.status = AITask.Status.FAILED
task.error_code = "content_rejected"
task.error_message = message[:2000]
task.completed_at = timezone.now()
task.save(update_fields=["status", "error_code", "error_message", "completed_at", "updated_at"])
_notify_failure(task, raw=message, hint=message)
return task
def _dispatch_free_video_provider(
*,
task,
built,
model_config,
aspect_ratio,
duration,
resolution,
generate_audio,
seed,
search_mode,
feature,
mode,
poll_countdown=30,
):
"""RESERVED 任务调火山创建。失败退费。"""
try: try:
from .services import execute_routed_video_submit from .services import execute_routed_video_submit
@@ -558,8 +691,6 @@ def submit_free_video(*, team, user, params: dict) -> AITask:
request_summary={"feature": feature, "mode": mode}, request_summary={"feature": feature, "mode": mode},
) )
response, provider_task_id = routed.value response, provider_task_id = routed.value
# execute_model_call 以原子 F 表达式累计实际尝试的平台成本;刷新内存对象,
# 保证本方法返回值与数据库中的 AITask.base_cost 完全一致。
task.refresh_from_db(fields=["base_cost"]) task.refresh_from_db(fields=["base_cost"])
task.provider_task_id = provider_task_id task.provider_task_id = provider_task_id
task.response_payload = response task.response_payload = response
@@ -596,7 +727,7 @@ def submit_free_video(*, team, user, params: dict) -> AITask:
from apps.assets.free_asset_state import mark_remote_asset_unavailable from apps.assets.free_asset_state import mark_remote_asset_unavailable
mark_remote_asset_unavailable( mark_remote_asset_unavailable(
team=team, team=task.team,
local_asset_id=target["local_asset_id"], local_asset_id=target["local_asset_id"],
remote_asset_id=target["remote_asset_id"], remote_asset_id=target["remote_asset_id"],
) )
@@ -616,11 +747,10 @@ def submit_free_video(*, team, user, params: dict) -> AITask:
logger.warning("free video create failed: %s", exc) logger.warning("free video create failed: %s", exc)
return task return task
# worker 兜底轮询(自重排);派发失败仅 log,前端主动 poll 仍能收尾
try: try:
from .tasks import poll_free_video_task from .tasks import poll_free_video_task
poll_free_video_task.apply_async(args=[str(task.id), 0], countdown=30) poll_free_video_task.apply_async(args=[str(task.id), 0], countdown=poll_countdown)
except Exception: # noqa: BLE001 except Exception: # noqa: BLE001
logger.error("poll_free_video_task enqueue failed; relying on client polling", exc_info=True) logger.error("poll_free_video_task enqueue failed; relying on client polling", exc_info=True)
+40 -54
View File
@@ -52,6 +52,11 @@ from apps.projects.models import (
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
# 火山官方直连(SeeDream 生图 / Seedance 视频 / 豆包文本)走 ARK SDK;其余 provider 一律
# 视为「OpenAI 兼容中转站」走通用适配器。加/换中转站 = DB 加一行 ModelProvider,零改代码。
# 注意:DB 里火山 provider 实际命名为 "volcengine"(豆包),必须包含,否则会被错路由到中转站。
OFFICIAL_DIRECT_PROVIDERS = {"volcengine", "volcano", "ark", "volcano_ark", "doubao"}
def get_default_model(capability: str) -> ModelConfig: def get_default_model(capability: str) -> ModelConfig:
qs = ( qs = (
@@ -62,21 +67,24 @@ def get_default_model(capability: str) -> ModelConfig:
return qs.filter(is_default=True).order_by("created_at").first() or qs.order_by("created_at").first() return qs.filter(is_default=True).order_by("created_at").first() or qs.order_by("created_at").first()
def get_storyboard_image_model() -> ModelConfig: def get_storyboard_image_model() -> ModelConfig | None:
"""故事板出图钉 YunQi gpt-image-2 多图 edits(与手工测通的 curl 同一条链路)。 """故事板出图只走 GPT 图像模型(gpt-image / gpt-image-2)。
找不到再回落默认图像模型,避免测试/未 seed 环境直接挂。"""
pinned = ( 不回落默认图像模型,也不走火山 Seedream:用户明确要求故事板无论怎样都不改成其他生图模型。
"""
qs = (
ModelConfig.objects.select_related("provider") ModelConfig.objects.select_related("provider")
.filter( .filter(
capability=ModelConfig.Capability.IMAGE, capability=ModelConfig.Capability.IMAGE,
status=ModelConfig.Status.ACTIVE, status=ModelConfig.Status.ACTIVE,
provider__status="active", provider__status="active",
provider__name="yunqi", name__icontains="gpt-image",
name="gpt-image-2",
) )
.first()
) )
return pinned or get_default_model(ModelConfig.Capability.IMAGE) return (
qs.filter(name="gpt-image-2").order_by("created_at").first()
or qs.order_by("created_at").first()
)
def resolve_image_model(key: str | None) -> "ModelConfig | None": def resolve_image_model(key: str | None) -> "ModelConfig | None":
@@ -101,12 +109,6 @@ def resolve_image_model(key: str | None) -> "ModelConfig | None":
return qs.filter(name=key).first() return qs.filter(name=key).first()
# 火山官方直连(SeeDream 生图 / Seedance 视频 / 豆包文本)走 ARK SDK;其余 provider 一律
# 视为「OpenAI 兼容中转站」走通用适配器。加/换中转站 = DB 加一行 ModelProvider,零改代码。
# 注意:DB 里火山 provider 实际命名为 "volcengine"(豆包),必须包含,否则会被错路由到中转站。
OFFICIAL_DIRECT_PROVIDERS = {"volcengine", "volcano", "ark", "volcano_ark", "doubao"}
def public_model_name(model_config: ModelConfig) -> str: def public_model_name(model_config: ModelConfig) -> str:
"""普通用户公开名称保持稳定;Fallback 的真实模型只在管理员尝试链中展示。""" """普通用户公开名称保持稳定;Fallback 的真实模型只在管理员尝试链中展示。"""
@@ -2781,7 +2783,7 @@ def submit_storyboard(*, project, user, prompt: str = "", shot_ids: list | None
if adopted_script is None: if adopted_script is None:
raise ValueError("script must be adopted before generating storyboard") raise ValueError("script must be adopted before generating storyboard")
if get_storyboard_image_model() is None: if get_storyboard_image_model() is None:
raise ValueError("no active image model configured") raise ValueError("故事板只使用 GPT 图像模型,当前没有启用 gpt-image-2")
if prompt: if prompt:
meta = dict(project.metadata or {}) meta = dict(project.metadata or {})
if meta.get("storyboard_prompt") != prompt: if meta.get("storyboard_prompt") != prompt:
@@ -3071,7 +3073,6 @@ def _storyboard_shot_worker(task_id, shot_id, user_id) -> None:
user = User.objects.get(id=user_id) user = User.objects.get(id=user_id)
project = shot.project project = shot.project
segment = shot.script_segment segment = shot.script_segment
model_config = task.model_config
reservation = task.credit_reservation reservation = task.credit_reservation
extra_prompt = (project.metadata or {}).get("storyboard_prompt", "") or "" extra_prompt = (project.metadata or {}).get("storyboard_prompt", "") or ""
spec = project_output_spec(project) spec = project_output_spec(project)
@@ -3080,11 +3081,14 @@ def _storyboard_shot_worker(task_id, shot_id, user_id) -> None:
task.status = AITask.Status.SUBMITTED task.status = AITask.Status.SUBMITTED
task.save(update_fields=["status", "updated_at"]) task.save(update_fields=["status", "updated_at"])
try: try:
use_model_routing = bool((task.request_payload or {}).get("model_routing_v1")) # 故事板无论任务上挂了什么模型、是否开了路由,都只走 GPT 图像;失败也不换 Seedream。
provider = None if use_model_routing else get_image_provider(model_config) model_config = get_storyboard_image_model()
if model_config is None:
raise ValueError("故事板只使用 GPT 图像模型,当前没有启用 gpt-image-2")
provider = get_image_provider(model_config)
refs = _storyboard_reference_images(project, segment) if segment is not None else [] refs = _storyboard_reference_images(project, segment) if segment is not None else []
ref_urls = [r["url"] for r in refs] ref_urls = [r["url"] for r in refs]
if ref_urls and (use_model_routing or hasattr(provider, "image_edit")): if ref_urls and hasattr(provider, "image_edit"):
# gpt-image-2 多图参考:必须用 refs 版提示词(点名「参考图N=角色/场景/商品」+锁脸锁商品) # gpt-image-2 多图参考:必须用 refs 版提示词(点名「参考图N=角色/场景/商品」+锁脸锁商品)
frame_prompt = build_storyboard_frame_prompt_refs(project, segment, refs, extra_prompt) frame_prompt = build_storyboard_frame_prompt_refs(project, segment, refs, extra_prompt)
else: else:
@@ -3093,40 +3097,24 @@ def _storyboard_shot_worker(task_id, shot_id, user_id) -> None:
else (task.request_payload.get("prompt") or "") else (task.request_payload.get("prompt") or "")
) )
if use_model_routing: if ref_urls and hasattr(provider, "image_edit"):
routed = execute_routed_image_request( response = _call_image_with_retry(
task=task, lambda: provider.image_edit(
primary_model=model_config, model=model_config.name,
prompt=frame_prompt, prompt=frame_prompt,
reference_images=ref_urls, images=ref_urls,
aspect_ratio=frame_ratio, size=frame_size,
edit_size=frame_size, )
direct_size=frame_size,
request_summary={
"storyboard_shot": str(shot.id),
"storyboard_sort_order": shot.sort_order,
},
) )
response, media = routed.value
else: else:
if ref_urls and hasattr(provider, "image_edit"): response = _call_image_with_retry(
response = _call_image_with_retry( lambda: provider.image_generation(
lambda: provider.image_edit( model=model_config.name,
model=model_config.name, endpoint=model_config.endpoint,
prompt=frame_prompt, prompt=frame_prompt,
images=ref_urls,
size=frame_size,
)
) )
else: )
response = _call_image_with_retry( media = provider.extract_first_media_url(response)
lambda: provider.image_generation(
model=model_config.name,
endpoint=model_config.endpoint,
prompt=frame_prompt,
)
)
media = provider.extract_first_media_url(response)
asset = _store_generated_media( asset = _store_generated_media(
team=project.team, user=user, project=project, task=task, media=media, team=project.team, user=user, project=project, task=task, media=media,
name=f"{project.name}-storyboard-{shot.sort_order + 1}", name=f"{project.name}-storyboard-{shot.sort_order + 1}",
@@ -3206,6 +3194,8 @@ def poll_storyboard(*, project, user) -> dict:
if v if v
} }
model_config = get_storyboard_image_model() model_config = get_storyboard_image_model()
if model_config is None:
return {"status": "failed", "done": done, "total": total, "error": "故事板只使用 GPT 图像模型,当前没有启用 gpt-image-2"}
extra_prompt = (project.metadata or {}).get("storyboard_prompt", "") or "" extra_prompt = (project.metadata or {}).get("storyboard_prompt", "") or ""
spawnable = [s for s in active if str(s.id) not in inflight_shot_ids] spawnable = [s for s in active if str(s.id) not in inflight_shot_ids]
slots = max(0, STORYBOARD_MAX_PARALLEL - len(inflight_shot_ids)) slots = max(0, STORYBOARD_MAX_PARALLEL - len(inflight_shot_ids))
@@ -3217,12 +3207,8 @@ def poll_storyboard(*, project, user) -> dict:
"model": model_config.name, "endpoint": model_config.endpoint, "model": model_config.name, "endpoint": model_config.endpoint,
"prompt": build_storyboard_frame_prompt(project, segment, extra_prompt) if segment is not None else "", "prompt": build_storyboard_frame_prompt(project, segment, extra_prompt) if segment is not None else "",
"storyboard_shot": str(shot.id), "storyboard_shot": str(shot.id),
"model_routing_v1": True,
}, },
) )
# 真实平台成本由每条 AIModelAttempt 按实际模型累加,避免 Fallback 后仍记默认模型旧成本。
task.base_cost = Decimal("0")
task.save(update_fields=["base_cost", "updated_at"])
StoryboardShot.objects.filter(id=shot.id).update(status=StoryboardShot.Status.RUNNING, updated_at=timezone.now()) StoryboardShot.objects.filter(id=shot.id).update(status=StoryboardShot.Status.RUNNING, updated_at=timezone.now())
threading.Thread( threading.Thread(
target=_storyboard_shot_worker, args=(str(task.id), str(shot.id), str(user.id)), daemon=True target=_storyboard_shot_worker, args=(str(task.id), str(shot.id), str(user.id)), daemon=True
+8 -2
View File
@@ -78,12 +78,16 @@ def poll_free_video_task(self, task_id: str, attempt: int = 0) -> str:
与前端主动 poll 并存不双扣。轮询本身出错不重试(max_retries=0),下一次自重排继续。""" 与前端主动 poll 并存不双扣。轮询本身出错不重试(max_retries=0),下一次自重排继续。"""
from apps.ai.free_video import finalize_free_video from apps.ai.free_video import finalize_free_video
from apps.ai.models import AITask from apps.ai.models import AITask
from apps.ai.video_replace import advance_video_replace, is_video_replace_task
task = AITask.objects.select_related("model_config", "model_config__provider", "team").filter(id=task_id).first() task = AITask.objects.select_related("model_config", "model_config__provider", "team").filter(id=task_id).first()
if task is None: if task is None:
return task_id return task_id
try: try:
task = finalize_free_video(task=task) if is_video_replace_task(task) and task.status == AITask.Status.CREATED:
task = advance_video_replace(task)
else:
task = finalize_free_video(task=task)
except Exception: # noqa: BLE001 — 单次轮询失败(网络抖动等)不终结任务,等下一轮 except Exception: # noqa: BLE001 — 单次轮询失败(网络抖动等)不终结任务,等下一轮
import logging import logging
@@ -93,7 +97,9 @@ def poll_free_video_task(self, task_id: str, attempt: int = 0) -> str:
if getattr(dj_settings, "CELERY_TASK_ALWAYS_EAGER", False): if getattr(dj_settings, "CELERY_TASK_ALWAYS_EAGER", False):
return task_id return task_id
if task.status in (AITask.Status.SUBMITTED, AITask.Status.POLLING) and attempt < 60: if task.status == AITask.Status.CREATED and attempt < 60:
poll_free_video_task.apply_async(args=[task_id, attempt + 1], countdown=8 if attempt < 12 else 30)
elif task.status in (AITask.Status.SUBMITTED, AITask.Status.POLLING) and attempt < 60:
poll_free_video_task.apply_async(args=[task_id, attempt + 1], countdown=30) poll_free_video_task.apply_async(args=[task_id, attempt + 1], countdown=30)
return task_id return task_id
+45 -55
View File
@@ -33,7 +33,10 @@ def _metadata(*, outbound=True, base_cost="0.50", max_refs=9):
} }
@override_settings(STORYBOARD_MAX_PARALLEL=4) @override_settings(
STORYBOARD_MAX_PARALLEL=4,
CACHES={"default": {"BACKEND": "django.core.cache.backends.locmem.LocMemCache"}},
)
class StoryboardRoutingTests(TestCase): class StoryboardRoutingTests(TestCase):
def setUp(self): def setUp(self):
ModelConfig.objects.filter(capability=ModelConfig.Capability.IMAGE).update( ModelConfig.objects.filter(capability=ModelConfig.Capability.IMAGE).update(
@@ -158,20 +161,15 @@ class StoryboardRoutingTests(TestCase):
return CreditLedger.objects.filter(task=task, ledger_type=ledger_type).count() return CreditLedger.objects.filter(task=task, ledger_type=ledger_type).count()
def test_multi_reference_success_records_9_16_and_adopts_one_shot_version(self): def test_multi_reference_success_records_9_16_and_adopts_one_shot_version(self):
primary = self.model(self.provider("storyboard-primary", 20), "storyboard-primary") primary = self.model(self.provider("storyboard-primary", 20), "gpt-image-2")
task = self.enqueue() task = self.enqueue()
task.refresh_from_db() task.refresh_from_db()
self.shot.refresh_from_db() self.shot.refresh_from_db()
self.assertEqual(task.status, AITask.Status.SUCCEEDED) self.assertEqual(task.status, AITask.Status.SUCCEEDED)
self.assertTrue(task.request_payload["model_routing_v1"]) self.assertNotIn("model_routing_v1", task.request_payload)
self.assertEqual(task.base_cost, Decimal("0.5000")) self.assertEqual(task.model_config_id, primary.id)
attempt = task.model_attempts.get() self.assertEqual(task.model_config.name, "gpt-image-2")
self.assertEqual(attempt.model_config_id, primary.id)
self.assertEqual(attempt.operation, "image_edit")
self.assertEqual(attempt.request_summary["reference_images"], 2)
self.assertEqual(attempt.request_summary["aspect_ratio"], "9:16")
self.assertEqual(attempt.request_summary["storyboard_shot"], str(self.shot.id))
call = self.provider_mocks[primary.id].image_edit.call_args call = self.provider_mocks[primary.id].image_edit.call_args
self.assertEqual( self.assertEqual(
call.kwargs["images"], call.kwargs["images"],
@@ -179,6 +177,7 @@ class StoryboardRoutingTests(TestCase):
) )
self.assertEqual(call.kwargs["prompt"], "带参考图编号与锁定约束的故事板提示词") self.assertEqual(call.kwargs["prompt"], "带参考图编号与锁定约束的故事板提示词")
self.assertEqual(call.kwargs["size"], "1024x1536") self.assertEqual(call.kwargs["size"], "1024x1536")
self.assertEqual(call.kwargs["model"], "gpt-image-2")
self.assertIsNotNone(self.shot.adopted_version_id) self.assertIsNotNone(self.shot.adopted_version_id)
self.assertEqual(self.shot.adopted_version.task_id, task.id) self.assertEqual(self.shot.adopted_version.task_id, task.id)
self.assertEqual(self.shot.versions.count(), 1) self.assertEqual(self.shot.versions.count(), 1)
@@ -188,44 +187,40 @@ class StoryboardRoutingTests(TestCase):
def test_wizard_aspect_ratio_drives_storyboard_size(self): def test_wizard_aspect_ratio_drives_storyboard_size(self):
self.project.metadata = {"wizard": {"aspect_ratio": "16:9", "resolution": "480p"}} self.project.metadata = {"wizard": {"aspect_ratio": "16:9", "resolution": "480p"}}
self.project.save(update_fields=["metadata"]) self.project.save(update_fields=["metadata"])
primary = self.model(self.provider("storyboard-wide", 20), "storyboard-wide") primary = self.model(self.provider("storyboard-wide", 20), "gpt-image-2")
task = self.enqueue() 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 call = self.provider_mocks[primary.id].image_edit.call_args
self.assertEqual(call.kwargs["size"], "1536x864") self.assertEqual(call.kwargs["size"], "1536x864")
def test_retry_then_dynamic_candidate_success_charges_once(self): def test_gpt_failure_does_not_fallback_to_seedream(self):
primary = self.model( primary = self.model(
self.provider("storyboard-fallback-primary", 100), self.provider("storyboard-gpt", 100),
"storyboard-primary", "gpt-image-2",
base_cost="0.25", base_cost="0.25",
) )
candidate = self.model( seedream = self.model(
self.provider("volcano", 10), self.provider("storyboard-volcano-fallback", 10),
"storyboard-candidate", "seedream-4-5-251128",
outbound=False, outbound=False,
base_cost="0.75", base_cost="0.75",
) )
self.provider_mocks[primary.id] = self._new_provider_mock() self.provider_mocks[primary.id] = self._new_provider_mock()
self.provider_mocks[primary.id].image_edit.side_effect = requests.ConnectionError("offline") self.provider_mocks[primary.id].image_edit.side_effect = requests.ConnectionError("offline")
self.provider_mocks[seedream.id] = self._new_provider_mock()
task = self.enqueue() task = self.enqueue()
task.refresh_from_db() task.refresh_from_db()
attempts = list(task.model_attempts.all()) self.assertEqual(task.status, AITask.Status.FAILED)
self.assertEqual(task.status, AITask.Status.SUCCEEDED) self.assertEqual(task.model_config_id, primary.id)
self.assertEqual([a.model_config_id for a in attempts], [primary.id, primary.id, candidate.id]) self.provider_mocks[seedream.id].image_edit.assert_not_called()
self.assertTrue(attempts[-1].is_fallback) self.provider_mocks[seedream.id].image_generation.assert_not_called()
self.assertEqual(task.base_cost, Decimal("0.7500"))
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1) self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1) self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 0)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 0) self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 1)
def test_all_candidates_fail_releases_and_keeps_old_adopted_version(self): def test_all_candidates_fail_releases_and_keeps_old_adopted_version(self):
primary = self.model(self.provider("storyboard-all-primary", 100), "storyboard-primary") primary = self.model(self.provider("storyboard-all-primary", 100), "gpt-image-2")
candidate = self.model( self.model(self.provider("storyboard-volcano-keep", 10), "seedream-4-5-251128", outbound=False)
self.provider("volcano", 10), "storyboard-candidate", outbound=False
)
old_asset = Asset.objects.create( old_asset = Asset.objects.create(
team=self.team, team=self.team,
created_by=self.user, created_by=self.user,
@@ -242,16 +237,14 @@ class StoryboardRoutingTests(TestCase):
) )
self.shot.adopted_version = old_version self.shot.adopted_version = old_version
self.shot.save(update_fields=["adopted_version", "updated_at"]) self.shot.save(update_fields=["adopted_version", "updated_at"])
for model in (primary, candidate): self.provider_mocks[primary.id] = self._new_provider_mock()
self.provider_mocks[model.id] = self._new_provider_mock() self.provider_mocks[primary.id].image_edit.side_effect = requests.ConnectionError("offline")
self.provider_mocks[model.id].image_edit.side_effect = requests.ConnectionError("offline")
task = self.enqueue() task = self.enqueue()
task.refresh_from_db() task.refresh_from_db()
self.shot.refresh_from_db() self.shot.refresh_from_db()
self.assertEqual(task.status, AITask.Status.FAILED) self.assertEqual(task.status, AITask.Status.FAILED)
self.assertEqual(task.model_attempts.count(), 3)
self.assertEqual(self.shot.adopted_version_id, old_version.id) self.assertEqual(self.shot.adopted_version_id, old_version.id)
self.assertEqual(self.shot.versions.count(), 1) self.assertEqual(self.shot.versions.count(), 1)
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1) self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
@@ -259,38 +252,35 @@ class StoryboardRoutingTests(TestCase):
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 1) self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 1)
self.assertEqual(CreditAccount.objects.get(team=self.team).reserved_balance, Decimal("0")) self.assertEqual(CreditAccount.objects.get(team=self.team).reserved_balance, Decimal("0"))
def test_no_reference_uses_image_generation_and_keeps_vertical_capability(self): def test_no_reference_uses_image_generation(self):
primary = self.model(self.provider("storyboard-no-ref", 20), "storyboard-no-ref") primary = self.model(self.provider("storyboard-no-ref", 20), "gpt-image-2")
self.refs = [] self.refs = []
task = self.enqueue() task = self.enqueue()
task.refresh_from_db() task.refresh_from_db()
attempt = task.model_attempts.get() self.provider_mocks[primary.id].image_edit.assert_not_called()
self.assertEqual(attempt.operation, "image_generate")
self.assertEqual(attempt.request_summary["reference_images"], 0)
self.assertEqual(attempt.request_summary["aspect_ratio"], "9:16")
call = self.provider_mocks[primary.id].image_generation.call_args call = self.provider_mocks[primary.id].image_generation.call_args
self.assertEqual(call.kwargs["prompt"], "无参考图故事板提示词") self.assertEqual(call.kwargs["prompt"], "无参考图故事板提示词")
self.assertNotIn("size", call.kwargs) self.assertEqual(call.kwargs["model"], "gpt-image-2")
def test_direct_primary_uses_image_generation_with_all_references(self): def test_seedream_is_ignored_even_if_it_is_the_default_image_model(self):
primary = self.model( seedream = self.model(
self.provider("volcano", 10), "seedream-storyboard", outbound=False self.provider("storyboard-volcano-default", 10), "seedream-4-5-251128", outbound=False
) )
self.provider_mocks[primary.id] = self._new_generation_only_provider_mock() gpt = self.model(self.provider("storyboard-gpt-default", 20), "gpt-image-2")
seedream.is_default = True
seedream.save(update_fields=["is_default"])
self.provider_mocks[seedream.id] = self._new_generation_only_provider_mock()
task = self.enqueue() task = self.enqueue()
task.refresh_from_db() task.refresh_from_db()
call = self.provider_mocks[primary.id].image_generation.call_args self.assertEqual(task.status, AITask.Status.SUCCEEDED)
self.assertEqual( self.assertEqual(task.model_config_id, gpt.id)
call.kwargs["image"], self.provider_mocks[seedream.id].image_generation.assert_not_called()
["http://example.test/person.png", "http://example.test/product.png"], self.provider_mocks[gpt.id].image_edit.assert_called_once()
)
self.assertEqual(call.kwargs["size"], "1024x1536")
self.assertEqual(task.model_attempts.get().public_model_name, primary.display_name)
def test_poll_creates_one_independent_task_and_reservation_per_shot(self): def test_poll_creates_one_independent_task_and_reservation_per_shot(self):
self.model(self.provider("storyboard-batch", 20), "storyboard-batch") gpt = self.model(self.provider("storyboard-batch", 20), "gpt-image-2")
second_segment = ScriptSegment.objects.create( second_segment = ScriptSegment.objects.create(
script_version=self.segment.script_version, script_version=self.segment.script_version,
sort_order=1, sort_order=1,
@@ -317,6 +307,6 @@ class StoryboardRoutingTests(TestCase):
{task.request_payload["storyboard_shot"] for task in tasks}, {task.request_payload["storyboard_shot"] for task in tasks},
{str(self.shot.id), str(second_shot.id)}, {str(self.shot.id), str(second_shot.id)},
) )
self.assertTrue(all(task.request_payload["model_routing_v1"] for task in tasks)) self.assertTrue(all("model_routing_v1" not in task.request_payload for task in tasks))
self.assertTrue(all(task.base_cost == Decimal("0") for task in tasks)) self.assertTrue(all(task.model_config_id == gpt.id for task in tasks))
self.assertTrue(all(self.ledger_count(task, CreditLedger.Type.RESERVE) == 1 for task in tasks)) self.assertTrue(all(self.ledger_count(task, CreditLedger.Type.RESERVE) == 1 for task in tasks))
+116 -3
View File
@@ -3,6 +3,7 @@
运行:DB_ENGINE=sqlite python manage.py test apps.ai.test_video_replace --settings=airshelf.settings.test 运行:DB_ENGINE=sqlite python manage.py test apps.ai.test_video_replace --settings=airshelf.settings.test
""" """
from unittest.mock import MagicMock, patch from unittest.mock import MagicMock, patch
from uuid import uuid4
from django.test import TestCase from django.test import TestCase
from rest_framework.test import APIClient from rest_framework.test import APIClient
@@ -11,13 +12,13 @@ from apps.accounts.models import Team, TeamMember, User
from apps.ai.free_video import submit_free_video from apps.ai.free_video import submit_free_video
from apps.ai.models import AITask, ModelConfig from apps.ai.models import AITask, ModelConfig
from apps.ai.test_free_video import STANDARD, _ark_create_response from apps.ai.test_free_video import STANDARD, _ark_create_response
from apps.ai.video_replace import PRODUCT_PROMPT, submit_video_replace from apps.ai.video_replace import PRODUCT_PROMPT, REVIEW_FAILED, REVIEW_UNAVAILABLE, advance_video_replace, submit_video_replace
from apps.assets.models import Asset, AssetFile, Model from apps.assets.models import Asset, AssetFile, Model
from apps.billing.models import CreditAccount from apps.billing.models import CreditAccount, CreditReservation
from apps.products.models import Product, ProductImage from apps.products.models import Product, ProductImage
def _asset(team, user, *, kind=Asset.Type.IMAGE, name="素材", duration_ms=None, preview="http://tos/1.png"): def _asset(team, user, *, kind=Asset.Type.IMAGE, name="素材", duration_ms=None, preview="http://tos/1.png", review_status="active", review_remote_id=None):
asset = Asset.objects.create( asset = Asset.objects.create(
team=team, team=team,
created_by=user, created_by=user,
@@ -25,6 +26,8 @@ def _asset(team, user, *, kind=Asset.Type.IMAGE, name="素材", duration_ms=None
asset_type=kind, asset_type=kind,
source=Asset.Source.AI_GENERATED, source=Asset.Source.AI_GENERATED,
category=Asset.Category.UPLOAD, category=Asset.Category.UPLOAD,
review_status=review_status,
review_remote_id=review_remote_id if review_remote_id is not None else (f"asset-{uuid4().hex[:12]}" if review_status == "active" else ""),
) )
AssetFile.objects.create( AssetFile.objects.create(
asset=asset, asset=asset,
@@ -84,6 +87,14 @@ class SubmitVideoReplaceTests(TestCase):
roles = [item.get("role") for item in content] roles = [item.get("role") for item in content]
self.assertIn("reference_video", roles) self.assertIn("reference_video", roles)
self.assertIn("reference_image", roles) self.assertIn("reference_image", roles)
urls = []
for item in content:
if item.get("type") == "video_url":
urls.append(item["video_url"]["url"])
elif item.get("type") == "image_url":
urls.append(item["image_url"]["url"])
self.assertTrue(urls)
self.assertTrue(all(url.startswith("asset://") for url in urls))
def test_character_library_and_temp_are_exclusive(self): def test_character_library_and_temp_are_exclusive(self):
portrait = _asset(self.team, self.user, name="模特.png", preview="http://tos/model.png") portrait = _asset(self.team, self.user, name="模特.png", preview="http://tos/model.png")
@@ -234,3 +245,105 @@ class VideoReplaceApiTests(TestCase):
def test_validation_error_returns_400(self): def test_validation_error_returns_400(self):
resp = self.client.post("/api/ai/video-replace/", {"replace_mode": "product"}, format="json") resp = self.client.post("/api/ai/video-replace/", {"replace_mode": "product"}, format="json")
self.assertEqual(resp.status_code, 400) self.assertEqual(resp.status_code, 400)
class VideoReplaceReviewGateTests(TestCase):
def setUp(self):
self.user = User.objects.create_user(username="vrrev", password="p")
self.team = Team.objects.create(name="VRR", owner=self.user)
CreditAccount.objects.create(team=self.team, balance="10000.0000")
self.provider = MagicMock()
self.provider.create_video_task.return_value = _ark_create_response()
patch("apps.ai.services.build_provider", return_value=self.provider).start()
patch("apps.ai.tasks.poll_free_video_task.apply_async").start()
self.addCleanup(patch.stopall)
self.video = _asset(
self.team, self.user, kind=Asset.Type.VIDEO, name="真人.mp4", duration_ms=8000,
preview="http://tos/video.mp4", review_status="", review_remote_id="",
)
self.image = _asset(
self.team, self.user, name="乔尔.png", preview="http://tos/face.png",
review_status="", review_remote_id="",
)
self.product = Product.objects.create(team=self.team, created_by=self.user, title="净颜精华", cover_asset=self.image)
ProductImage.objects.create(product=self.product, asset=self.image, is_primary=True)
def _submit(self, **over):
params = {
"replace_mode": "product",
"video_asset_id": str(self.video.id),
"product_id": str(self.product.id),
"model": STANDARD,
"aspect_ratio": "9:16",
"resolution": "480p",
"duration": 4,
}
params.update(over)
return submit_video_replace(team=self.team, user=self.user, params=params)
def test_review_disabled_without_remote_id_is_rejected(self):
with patch("apps.assets.assets_client.is_enabled", return_value=False):
with self.assertRaisesMessage(ValueError, REVIEW_UNAVAILABLE):
self._submit()
self.assertFalse(self.provider.create_video_task.called)
self.assertEqual(CreditReservation.objects.filter(team=self.team).count(), 0)
def test_unreviewed_assets_create_pending_task_without_charging(self):
with patch("apps.assets.assets_client.is_enabled", return_value=True), patch(
"apps.assets.review.get_or_create_team_group"
) as grp, patch("apps.assets.assets_client.create_asset", side_effect=["Asset-vid", "Asset-img"]) as create, patch(
"apps.assets.assets_client.get_asset", return_value={"Status": "Processing"}
):
grp.return_value = MagicMock(remote_group_id="Group-1")
task = self._submit()
self.assertEqual(task.status, AITask.Status.CREATED)
self.assertTrue((task.request_payload or {}).get("review_pending"))
self.assertEqual(CreditReservation.objects.filter(task=task).count(), 0)
self.assertFalse(self.provider.create_video_task.called)
passed_types = [call.kwargs.get("asset_type") for call in create.call_args_list]
self.assertIn("Video", passed_types)
self.assertIn("Image", passed_types)
def test_advance_starts_generation_after_review_passes(self):
with patch("apps.assets.assets_client.is_enabled", return_value=True), patch(
"apps.assets.review.get_or_create_team_group"
) as grp, patch("apps.assets.assets_client.create_asset", side_effect=["Asset-vid", "Asset-img"]), patch(
"apps.assets.assets_client.get_asset", return_value={"Status": "Processing"}
):
grp.return_value = MagicMock(remote_group_id="Group-1")
task = self._submit()
self.video.refresh_from_db()
self.image.refresh_from_db()
self.video.review_status = "active"
self.video.review_remote_id = "asset-vid"
self.video.save(update_fields=["review_status", "review_remote_id"])
self.image.review_status = "active"
self.image.review_remote_id = "asset-img"
self.image.save(update_fields=["review_status", "review_remote_id"])
with patch("apps.assets.assets_client.is_enabled", return_value=True), patch(
"apps.assets.assets_client.get_asset", return_value={"Status": "Active"}
):
task = advance_video_replace(task)
self.assertEqual(task.status, AITask.Status.SUBMITTED)
self.assertTrue(CreditReservation.objects.filter(task=task).exists())
self.assertTrue(self.provider.create_video_task.called)
content = self.provider.create_video_task.call_args.kwargs.get("content_items") or []
urls = [item.get("video_url", item.get("image_url", {})).get("url") for item in content]
self.assertTrue(all(str(url).startswith("asset://") for url in urls if url))
def test_advance_fails_review_without_charging(self):
with patch("apps.assets.assets_client.is_enabled", return_value=True), patch(
"apps.assets.review.get_or_create_team_group"
) as grp, patch("apps.assets.assets_client.create_asset", side_effect=["Asset-vid", "Asset-img"]), patch(
"apps.assets.assets_client.get_asset", return_value={"Status": "Processing"}
):
grp.return_value = MagicMock(remote_group_id="Group-1")
task = self._submit()
with patch("apps.assets.assets_client.is_enabled", return_value=True), patch(
"apps.assets.assets_client.get_asset", return_value={"Status": "Failed", "ErrorMessage": "real person"}
):
task = advance_video_replace(task)
self.assertEqual(task.status, AITask.Status.FAILED)
self.assertIn("合规审核", task.error_message or REVIEW_FAILED)
self.assertEqual(CreditReservation.objects.filter(task=task).count(), 0)
self.assertFalse(self.provider.create_video_task.called)
+322 -26
View File
@@ -2,23 +2,47 @@
不新建任务类型、不接检测/抠图。提交仍走 submit_free_video,只把 不新建任务类型、不接检测/抠图。提交仍走 submit_free_video,只把
feature=video_replace 和 replace_mode 写进 payload,提示词由后端写死。 feature=video_replace 和 replace_mode 写进 payload,提示词由后端写死。
真人参考必须先送火山素材库审核,过审后用 asset:// 生成;审核失败不扣费。
""" """
from __future__ import annotations from __future__ import annotations
import logging
import uuid import uuid
from decimal import Decimal
from django.conf import settings
from django.db import transaction
from django.db.models import Q from django.db.models import Q
from apps.assets.models import Asset, Model from apps.assets.models import Asset, Model
from apps.billing.models import CreditAccount
from apps.billing.pricing import quote_video_estimate, video_reserve_amount
from apps.products.models import Product from apps.products.models import Product
from .free_video import HIGH_RES_MODEL, serialize_free_video_task, submit_free_video from .free_video import (
FREE_VIDEO_MODELS,
HIGH_RES_MODEL,
IN_FLIGHT_STATUSES,
RATIOS,
RESOLUTIONS,
_reap_stale_free_video_tasks,
serialize_free_video_task,
start_pending_free_video,
submit_free_video,
)
from .media_probe import REF_DURATION_MAX from .media_probe import REF_DURATION_MAX
from .models import AITask, ModelConfig
logger = logging.getLogger(__name__)
FEATURE = "video_replace" FEATURE = "video_replace"
REPLACE_MODES = {"product", "character"} REPLACE_MODES = {"product", "character"}
MAX_IMAGES = 9 MAX_IMAGES = 9
LEGACY_PROMPT_PREFIX = "[视频复刻]" LEGACY_PROMPT_PREFIX = "[视频复刻]"
REVIEW_UNAVAILABLE = "素材审核服务暂不可用,请稍后重试"
REVIEW_FAILED = "参考素材未通过真人合规审核,请更换视频或图片后重试"
REVIEW_SUBMIT_FAILED = "素材提交审核失败,请稍后重试"
PRODUCT_PROMPT = ( PRODUCT_PROMPT = (
"使用@参考视频作为镜头、节奏与口播氛围基准," "使用@参考视频作为镜头、节奏与口播氛围基准,"
@@ -49,6 +73,7 @@ def serialize_video_replace_task(task, *, include_deleted_assets: bool = False)
data = serialize_free_video_task(task, include_deleted_assets=include_deleted_assets) data = serialize_free_video_task(task, include_deleted_assets=include_deleted_assets)
payload = task.request_payload or {} payload = task.request_payload or {}
replace_mode = payload.get("replace_mode") or _legacy_replace_mode(payload.get("prompt") or "") replace_mode = payload.get("replace_mode") or _legacy_replace_mode(payload.get("prompt") or "")
reviewing = task.status == AITask.Status.CREATED and bool(payload.get("review_pending"))
data.update({ data.update({
"feature": FEATURE, "feature": FEATURE,
"replace_mode": replace_mode, "replace_mode": replace_mode,
@@ -56,12 +81,13 @@ def serialize_video_replace_task(task, *, include_deleted_assets: bool = False)
"subject_source": payload.get("subject_source") or "", "subject_source": payload.get("subject_source") or "",
"product_id": payload.get("product_id") or "", "product_id": payload.get("product_id") or "",
"model_id": payload.get("model_id") or "", "model_id": payload.get("model_id") or "",
"review_stage": "reviewing" if reviewing else "",
}) })
return data return data
def submit_video_replace(*, team, user, params: dict): def submit_video_replace(*, team, user, params: dict):
"""校验素材 → 套提示词 → 复用 free_video 提交。失败抛 ValueError。""" """校验素材 → 送审 → 已过审则直接生成,否则建 CREATED 任务等绿盾。失败抛 ValueError。"""
replace_mode = str(params.get("replace_mode") or "").strip() replace_mode = str(params.get("replace_mode") or "").strip()
if replace_mode not in REPLACE_MODES: if replace_mode not in REPLACE_MODES:
raise ValueError("请选择替换商品或替换角色") raise ValueError("请选择替换商品或替换角色")
@@ -107,35 +133,305 @@ def submit_video_replace(*, team, user, params: dict):
_owned_ref(video, kind="video", role="reference_video", label="参考视频"), _owned_ref(video, kind="video", role="reference_video", label="参考视频"),
*image_refs, *image_refs,
] ]
return submit_free_video( review_state = _ensure_replace_refs_reviewed(team, references)
team=team, if review_state == "failed":
user=user, raise ValueError(REVIEW_FAILED)
params={ references = _refresh_replace_refs(team, references)
"prompt": prompt, extra = {
"mode": "universal", "replace_mode": replace_mode,
"model": str(params.get("model") or HIGH_RES_MODEL), "subject_name": subject_name,
"aspect_ratio": str(params.get("aspect_ratio") or "9:16"), "subject_source": subject_source,
"resolution": str(params.get("resolution") or "720p"), "product_id": str(product_id) if product_id else "",
"duration": duration, "model_id": str(model_id) if model_id else "",
"seed": params.get("seed", -1), "review_pending": review_state != "ready",
"generate_audio": True, }
"references": references, submit_params = {
"feature": FEATURE, "prompt": prompt,
"extra_payload": { "mode": "universal",
"replace_mode": replace_mode, "model": str(params.get("model") or HIGH_RES_MODEL),
"subject_name": subject_name, "aspect_ratio": str(params.get("aspect_ratio") or "9:16"),
"subject_source": subject_source, "resolution": str(params.get("resolution") or "720p"),
"product_id": str(product_id) if product_id else "", "duration": duration,
"model_id": str(model_id) if model_id else "", "seed": params.get("seed", -1),
}, "generate_audio": True,
}, "references": references,
) "feature": FEATURE,
"extra_payload": extra,
}
if review_state == "ready":
_assert_replace_refs_ready(references)
return submit_free_video(team=team, user=user, params=submit_params)
return _create_reviewing_task(team=team, user=user, params=submit_params)
def advance_video_replace(task):
"""轮询审核中的复刻任务:失败则结束(不扣费),过审则预留积分并提交火山。"""
if not is_video_replace_task(task):
return task
if task.status != AITask.Status.CREATED:
from .free_video import finalize_free_video
return finalize_free_video(task=task)
payload = task.request_payload or {}
references = list(payload.get("references") or [])
try:
state = _ensure_replace_refs_reviewed(task.team, references)
except ValueError as exc:
return _fail_reviewing_task(task, str(exc))
if state == "failed":
return _fail_reviewing_task(task, REVIEW_FAILED)
if state != "ready":
return task
refreshed = _refresh_replace_refs(task.team, references)
try:
_assert_replace_refs_ready(refreshed)
except ValueError as exc:
return _fail_reviewing_task(task, str(exc))
with transaction.atomic():
locked = AITask.objects.select_for_update().get(id=task.id)
if locked.status != AITask.Status.CREATED:
return locked
next_payload = dict(locked.request_payload or {})
next_payload["references"] = refreshed
next_payload["review_pending"] = False
locked.request_payload = next_payload
locked.save(update_fields=["request_payload", "updated_at"])
task = locked
return start_pending_free_video(task)
def _legacy_replace_mode(prompt: str) -> str: def _legacy_replace_mode(prompt: str) -> str:
return "character" if prompt.startswith("[视频复刻·角色]") else "product" return "character" if prompt.startswith("[视频复刻·角色]") else "product"
def _fail_reviewing_task(task, message: str):
from .free_video import _fail_pending_free_video
return _fail_pending_free_video(task, message)
def _enqueue_replace_review_poll(task):
try:
from .tasks import poll_free_video_task
poll_free_video_task.apply_async(args=[str(task.id), 0], countdown=8)
except Exception: # noqa: BLE001
logger.error("video replace review poll enqueue failed; relying on client polling", exc_info=True)
def _create_reviewing_task(*, team, user, params: dict):
"""审核未完成:只建 CREATED 任务,不预留积分。"""
model_name = str(params.get("model") or HIGH_RES_MODEL)
aspect_ratio = str(params.get("aspect_ratio") or "9:16")
resolution = str(params.get("resolution") or "720p")
try:
duration = int(params.get("duration") or 5)
except (TypeError, ValueError):
raise ValueError("时长参数无效")
if model_name not in FREE_VIDEO_MODELS:
raise ValueError("模型无效")
if aspect_ratio not in RATIOS:
raise ValueError("画面比例无效")
if resolution not in RESOLUTIONS:
raise ValueError("分辨率无效")
if not 4 <= duration <= 15:
raise ValueError("视频时长需在 4-15 秒之间")
model_config = (
ModelConfig.objects.select_related("provider")
.filter(name=model_name, capability=ModelConfig.Capability.VIDEO, status=ModelConfig.Status.ACTIVE)
.first()
)
if model_config is None:
raise ValueError("视频模型未配置,请联系管理员")
_reap_stale_free_video_tasks(team=team)
max_concurrent = int(getattr(settings, "FREE_VIDEO_MAX_CONCURRENT", 3))
in_flight = AITask.objects.filter(
team=team, task_type=AITask.Type.FREE_VIDEO, status__in=IN_FLIGHT_STATUSES
).count()
if in_flight >= max_concurrent:
raise ValueError(f"当前有 {in_flight} 个视频任务进行中(上限 {max_concurrent}),请等待完成后再提交")
references = params.get("references") or []
tokens, quote = quote_video_estimate(
model_config,
aspect_ratio=aspect_ratio,
resolution=resolution,
duration=duration,
references=references,
team=team,
)
reserve_amount = video_reserve_amount(quote.points)
account = CreditAccount.objects.filter(team=team).first()
available = (account.balance - account.reserved_balance) if account else Decimal("0")
if available < reserve_amount:
raise ValueError("团队余额不足,请充值后重试")
try:
seed = int(params.get("seed") if params.get("seed") is not None else -1)
except (TypeError, ValueError):
seed = -1
extra = params.get("extra_payload") if isinstance(params.get("extra_payload"), dict) else {}
request_payload = {
"feature": FEATURE,
"mode": "universal",
"model": model_name,
"endpoint": model_config.endpoint,
"prompt": params.get("prompt") or "",
"api_prompt": "",
"aspect_ratio": aspect_ratio,
"resolution": resolution,
"duration": duration,
"seed": seed,
"generate_audio": True,
"search_mode": "off",
"estimated_tokens": tokens,
"price_multiplier": quote.meta.get("price_multiplier", "1"),
"points_per_yuan_snapshot": quote.meta.get("rate", ""),
"references": references,
"model_routing_v1": True,
"review_pending": True,
}
for key, value in extra.items():
if key in request_payload or value in (None, ""):
continue
request_payload[key] = value
task = AITask.objects.create(
team=team,
created_by=user,
project=None,
task_type=AITask.Type.FREE_VIDEO,
status=AITask.Status.CREATED,
model_config=model_config,
idempotency_key=f"free_video:{team.id}:{uuid.uuid4()}",
request_payload=request_payload,
estimated_cost=quote.points,
base_cost=Decimal("0"),
)
_enqueue_replace_review_poll(task)
return task
def _replace_ref_assets(team, references: list) -> list[tuple[dict, Asset]]:
out = []
seen = set()
for ref in references or []:
raw_id = ref.get("asset_id")
if not raw_id:
continue
try:
parsed = uuid.UUID(str(raw_id))
except (TypeError, ValueError):
continue
if parsed in seen:
continue
seen.add(parsed)
asset = _load_replace_asset(team, parsed, ref.get("label") or "参考素材")
out.append((ref, asset))
return out
def _load_replace_asset(team, asset_id: uuid.UUID, label: str) -> Asset:
asset = Asset.objects.filter(id=asset_id, is_deleted=False, purged_at__isnull=True).first()
if asset is None:
raise ValueError(f"{label}不存在或已被删除")
if asset.team_id == team.id:
return asset
if Model.objects.filter(Q(is_official=True), Q(portrait_asset=asset) | Q(triview_asset=asset)).exists():
return asset
raise ValueError(f"{label}不存在或已被删除")
def _ensure_replace_refs_reviewed(team, references: list) -> str:
"""送审/轮询全部参考素材。返回 ready / pending / failed;审核未配置且无 remote_id 抛错。"""
from apps.assets import assets_client
from apps.assets.review import poll_asset_review, submit_asset_for_review
pairs = _replace_ref_assets(team, references)
if not pairs:
raise ValueError("请先上传参考视频")
states = []
for ref, asset in pairs:
label = ref.get("label") or asset.name or "参考素材"
if asset.review_status == "active" and asset.review_remote_id:
states.append("ready")
continue
if not assets_client.is_enabled():
raise ValueError(REVIEW_UNAVAILABLE)
if asset.review_status == "processing" and asset.review_remote_id:
poll_asset_review(asset)
asset.refresh_from_db(fields=["review_status", "review_remote_id", "review_error"])
elif asset.review_status != "active" or not asset.review_remote_id:
ok = submit_asset_for_review(asset, force=True)
asset.refresh_from_db(fields=["review_status", "review_remote_id", "review_error"])
if not ok and not asset.review_remote_id:
raise ValueError(f"{label}{REVIEW_SUBMIT_FAILED}")
if asset.review_status == "processing" and asset.review_remote_id:
poll_asset_review(asset)
asset.refresh_from_db(fields=["review_status", "review_remote_id", "review_error"])
if asset.review_status == "active" and asset.review_remote_id:
states.append("ready")
elif asset.review_status == "failed":
states.append("failed")
else:
states.append("pending")
if any(state == "failed" for state in states):
return "failed"
if all(state == "ready" for state in states):
return "ready"
return "pending"
def _refresh_replace_refs(team, references: list) -> list:
"""同一团队走 source=asset;官方跨团队素材把过审 id 写成 resolved_url=asset://。"""
out = []
for ref in references or []:
item = dict(ref)
raw_id = item.get("asset_id")
if not raw_id:
out.append(item)
continue
try:
parsed = uuid.UUID(str(raw_id))
except (TypeError, ValueError):
out.append(item)
continue
try:
asset = _load_replace_asset(team, parsed, item.get("label") or "参考素材")
except ValueError:
out.append(item)
continue
if asset.team_id == team.id:
item["source"] = "asset"
item.pop("resolved_url", None)
elif asset.review_remote_id:
item["source"] = "upload"
item["resolved_url"] = f"asset://{asset.review_remote_id}"
out.append(item)
return out
def _assert_replace_refs_ready(references: list) -> None:
missing = []
for ref in references or []:
raw_id = ref.get("asset_id")
if not raw_id:
continue
try:
parsed = uuid.UUID(str(raw_id))
except (TypeError, ValueError):
continue
asset = Asset.objects.filter(id=parsed).first()
if asset is None or not asset.review_remote_id or asset.review_status != "active":
missing.append(ref.get("label") or "参考素材")
if missing:
raise ValueError("素材尚未完成合规审核,请稍后再试")
def _optional_uuid(value, label: str): def _optional_uuid(value, label: str):
text = str(value or "").strip() text = str(value or "").strip()
if not text: if not text:
@@ -206,7 +502,7 @@ def _owned_ref(asset: Asset, *, kind: str, role: str, label: str) -> dict:
"type": kind, "type": kind,
"role": role, "role": role,
"label": label, "label": label,
"source": "upload", "source": "asset",
"asset_id": str(asset.id), "asset_id": str(asset.id),
} }
seconds = _asset_duration_seconds(asset) seconds = _asset_duration_seconds(asset)
+6 -5
View File
@@ -874,6 +874,8 @@ class VideoReplaceView(APIView):
"user_credit_insufficient" if "余额不足" in message "user_credit_insufficient" if "余额不足" in message
else "model_unavailable" if "模型未配置" in message else "model_unavailable" if "模型未配置" in message
else "provider_rate_limited" if "任务进行中" in message else "provider_rate_limited" if "任务进行中" in message
else "provider_unavailable" if "审核服务" in message
else "content_rejected" if "合规审核" in message
else "invalid_input" else "invalid_input"
) )
public_error = classify_generation_error( public_error = classify_generation_error(
@@ -911,19 +913,18 @@ class VideoReplaceView(APIView):
class VideoReplacePollView(APIView): class VideoReplacePollView(APIView):
"""POST /api/ai/video-replace/<id>/poll/ —— 与自由创作共用 finalize,只认复刻任务""" """POST /api/ai/video-replace/<id>/poll/ —— 审核中推进送审,生成中走 finalize。"""
def post(self, request, task_id): def post(self, request, task_id):
from .free_video import finalize_free_video from .video_replace import advance_video_replace, is_video_replace_task, serialize_video_replace_task
from .video_replace import is_video_replace_task, serialize_video_replace_task
team = get_current_team(request.user) team = get_current_team(request.user)
task = _free_video_task_queryset(team).filter(id=task_id).first() task = _free_video_task_queryset(team).filter(id=task_id).first()
if task is None or not is_video_replace_task(task): if task is None or not is_video_replace_task(task):
return Response({"detail": "任务不存在"}, status=status.HTTP_404_NOT_FOUND) return Response({"detail": "任务不存在"}, status=status.HTTP_404_NOT_FOUND)
if task.status in (AITask.Status.SUBMITTED, AITask.Status.POLLING): if task.status in (AITask.Status.CREATED, AITask.Status.SUBMITTED, AITask.Status.POLLING):
try: try:
task = finalize_free_video(task=task) task = advance_video_replace(task)
except Exception: # noqa: BLE001 — 单次轮询失败不终结任务 except Exception: # noqa: BLE001 — 单次轮询失败不终结任务
logger.warning("video replace poll failed for %s", task_id, exc_info=True) logger.warning("video replace poll failed for %s", task_id, exc_info=True)
task = _free_video_task_queryset(team).get(id=task.id) task = _free_video_task_queryset(team).get(id=task.id)
+19 -3
View File
@@ -59,6 +59,15 @@ def get_or_create_team_group(team) -> AssetReviewGroup:
return grp return grp
def _volcano_asset_type(asset: Asset) -> str:
"""火山 CreateAsset 的 AssetType:视频必须传 Video,不能默认 Image。"""
if asset.asset_type == Asset.Type.VIDEO:
return "Video"
if asset.asset_type == Asset.Type.AUDIO:
return "Audio"
return "Image"
def submit_asset_for_review(asset: Asset, *, force: bool = False) -> bool: def submit_asset_for_review(asset: Asset, *, force: bool = False) -> bool:
"""真人资产送审:建组(若无)→ 传素材 → 标 processing。出错只记日志,不抛。 """真人资产送审:建组(若无)→ 传素材 → 标 processing。出错只记日志,不抛。
返回是否真正进入审核(True=已标 processing;False=未送审/未配置/失败), 返回是否真正进入审核(True=已标 processing;False=未送审/未配置/失败),
@@ -70,12 +79,19 @@ def submit_asset_for_review(asset: Asset, *, force: bool = False) -> bool:
return False return False
if not force and asset.category not in Asset.REVIEW_CATEGORIES: if not force and asset.category not in Asset.REVIEW_CATEGORIES:
return False return False
if asset.review_remote_id and asset.review_status in ("active", "processing"):
return True
url = _asset_url(asset) url = _asset_url(asset)
if not url: if not url:
return False return False
try: try:
grp = get_or_create_team_group(asset.team) grp = get_or_create_team_group(asset.team)
remote_id = assets_client.create_asset(group_id=grp.remote_group_id, image_url=url, name=(asset.name or "person")[:64]) remote_id = assets_client.create_asset(
group_id=grp.remote_group_id,
image_url=url,
name=(asset.name or "person")[:64],
asset_type=_volcano_asset_type(asset),
)
if not remote_id: if not remote_id:
# 火山没回 Id:不要标 processing(否则 remote_id 为空、poll 永远早退、卡死黄),留空可重试 # 火山没回 Id:不要标 processing(否则 remote_id 为空、poll 永远早退、卡死黄),留空可重试
logger.warning("create_asset 返回空 id,asset %s 暂不送审(可重试)", asset.id) logger.warning("create_asset 返回空 id,asset %s 暂不送审(可重试)", asset.id)
@@ -128,11 +144,11 @@ def _unsubmitted_review_qs(*, team=None):
def _processing_review_qs(*, team=None): def _processing_review_qs(*, team=None):
# 不限 REVIEW_CATEGORIES:force 送审的上传视频/临时图也是 processing,worker 得盯到绿/红
qs = Asset.objects.filter( qs = Asset.objects.filter(
category__in=Asset.REVIEW_CATEGORIES,
review_status="processing", review_status="processing",
is_deleted=False, is_deleted=False,
) ).exclude(review_remote_id="")
if team is not None: if team is not None:
qs = qs.filter(team=team) qs = qs.filter(team=team)
return qs return qs
+25 -6
View File
@@ -12,6 +12,7 @@ import {
Package, Package,
RefreshCw, RefreshCw,
Replace, Replace,
ShieldCheck,
Upload, Upload,
UserRound, UserRound,
X, X,
@@ -53,6 +54,8 @@ const REPLACE_MODE_COPY = {
temporaryFallback: "临时商品素材", temporaryFallback: "临时商品素材",
generatingTitle: "正在进行商品复刻", generatingTitle: "正在进行商品复刻",
generatingCopy: "正在匹配商品外观与原片镜头", generatingCopy: "正在匹配商品外观与原片镜头",
reviewingTitle: "正在审核参考素材",
reviewingCopy: "真人视频需先通过合规审核,通过后自动开始复刻",
resultTitle: "商品复刻已完成", resultTitle: "商品复刻已完成",
resultPreview: "商品复刻预览", resultPreview: "商品复刻预览",
consistency: "商品一致性检查通过", consistency: "商品一致性检查通过",
@@ -76,6 +79,8 @@ const REPLACE_MODE_COPY = {
temporaryFallback: "临时角色素材", temporaryFallback: "临时角色素材",
generatingTitle: "正在进行角色复刻", generatingTitle: "正在进行角色复刻",
generatingCopy: "正在匹配角色外观、表情与原片动作", generatingCopy: "正在匹配角色外观、表情与原片动作",
reviewingTitle: "正在审核参考素材",
reviewingCopy: "真人视频需先通过合规审核,通过后自动开始复刻",
resultTitle: "角色复刻已完成", resultTitle: "角色复刻已完成",
resultPreview: "角色复刻预览", resultPreview: "角色复刻预览",
consistency: "角色一致性检查通过", consistency: "角色一致性检查通过",
@@ -254,6 +259,7 @@ export function VideoReplacePage({
const videoInputRef = useRef<HTMLInputElement>(null); const videoInputRef = useRef<HTMLInputElement>(null);
const tempInputRef = useRef<HTMLInputElement>(null); const tempInputRef = useRef<HTMLInputElement>(null);
const completedNoticeRef = useRef(""); const completedNoticeRef = useRef("");
const wasReviewingRef = useRef(false);
const videoConfigs = useMemo( const videoConfigs = useMemo(
() => modelConfigs.filter((config) => config.capability === "video" && config.status === "active"), () => modelConfigs.filter((config) => config.capability === "video" && config.status === "active"),
@@ -307,15 +313,17 @@ export function VideoReplacePage({
}, billingRates); }, billingRates);
const points = estimated.points || 220; const points = estimated.points || 220;
const generating = Boolean(job && isInFlight(job.status)) || submitting || videoUploading; const generating = Boolean(job && isInFlight(job.status)) || submitting || videoUploading;
const reviewing = submitting || Boolean(job && (job.review_stage === "reviewing" || job.status === "created"));
const hasResult = Boolean(job && job.status === "succeeded" && job.video_url); const hasResult = Boolean(job && job.status === "succeeded" && job.video_url);
const resultCopy = hasResult && job ? REPLACE_MODE_COPY[modeFromTask(job)] : copy; const resultCopy = hasResult && job ? REPLACE_MODE_COPY[modeFromTask(job)] : copy;
const panelClass = [ const panelClass = [
"video-result-panel replace-result-panel", "video-result-panel replace-result-panel",
generating ? "is-generating" : "", generating ? "is-generating" : "",
reviewing ? "is-reviewing" : "",
hasResult ? "has-result" : "", hasResult ? "has-result" : "",
].filter(Boolean).join(" "); ].filter(Boolean).join(" ");
const generateLabel = generating const generateLabel = generating
? "正在复刻…" ? (reviewing ? "正在审核素材…" : "正在复刻…")
: hasResult : hasResult
? `再次${copy.modeLabel} · 消耗 ${points} 积分` ? `再次${copy.modeLabel} · 消耗 ${points} 积分`
: `开始${copy.modeLabel} · 消耗 ${points} 积分`; : `开始${copy.modeLabel} · 消耗 ${points} 积分`;
@@ -364,8 +372,14 @@ export function VideoReplacePage({
const data = await api.pollVideoReplace(jobId); const data = await api.pollVideoReplace(jobId);
if (cancelled) return; if (cancelled) return;
setJob(data.task); setJob(data.task);
const stillReviewing = data.task.review_stage === "reviewing" || data.task.status === "created";
if (stillReviewing) wasReviewingRef.current = true;
else if (wasReviewingRef.current && isInFlight(data.task.status)) {
wasReviewingRef.current = false;
onNotify("success", "素材审核已通过,正在复刻");
}
if (isInFlight(data.task.status)) { if (isInFlight(data.task.status)) {
timer = window.setTimeout(poll, 2500); timer = window.setTimeout(poll, stillReviewing ? 2000 : 2500);
return; return;
} }
if (data.task.status === "succeeded") { if (data.task.status === "succeeded") {
@@ -575,7 +589,12 @@ export function VideoReplacePage({
setJobId(data.task.id); setJobId(data.task.id);
rememberJob(data.task.id); rememberJob(data.task.id);
completedNoticeRef.current = ""; completedNoticeRef.current = "";
onNotify("success", "视频复刻任务已开始"); wasReviewingRef.current = data.task.review_stage === "reviewing" || data.task.status === "created";
if (wasReviewingRef.current) {
onNotify("success", "已提交素材审核,通过后自动开始复刻");
} else {
onNotify("success", "视频复刻任务已开始");
}
if (!isInFlight(data.task.status) && data.task.status === "succeeded") { if (!isInFlight(data.task.status) && data.task.status === "succeeded") {
onNotify("success", "视频复刻成片已生成"); onNotify("success", "视频复刻成片已生成");
void loadHistory(); void loadHistory();
@@ -879,11 +898,11 @@ export function VideoReplacePage({
<div className="replace-generating-state" role="status" aria-live="polite"> <div className="replace-generating-state" role="status" aria-live="polite">
<div className="replace-generating-content"> <div className="replace-generating-content">
<div className="replace-generating-visual"> <div className="replace-generating-visual">
<span className="replace-generating-frame"><Clapperboard /></span> <span className="replace-generating-frame">{reviewing ? <ShieldCheck /> : <Clapperboard />}</span>
<span className="replace-generating-product">{replaceMode === "character" ? <UserRound /> : <Package />}</span> <span className="replace-generating-product">{replaceMode === "character" ? <UserRound /> : <Package />}</span>
</div> </div>
<strong>{copy.generatingTitle}</strong> <strong>{reviewing ? copy.reviewingTitle : copy.generatingTitle}</strong>
<span>{copy.generatingCopy}</span> <span>{reviewing ? copy.reviewingCopy : copy.generatingCopy}</span>
<div className="replace-generating-bar" aria-hidden="true"><span /></div> <div className="replace-generating-bar" aria-hidden="true"><span /></div>
</div> </div>
</div> </div>
+1
View File
@@ -662,6 +662,7 @@ export type FreeVideoTask = {
subject_source?: "library" | "temporary" | ""; subject_source?: "library" | "temporary" | "";
product_id?: string; product_id?: string;
model_id?: string; model_id?: string;
review_stage?: "reviewing" | "";
references: FreeVideoRef[]; references: FreeVideoRef[];
estimated_tokens: number; estimated_tokens: number;
actual_tokens: number; actual_tokens: number;