完善视频复刻
This commit is contained in:
@@ -47,6 +47,7 @@ RATIOS = {"21:9", "16:9", "4:3", "1:1", "3:4", "9:16"}
|
||||
RESOLUTIONS = {"480p", "720p", "1080p", "4k"}
|
||||
MODES = {"universal", "keyframe"}
|
||||
IN_FLIGHT_STATUSES = (
|
||||
AITask.Status.CREATED, # 视频复刻审核中也占并发,避免连点刷出一堆待审任务
|
||||
AITask.Status.RESERVED,
|
||||
AITask.Status.SUBMITTED,
|
||||
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)
|
||||
continue
|
||||
|
||||
# 直传素材(已上传 TOS 的直链)
|
||||
# 直传素材(已上传 TOS 的直链)。resolved_url 可覆盖为 asset://(官方模特跨团队等)
|
||||
push_url = str(ref.get("resolved_url") or url)
|
||||
if ref_type == "image":
|
||||
# 参考图模式下所有图 role 必须 reference_image;keyframe 用 first_frame/last_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":
|
||||
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":
|
||||
asset_type = _push("audio", url, role or "reference_audio", duration)
|
||||
asset_type = _push("audio", push_url, role or "reference_audio", duration)
|
||||
else:
|
||||
logger.warning("unknown ref_type=%s url=%s label=%s, skipped", ref_type, url, label)
|
||||
continue
|
||||
@@ -366,6 +368,11 @@ def _reap_stale_free_video_tasks(*, team) -> None:
|
||||
|
||||
video_policy = load_model_routing_policy().video
|
||||
buckets = [
|
||||
(
|
||||
[AITask.Status.CREATED],
|
||||
{"updated_at__lt": now - timedelta(minutes=16)},
|
||||
"素材审核超时(自动回收)",
|
||||
),
|
||||
(
|
||||
[AITask.Status.RESERVED],
|
||||
{"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.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:
|
||||
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},
|
||||
)
|
||||
response, provider_task_id = routed.value
|
||||
# execute_model_call 以原子 F 表达式累计实际尝试的平台成本;刷新内存对象,
|
||||
# 保证本方法返回值与数据库中的 AITask.base_cost 完全一致。
|
||||
task.refresh_from_db(fields=["base_cost"])
|
||||
task.provider_task_id = provider_task_id
|
||||
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
|
||||
|
||||
mark_remote_asset_unavailable(
|
||||
team=team,
|
||||
team=task.team,
|
||||
local_asset_id=target["local_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)
|
||||
return task
|
||||
|
||||
# worker 兜底轮询(自重排);派发失败仅 log,前端主动 poll 仍能收尾
|
||||
try:
|
||||
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
|
||||
logger.error("poll_free_video_task enqueue failed; relying on client polling", exc_info=True)
|
||||
|
||||
|
||||
@@ -52,6 +52,11 @@ from apps.projects.models import (
|
||||
|
||||
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:
|
||||
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()
|
||||
|
||||
|
||||
def get_storyboard_image_model() -> ModelConfig:
|
||||
"""故事板出图钉 YunQi gpt-image-2 多图 edits(与手工测通的 curl 同一条链路)。
|
||||
找不到再回落默认图像模型,避免测试/未 seed 环境直接挂。"""
|
||||
pinned = (
|
||||
def get_storyboard_image_model() -> ModelConfig | None:
|
||||
"""故事板出图只走 GPT 图像模型(gpt-image / gpt-image-2)。
|
||||
|
||||
不回落默认图像模型,也不走火山 Seedream:用户明确要求故事板无论怎样都不改成其他生图模型。
|
||||
"""
|
||||
qs = (
|
||||
ModelConfig.objects.select_related("provider")
|
||||
.filter(
|
||||
capability=ModelConfig.Capability.IMAGE,
|
||||
status=ModelConfig.Status.ACTIVE,
|
||||
provider__status="active",
|
||||
provider__name="yunqi",
|
||||
name="gpt-image-2",
|
||||
name__icontains="gpt-image",
|
||||
)
|
||||
.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":
|
||||
@@ -101,12 +109,6 @@ def resolve_image_model(key: str | None) -> "ModelConfig | None":
|
||||
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:
|
||||
"""普通用户公开名称保持稳定;Fallback 的真实模型只在管理员尝试链中展示。"""
|
||||
|
||||
@@ -2781,7 +2783,7 @@ def submit_storyboard(*, project, user, prompt: str = "", shot_ids: list | None
|
||||
if adopted_script is None:
|
||||
raise ValueError("script must be adopted before generating storyboard")
|
||||
if get_storyboard_image_model() is None:
|
||||
raise ValueError("no active image model configured")
|
||||
raise ValueError("故事板只使用 GPT 图像模型,当前没有启用 gpt-image-2")
|
||||
if prompt:
|
||||
meta = dict(project.metadata or {})
|
||||
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)
|
||||
project = shot.project
|
||||
segment = shot.script_segment
|
||||
model_config = task.model_config
|
||||
reservation = task.credit_reservation
|
||||
extra_prompt = (project.metadata or {}).get("storyboard_prompt", "") or ""
|
||||
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.save(update_fields=["status", "updated_at"])
|
||||
try:
|
||||
use_model_routing = bool((task.request_payload or {}).get("model_routing_v1"))
|
||||
provider = None if use_model_routing else get_image_provider(model_config)
|
||||
# 故事板无论任务上挂了什么模型、是否开了路由,都只走 GPT 图像;失败也不换 Seedream。
|
||||
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 []
|
||||
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=角色/场景/商品」+锁脸锁商品)
|
||||
frame_prompt = build_storyboard_frame_prompt_refs(project, segment, refs, extra_prompt)
|
||||
else:
|
||||
@@ -3093,40 +3097,24 @@ def _storyboard_shot_worker(task_id, shot_id, user_id) -> None:
|
||||
else (task.request_payload.get("prompt") or "")
|
||||
)
|
||||
|
||||
if use_model_routing:
|
||||
routed = execute_routed_image_request(
|
||||
task=task,
|
||||
primary_model=model_config,
|
||||
prompt=frame_prompt,
|
||||
reference_images=ref_urls,
|
||||
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,
|
||||
},
|
||||
if ref_urls and hasattr(provider, "image_edit"):
|
||||
response = _call_image_with_retry(
|
||||
lambda: provider.image_edit(
|
||||
model=model_config.name,
|
||||
prompt=frame_prompt,
|
||||
images=ref_urls,
|
||||
size=frame_size,
|
||||
)
|
||||
)
|
||||
response, media = routed.value
|
||||
else:
|
||||
if ref_urls and hasattr(provider, "image_edit"):
|
||||
response = _call_image_with_retry(
|
||||
lambda: provider.image_edit(
|
||||
model=model_config.name,
|
||||
prompt=frame_prompt,
|
||||
images=ref_urls,
|
||||
size=frame_size,
|
||||
)
|
||||
response = _call_image_with_retry(
|
||||
lambda: provider.image_generation(
|
||||
model=model_config.name,
|
||||
endpoint=model_config.endpoint,
|
||||
prompt=frame_prompt,
|
||||
)
|
||||
else:
|
||||
response = _call_image_with_retry(
|
||||
lambda: provider.image_generation(
|
||||
model=model_config.name,
|
||||
endpoint=model_config.endpoint,
|
||||
prompt=frame_prompt,
|
||||
)
|
||||
)
|
||||
media = provider.extract_first_media_url(response)
|
||||
)
|
||||
media = provider.extract_first_media_url(response)
|
||||
asset = _store_generated_media(
|
||||
team=project.team, user=user, project=project, task=task, media=media,
|
||||
name=f"{project.name}-storyboard-{shot.sort_order + 1}",
|
||||
@@ -3206,6 +3194,8 @@ def poll_storyboard(*, project, user) -> dict:
|
||||
if v
|
||||
}
|
||||
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 ""
|
||||
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))
|
||||
@@ -3217,12 +3207,8 @@ def poll_storyboard(*, project, user) -> dict:
|
||||
"model": model_config.name, "endpoint": model_config.endpoint,
|
||||
"prompt": build_storyboard_frame_prompt(project, segment, extra_prompt) if segment is not None else "",
|
||||
"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())
|
||||
threading.Thread(
|
||||
target=_storyboard_shot_worker, args=(str(task.id), str(shot.id), str(user.id)), daemon=True
|
||||
|
||||
@@ -78,12 +78,16 @@ def poll_free_video_task(self, task_id: str, attempt: int = 0) -> str:
|
||||
与前端主动 poll 并存不双扣。轮询本身出错不重试(max_retries=0),下一次自重排继续。"""
|
||||
from apps.ai.free_video import finalize_free_video
|
||||
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()
|
||||
if task is None:
|
||||
return task_id
|
||||
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 — 单次轮询失败(网络抖动等)不终结任务,等下一轮
|
||||
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):
|
||||
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)
|
||||
return task_id
|
||||
|
||||
|
||||
@@ -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):
|
||||
def setUp(self):
|
||||
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()
|
||||
|
||||
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.refresh_from_db()
|
||||
self.shot.refresh_from_db()
|
||||
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
|
||||
self.assertTrue(task.request_payload["model_routing_v1"])
|
||||
self.assertEqual(task.base_cost, Decimal("0.5000"))
|
||||
attempt = task.model_attempts.get()
|
||||
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))
|
||||
self.assertNotIn("model_routing_v1", task.request_payload)
|
||||
self.assertEqual(task.model_config_id, primary.id)
|
||||
self.assertEqual(task.model_config.name, "gpt-image-2")
|
||||
call = self.provider_mocks[primary.id].image_edit.call_args
|
||||
self.assertEqual(
|
||||
call.kwargs["images"],
|
||||
@@ -179,6 +177,7 @@ class StoryboardRoutingTests(TestCase):
|
||||
)
|
||||
self.assertEqual(call.kwargs["prompt"], "带参考图编号与锁定约束的故事板提示词")
|
||||
self.assertEqual(call.kwargs["size"], "1024x1536")
|
||||
self.assertEqual(call.kwargs["model"], "gpt-image-2")
|
||||
self.assertIsNotNone(self.shot.adopted_version_id)
|
||||
self.assertEqual(self.shot.adopted_version.task_id, task.id)
|
||||
self.assertEqual(self.shot.versions.count(), 1)
|
||||
@@ -188,44 +187,40 @@ class StoryboardRoutingTests(TestCase):
|
||||
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")
|
||||
primary = self.model(self.provider("storyboard-wide", 20), "gpt-image-2")
|
||||
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):
|
||||
def test_gpt_failure_does_not_fallback_to_seedream(self):
|
||||
primary = self.model(
|
||||
self.provider("storyboard-fallback-primary", 100),
|
||||
"storyboard-primary",
|
||||
self.provider("storyboard-gpt", 100),
|
||||
"gpt-image-2",
|
||||
base_cost="0.25",
|
||||
)
|
||||
candidate = self.model(
|
||||
self.provider("volcano", 10),
|
||||
"storyboard-candidate",
|
||||
seedream = self.model(
|
||||
self.provider("storyboard-volcano-fallback", 10),
|
||||
"seedream-4-5-251128",
|
||||
outbound=False,
|
||||
base_cost="0.75",
|
||||
)
|
||||
self.provider_mocks[primary.id] = self._new_provider_mock()
|
||||
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.refresh_from_db()
|
||||
attempts = list(task.model_attempts.all())
|
||||
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
|
||||
self.assertEqual([a.model_config_id for a in attempts], [primary.id, primary.id, candidate.id])
|
||||
self.assertTrue(attempts[-1].is_fallback)
|
||||
self.assertEqual(task.base_cost, Decimal("0.7500"))
|
||||
self.assertEqual(task.status, AITask.Status.FAILED)
|
||||
self.assertEqual(task.model_config_id, primary.id)
|
||||
self.provider_mocks[seedream.id].image_edit.assert_not_called()
|
||||
self.provider_mocks[seedream.id].image_generation.assert_not_called()
|
||||
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.RELEASE), 0)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 0)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 1)
|
||||
|
||||
def test_all_candidates_fail_releases_and_keeps_old_adopted_version(self):
|
||||
primary = self.model(self.provider("storyboard-all-primary", 100), "storyboard-primary")
|
||||
candidate = self.model(
|
||||
self.provider("volcano", 10), "storyboard-candidate", outbound=False
|
||||
)
|
||||
primary = self.model(self.provider("storyboard-all-primary", 100), "gpt-image-2")
|
||||
self.model(self.provider("storyboard-volcano-keep", 10), "seedream-4-5-251128", outbound=False)
|
||||
old_asset = Asset.objects.create(
|
||||
team=self.team,
|
||||
created_by=self.user,
|
||||
@@ -242,16 +237,14 @@ class StoryboardRoutingTests(TestCase):
|
||||
)
|
||||
self.shot.adopted_version = old_version
|
||||
self.shot.save(update_fields=["adopted_version", "updated_at"])
|
||||
for model in (primary, candidate):
|
||||
self.provider_mocks[model.id] = self._new_provider_mock()
|
||||
self.provider_mocks[model.id].image_edit.side_effect = requests.ConnectionError("offline")
|
||||
self.provider_mocks[primary.id] = self._new_provider_mock()
|
||||
self.provider_mocks[primary.id].image_edit.side_effect = requests.ConnectionError("offline")
|
||||
|
||||
task = self.enqueue()
|
||||
|
||||
task.refresh_from_db()
|
||||
self.shot.refresh_from_db()
|
||||
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.versions.count(), 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(CreditAccount.objects.get(team=self.team).reserved_balance, Decimal("0"))
|
||||
|
||||
def test_no_reference_uses_image_generation_and_keeps_vertical_capability(self):
|
||||
primary = self.model(self.provider("storyboard-no-ref", 20), "storyboard-no-ref")
|
||||
def test_no_reference_uses_image_generation(self):
|
||||
primary = self.model(self.provider("storyboard-no-ref", 20), "gpt-image-2")
|
||||
self.refs = []
|
||||
task = self.enqueue()
|
||||
|
||||
task.refresh_from_db()
|
||||
attempt = task.model_attempts.get()
|
||||
self.assertEqual(attempt.operation, "image_generate")
|
||||
self.assertEqual(attempt.request_summary["reference_images"], 0)
|
||||
self.assertEqual(attempt.request_summary["aspect_ratio"], "9:16")
|
||||
self.provider_mocks[primary.id].image_edit.assert_not_called()
|
||||
call = self.provider_mocks[primary.id].image_generation.call_args
|
||||
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):
|
||||
primary = self.model(
|
||||
self.provider("volcano", 10), "seedream-storyboard", outbound=False
|
||||
def test_seedream_is_ignored_even_if_it_is_the_default_image_model(self):
|
||||
seedream = self.model(
|
||||
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.refresh_from_db()
|
||||
call = self.provider_mocks[primary.id].image_generation.call_args
|
||||
self.assertEqual(
|
||||
call.kwargs["image"],
|
||||
["http://example.test/person.png", "http://example.test/product.png"],
|
||||
)
|
||||
self.assertEqual(call.kwargs["size"], "1024x1536")
|
||||
self.assertEqual(task.model_attempts.get().public_model_name, primary.display_name)
|
||||
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
|
||||
self.assertEqual(task.model_config_id, gpt.id)
|
||||
self.provider_mocks[seedream.id].image_generation.assert_not_called()
|
||||
self.provider_mocks[gpt.id].image_edit.assert_called_once()
|
||||
|
||||
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(
|
||||
script_version=self.segment.script_version,
|
||||
sort_order=1,
|
||||
@@ -317,6 +307,6 @@ class StoryboardRoutingTests(TestCase):
|
||||
{task.request_payload["storyboard_shot"] for task in tasks},
|
||||
{str(self.shot.id), str(second_shot.id)},
|
||||
)
|
||||
self.assertTrue(all(task.request_payload["model_routing_v1"] for task in tasks))
|
||||
self.assertTrue(all(task.base_cost == Decimal("0") for task in tasks))
|
||||
self.assertTrue(all("model_routing_v1" not in task.request_payload 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))
|
||||
|
||||
@@ -3,6 +3,7 @@
|
||||
运行:DB_ENGINE=sqlite python manage.py test apps.ai.test_video_replace --settings=airshelf.settings.test
|
||||
"""
|
||||
from unittest.mock import MagicMock, patch
|
||||
from uuid import uuid4
|
||||
|
||||
from django.test import TestCase
|
||||
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.models import AITask, ModelConfig
|
||||
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.billing.models import CreditAccount
|
||||
from apps.billing.models import CreditAccount, CreditReservation
|
||||
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(
|
||||
team=team,
|
||||
created_by=user,
|
||||
@@ -25,6 +26,8 @@ def _asset(team, user, *, kind=Asset.Type.IMAGE, name="素材", duration_ms=None
|
||||
asset_type=kind,
|
||||
source=Asset.Source.AI_GENERATED,
|
||||
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(
|
||||
asset=asset,
|
||||
@@ -84,6 +87,14 @@ class SubmitVideoReplaceTests(TestCase):
|
||||
roles = [item.get("role") for item in content]
|
||||
self.assertIn("reference_video", 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):
|
||||
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):
|
||||
resp = self.client.post("/api/ai/video-replace/", {"replace_mode": "product"}, format="json")
|
||||
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)
|
||||
|
||||
@@ -2,23 +2,47 @@
|
||||
|
||||
不新建任务类型、不接检测/抠图。提交仍走 submit_free_video,只把
|
||||
feature=video_replace 和 replace_mode 写进 payload,提示词由后端写死。
|
||||
|
||||
真人参考必须先送火山素材库审核,过审后用 asset:// 生成;审核失败不扣费。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import uuid
|
||||
from decimal import Decimal
|
||||
|
||||
from django.conf import settings
|
||||
from django.db import transaction
|
||||
from django.db.models import Q
|
||||
|
||||
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 .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 .models import AITask, ModelConfig
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
FEATURE = "video_replace"
|
||||
REPLACE_MODES = {"product", "character"}
|
||||
MAX_IMAGES = 9
|
||||
LEGACY_PROMPT_PREFIX = "[视频复刻]"
|
||||
REVIEW_UNAVAILABLE = "素材审核服务暂不可用,请稍后重试"
|
||||
REVIEW_FAILED = "参考素材未通过真人合规审核,请更换视频或图片后重试"
|
||||
REVIEW_SUBMIT_FAILED = "素材提交审核失败,请稍后重试"
|
||||
|
||||
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)
|
||||
payload = task.request_payload 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({
|
||||
"feature": FEATURE,
|
||||
"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 "",
|
||||
"product_id": payload.get("product_id") or "",
|
||||
"model_id": payload.get("model_id") or "",
|
||||
"review_stage": "reviewing" if reviewing else "",
|
||||
})
|
||||
return data
|
||||
|
||||
|
||||
def submit_video_replace(*, team, user, params: dict):
|
||||
"""校验素材 → 套提示词 → 复用 free_video 提交。失败抛 ValueError。"""
|
||||
"""校验素材 → 送审 → 已过审则直接生成,否则建 CREATED 任务等绿盾。失败抛 ValueError。"""
|
||||
replace_mode = str(params.get("replace_mode") or "").strip()
|
||||
if replace_mode not in REPLACE_MODES:
|
||||
raise ValueError("请选择替换商品或替换角色")
|
||||
@@ -107,35 +133,305 @@ def submit_video_replace(*, team, user, params: dict):
|
||||
_owned_ref(video, kind="video", role="reference_video", label="参考视频"),
|
||||
*image_refs,
|
||||
]
|
||||
return submit_free_video(
|
||||
team=team,
|
||||
user=user,
|
||||
params={
|
||||
"prompt": prompt,
|
||||
"mode": "universal",
|
||||
"model": 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"),
|
||||
"duration": duration,
|
||||
"seed": params.get("seed", -1),
|
||||
"generate_audio": True,
|
||||
"references": references,
|
||||
"feature": FEATURE,
|
||||
"extra_payload": {
|
||||
"replace_mode": replace_mode,
|
||||
"subject_name": subject_name,
|
||||
"subject_source": subject_source,
|
||||
"product_id": str(product_id) if product_id else "",
|
||||
"model_id": str(model_id) if model_id else "",
|
||||
},
|
||||
},
|
||||
)
|
||||
review_state = _ensure_replace_refs_reviewed(team, references)
|
||||
if review_state == "failed":
|
||||
raise ValueError(REVIEW_FAILED)
|
||||
references = _refresh_replace_refs(team, references)
|
||||
extra = {
|
||||
"replace_mode": replace_mode,
|
||||
"subject_name": subject_name,
|
||||
"subject_source": subject_source,
|
||||
"product_id": str(product_id) if product_id else "",
|
||||
"model_id": str(model_id) if model_id else "",
|
||||
"review_pending": review_state != "ready",
|
||||
}
|
||||
submit_params = {
|
||||
"prompt": prompt,
|
||||
"mode": "universal",
|
||||
"model": 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"),
|
||||
"duration": duration,
|
||||
"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:
|
||||
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):
|
||||
text = str(value or "").strip()
|
||||
if not text:
|
||||
@@ -206,7 +502,7 @@ def _owned_ref(asset: Asset, *, kind: str, role: str, label: str) -> dict:
|
||||
"type": kind,
|
||||
"role": role,
|
||||
"label": label,
|
||||
"source": "upload",
|
||||
"source": "asset",
|
||||
"asset_id": str(asset.id),
|
||||
}
|
||||
seconds = _asset_duration_seconds(asset)
|
||||
|
||||
@@ -874,6 +874,8 @@ class VideoReplaceView(APIView):
|
||||
"user_credit_insufficient" if "余额不足" in message
|
||||
else "model_unavailable" 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"
|
||||
)
|
||||
public_error = classify_generation_error(
|
||||
@@ -911,19 +913,18 @@ class VideoReplaceView(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):
|
||||
from .free_video import finalize_free_video
|
||||
from .video_replace import is_video_replace_task, serialize_video_replace_task
|
||||
from .video_replace import advance_video_replace, is_video_replace_task, serialize_video_replace_task
|
||||
|
||||
team = get_current_team(request.user)
|
||||
task = _free_video_task_queryset(team).filter(id=task_id).first()
|
||||
if task is None or not is_video_replace_task(task):
|
||||
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:
|
||||
task = finalize_free_video(task=task)
|
||||
task = advance_video_replace(task)
|
||||
except Exception: # noqa: BLE001 — 单次轮询失败不终结任务
|
||||
logger.warning("video replace poll failed for %s", task_id, exc_info=True)
|
||||
task = _free_video_task_queryset(team).get(id=task.id)
|
||||
|
||||
@@ -59,6 +59,15 @@ def get_or_create_team_group(team) -> AssetReviewGroup:
|
||||
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:
|
||||
"""真人资产送审:建组(若无)→ 传素材 → 标 processing。出错只记日志,不抛。
|
||||
返回是否真正进入审核(True=已标 processing;False=未送审/未配置/失败),
|
||||
@@ -70,12 +79,19 @@ def submit_asset_for_review(asset: Asset, *, force: bool = False) -> bool:
|
||||
return False
|
||||
if not force and asset.category not in Asset.REVIEW_CATEGORIES:
|
||||
return False
|
||||
if asset.review_remote_id and asset.review_status in ("active", "processing"):
|
||||
return True
|
||||
url = _asset_url(asset)
|
||||
if not url:
|
||||
return False
|
||||
try:
|
||||
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:
|
||||
# 火山没回 Id:不要标 processing(否则 remote_id 为空、poll 永远早退、卡死黄),留空可重试
|
||||
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):
|
||||
# 不限 REVIEW_CATEGORIES:force 送审的上传视频/临时图也是 processing,worker 得盯到绿/红
|
||||
qs = Asset.objects.filter(
|
||||
category__in=Asset.REVIEW_CATEGORIES,
|
||||
review_status="processing",
|
||||
is_deleted=False,
|
||||
)
|
||||
).exclude(review_remote_id="")
|
||||
if team is not None:
|
||||
qs = qs.filter(team=team)
|
||||
return qs
|
||||
|
||||
Reference in New Issue
Block a user