后端生成闸+多项修复;前端全站更新;QA 审计与报告
后端: - 新增 celery_health 生成前置闸——无 worker 在线时图片/视频生成入口 一律 503,防"提交到 ARK 后无人轮询、结果悬空+额度冻结"的数据丢失 - 拼接导出:帧率跟随源众数、字幕逐句重映射、原声/BGM 混音修复 - 资金核算审计脚本 + genesis 账目回填 + 滞留预留清理命令 - 接入 yunqi provider 与豆包 TTS 模型(catalog/migrations/bootstrap 命令) 前端:全站页面更新(pipeline/library/products/projects/team/account 等), 新增共享 pager 分页组件 QA:刷新 function-audit 全量输出,新增 full-qa 报告 文档:BP 产品介绍资料、design/CLAUDE.md Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
@@ -1,6 +1,6 @@
|
||||
from rest_framework import serializers
|
||||
|
||||
from apps.billing.models import CreditAccount
|
||||
from apps.billing.models import CreditAccount, CreditLedger
|
||||
|
||||
from .models import LoginSession, Team, TeamMember, User, UserPreference
|
||||
|
||||
@@ -82,10 +82,19 @@ class RegisterSerializer(serializers.Serializer):
|
||||
owner=user,
|
||||
)
|
||||
TeamMember.objects.create(team=team, user=user, role=TeamMember.Role.OWNER)
|
||||
CreditAccount.objects.create(
|
||||
team=team,
|
||||
balance=Decimal(str(settings.DEFAULT_TRIAL_CREDITS)),
|
||||
)
|
||||
trial = Decimal(str(settings.DEFAULT_TRIAL_CREDITS))
|
||||
CreditAccount.objects.create(team=team, balance=trial)
|
||||
# 开户赠送额度必须进流水,否则这笔钱在账单里查无凭证、无法对账(余额从第一条起就可追溯)
|
||||
if trial > 0:
|
||||
CreditLedger.objects.create(
|
||||
team=team,
|
||||
user=user,
|
||||
ledger_type=CreditLedger.Type.RECHARGE,
|
||||
amount=trial,
|
||||
balance_after=trial,
|
||||
reason="新用户试用额度赠送",
|
||||
metadata={"kind": "trial_grant"},
|
||||
)
|
||||
return {"user": user, "team": team}
|
||||
|
||||
|
||||
|
||||
@@ -1,8 +1,11 @@
|
||||
from decimal import Decimal
|
||||
|
||||
from django.conf import settings
|
||||
from django.test import TestCase
|
||||
from rest_framework.test import APIClient
|
||||
|
||||
from apps.accounts.models import Team, TeamMember, User
|
||||
from apps.billing.models import CreditAccount
|
||||
from apps.billing.models import CreditAccount, CreditLedger
|
||||
|
||||
|
||||
class AuthApiTests(TestCase):
|
||||
@@ -27,3 +30,20 @@ class AuthApiTests(TestCase):
|
||||
self.assertTrue(TeamMember.objects.filter(team=team, user=user, role=TeamMember.Role.OWNER).exists())
|
||||
self.assertTrue(CreditAccount.objects.filter(team=team).exists())
|
||||
|
||||
def test_register_trial_grant_is_recorded_in_ledger(self):
|
||||
"""开户赠送额度必须进流水:余额=赠送额,且有一条 balance_after=赠送额的 genesis 流水,账目从第一条起可对账。"""
|
||||
client = APIClient()
|
||||
client.post(
|
||||
"/api/auth/register/",
|
||||
{"username": "ledger-owner", "password": "strong-password", "team_name": "Ledger Team"},
|
||||
format="json",
|
||||
)
|
||||
team = Team.objects.get(name="Ledger Team")
|
||||
trial = Decimal(str(settings.DEFAULT_TRIAL_CREDITS))
|
||||
account = CreditAccount.objects.get(team=team)
|
||||
self.assertEqual(account.balance, trial)
|
||||
genesis = CreditLedger.objects.filter(team=team, ledger_type=CreditLedger.Type.RECHARGE)
|
||||
self.assertEqual(genesis.count(), 1)
|
||||
self.assertEqual(genesis.first().amount, trial)
|
||||
self.assertEqual(genesis.first().balance_after, trial) # 流水终点 == 账户余额,可对账
|
||||
|
||||
|
||||
@@ -4,6 +4,22 @@ VOLCANO_PROVIDER = {
|
||||
"base_url": "https://ark.cn-beijing.volces.com/api/v3",
|
||||
}
|
||||
|
||||
YUNQI_PROVIDER = {
|
||||
"name": "yunqi",
|
||||
"display_name": "YunQi AI(New API 网关)",
|
||||
"base_url": "https://www.yunqiai.chat/v1",
|
||||
}
|
||||
|
||||
YUNQI_MODELS = [
|
||||
{
|
||||
"display_name": "GPT-Image-2",
|
||||
"name": "gpt-image-2",
|
||||
"capability": "image",
|
||||
"endpoint": "images/generations",
|
||||
"metadata": {"modes": ["text"], "response_format": "b64_json", "source": "https://www.yunqiai.chat"},
|
||||
},
|
||||
]
|
||||
|
||||
VOLCANO_MODELS = [
|
||||
{
|
||||
"display_name": "Doubao-Seed-2.0-Pro",
|
||||
|
||||
+46
@@ -0,0 +1,46 @@
|
||||
# Generated by Django 5.1.15 on 2026-06-11 07:14
|
||||
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
("ai", "0002_initial"),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.AlterField(
|
||||
model_name="aitask",
|
||||
name="task_type",
|
||||
field=models.CharField(
|
||||
choices=[
|
||||
("script_generation", "Script Generation"),
|
||||
("script_optimization", "Script Optimization"),
|
||||
("product_image", "Product Image"),
|
||||
("person_image", "Person Image"),
|
||||
("scene_image", "Scene Image"),
|
||||
("storyboard", "Storyboard"),
|
||||
("video_segment", "Video Segment"),
|
||||
("voiceover", "Voiceover"),
|
||||
("export", "Export"),
|
||||
],
|
||||
max_length=48,
|
||||
),
|
||||
),
|
||||
migrations.AlterField(
|
||||
model_name="modelconfig",
|
||||
name="capability",
|
||||
field=models.CharField(
|
||||
choices=[
|
||||
("text", "Text"),
|
||||
("image", "Image"),
|
||||
("video", "Video"),
|
||||
("vision", "Vision"),
|
||||
("audio", "Audio"),
|
||||
("export", "Export"),
|
||||
],
|
||||
max_length=32,
|
||||
),
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,36 @@
|
||||
# 旁白配音(TTS):给 volcengine provider 补一条语音合成 ModelConfig。
|
||||
# unit_price=0 时计费走 estimate_cost 默认 ¥1/次(一次=整批旁白合成),与其它能力同口径。
|
||||
from django.db import migrations
|
||||
|
||||
|
||||
def seed_tts_model(apps, schema_editor):
|
||||
ModelProvider = apps.get_model("ai", "ModelProvider")
|
||||
ModelConfig = apps.get_model("ai", "ModelConfig")
|
||||
provider = ModelProvider.objects.filter(name="volcengine").first()
|
||||
if provider is None:
|
||||
return
|
||||
ModelConfig.objects.get_or_create(
|
||||
provider=provider,
|
||||
name="doubao-voice-bigtts",
|
||||
defaults={
|
||||
"display_name": "豆包语音合成",
|
||||
"capability": "audio",
|
||||
"endpoint": "api/v1/tts",
|
||||
"unit_price": "0.0000",
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def unseed_tts_model(apps, schema_editor):
|
||||
ModelConfig = apps.get_model("ai", "ModelConfig")
|
||||
ModelConfig.objects.filter(name="doubao-voice-bigtts").delete()
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
dependencies = [
|
||||
("ai", "0003_alter_aitask_task_type_alter_modelconfig_capability"),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.RunPython(seed_tts_model, unseed_tts_model),
|
||||
]
|
||||
@@ -24,6 +24,7 @@ class ModelConfig(TimeStampedModel):
|
||||
IMAGE = "image", "Image"
|
||||
VIDEO = "video", "Video"
|
||||
VISION = "vision", "Vision"
|
||||
AUDIO = "audio", "Audio"
|
||||
EXPORT = "export", "Export"
|
||||
|
||||
class Status(models.TextChoices):
|
||||
@@ -56,6 +57,7 @@ class AITask(TeamOwnedModel):
|
||||
SCENE_IMAGE = "scene_image", "Scene Image"
|
||||
STORYBOARD = "storyboard", "Storyboard"
|
||||
VIDEO_SEGMENT = "video_segment", "Video Segment"
|
||||
VOICEOVER = "voiceover", "Voiceover"
|
||||
EXPORT = "export", "Export"
|
||||
|
||||
class Status(models.TextChoices):
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
from .base import AIProvider, AIProviderResult
|
||||
from .volcano import VolcanoArkProvider
|
||||
from .volcano import TtsNotConfigured, VolcanoArkProvider, VolcanoTtsProvider
|
||||
from .yunqi import YunqiProvider
|
||||
|
||||
|
||||
__all__ = ["AIProvider", "AIProviderResult", "VolcanoArkProvider"]
|
||||
|
||||
__all__ = ["AIProvider", "AIProviderResult", "TtsNotConfigured", "VolcanoArkProvider", "VolcanoTtsProvider", "YunqiProvider"]
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
from dataclasses import dataclass
|
||||
import base64
|
||||
import uuid
|
||||
from io import BytesIO
|
||||
from typing import Any
|
||||
|
||||
@@ -171,3 +172,60 @@ class VolcanoArkProvider:
|
||||
content_type = header.split(";")[0].replace("data:", "") or "application/octet-stream"
|
||||
return BytesIO(base64.b64decode(raw)), content_type
|
||||
return BytesIO(base64.b64decode(media)), "image/png"
|
||||
|
||||
|
||||
class TtsNotConfigured(ValueError):
|
||||
"""语音合成凭证未配置(VOLC_TTS_APPID / VOLC_TTS_ACCESS_TOKEN)。"""
|
||||
|
||||
|
||||
@dataclass
|
||||
class VolcanoTtsProvider:
|
||||
"""豆包语音合成大模型(openspeech V1 HTTP,非流式)。
|
||||
注意:与 ARK 不是同一套钥匙,要在 火山引擎控制台→语音技术→语音合成大模型 创建应用拿 APPID/Access Token。"""
|
||||
|
||||
appid: str | None = None
|
||||
access_token: str | None = None
|
||||
cluster: str | None = None
|
||||
base_url: str | None = None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
cfg = getattr(settings, "VOLC_TTS", {}) or {}
|
||||
self.appid = self.appid or cfg.get("appid")
|
||||
self.access_token = self.access_token or cfg.get("access_token")
|
||||
self.cluster = self.cluster or cfg.get("cluster") or "volcano_tts"
|
||||
self.base_url = self.base_url or cfg.get("base_url") or "https://openspeech.bytedance.com/api/v1/tts"
|
||||
|
||||
@property
|
||||
def configured(self) -> bool:
|
||||
return bool(self.appid and self.access_token)
|
||||
|
||||
def synthesize(self, *, text: str, voice_type: str, speed_ratio: float = 1.0, uid: str = "airshelf") -> tuple[bytes, int]:
|
||||
"""合成一段语音。返回 (mp3 字节, 时长毫秒;接口没回时长则为 0)。"""
|
||||
if not self.configured:
|
||||
raise TtsNotConfigured(
|
||||
"语音合成未配置:请在后端环境变量设置 VOLC_TTS_APPID 和 VOLC_TTS_ACCESS_TOKEN"
|
||||
"(火山引擎控制台 → 语音技术 → 语音合成大模型 → 创建应用)"
|
||||
)
|
||||
body = {
|
||||
"app": {"appid": self.appid, "token": self.access_token, "cluster": self.cluster},
|
||||
"user": {"uid": uid},
|
||||
"audio": {"voice_type": voice_type, "encoding": "mp3", "speed_ratio": float(speed_ratio or 1.0)},
|
||||
"request": {"reqid": str(uuid.uuid4()), "text": text, "operation": "query"},
|
||||
}
|
||||
response = requests.post(
|
||||
self.base_url,
|
||||
headers={"Authorization": f"Bearer;{self.access_token}"},
|
||||
json=body,
|
||||
timeout=60,
|
||||
)
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
if data.get("code") != 3000 or not data.get("data"):
|
||||
raise ValueError(f"语音合成失败:{data.get('message') or data.get('code')}")
|
||||
audio = base64.b64decode(data["data"])
|
||||
duration_ms = 0
|
||||
try:
|
||||
duration_ms = int(float((data.get("addition") or {}).get("duration") or 0))
|
||||
except (TypeError, ValueError):
|
||||
duration_ms = 0
|
||||
return audio, duration_ms
|
||||
|
||||
@@ -0,0 +1,43 @@
|
||||
from typing import Any
|
||||
|
||||
import requests
|
||||
from django.conf import settings
|
||||
|
||||
from .volcano import VolcanoArkProvider
|
||||
|
||||
|
||||
class YunqiProvider(VolcanoArkProvider):
|
||||
"""YunQi AI(New API 网关)适配器:OpenAI 兼容 images/generations。
|
||||
|
||||
与火山 ARK 的差异:不接受 watermark/sequential_image_generation/response_format 私有参数;
|
||||
size 用 OpenAI 枚举(1024x1024/1024x1536/1536x1024);gpt-image-2 固定返回 b64_json,
|
||||
下游 extract_first_media_url / media_to_bytes 已兼容 base64,无需改动。
|
||||
"""
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
self.api_key = self.api_key or settings.YUNQI.get("api_key")
|
||||
self.base_url = self.base_url or settings.YUNQI.get("base_url")
|
||||
|
||||
def image_generation(
|
||||
self,
|
||||
*,
|
||||
model: str,
|
||||
prompt: str,
|
||||
endpoint: str = "images/generations",
|
||||
image: str | list[str] | None = None,
|
||||
size: str = "1024x1536",
|
||||
) -> dict[str, Any]:
|
||||
if not self.api_key:
|
||||
raise ValueError("YUNQI_API_KEY is not configured")
|
||||
if image:
|
||||
raise ValueError("gpt-image-2 via YunQi only supports text-to-image; reference image is not supported")
|
||||
body: dict[str, Any] = {"model": model, "prompt": prompt, "size": size, "n": 1}
|
||||
# 实测该网关生图平均延迟 75s+,超时须显著高于火山的 180s
|
||||
response = requests.post(
|
||||
f"{self.base_url.rstrip('/')}/{endpoint.lstrip('/')}",
|
||||
headers={"Authorization": f"Bearer {self.api_key}", "Content-Type": "application/json"},
|
||||
json=body,
|
||||
timeout=300,
|
||||
)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
@@ -10,7 +10,7 @@ from django.db import transaction
|
||||
from django.utils import timezone
|
||||
|
||||
from apps.ai.models import AITask, ModelConfig
|
||||
from apps.ai.providers import VolcanoArkProvider
|
||||
from apps.ai.providers import TtsNotConfigured, VolcanoArkProvider, VolcanoTtsProvider, YunqiProvider
|
||||
from apps.assets.models import Asset, AssetFile
|
||||
from apps.assets.storage import TosStorage
|
||||
from apps.billing.services.ledger import charge_reserved_credit, release_credit, reserve_credit
|
||||
@@ -22,6 +22,7 @@ from apps.projects.models import (
|
||||
ScriptVersion,
|
||||
StoryboardFrame,
|
||||
StoryboardVersion,
|
||||
Timeline,
|
||||
VideoSegment,
|
||||
VideoSegmentVersion,
|
||||
)
|
||||
@@ -36,6 +37,13 @@ def get_default_model(capability: str) -> ModelConfig:
|
||||
)
|
||||
|
||||
|
||||
def get_image_provider(model_config: ModelConfig):
|
||||
"""生图按 provider 分流:yunqi 走 OpenAI 兼容网关(gpt-image-2),其余沿用火山 ARK。"""
|
||||
if model_config.provider.name == "yunqi":
|
||||
return YunqiProvider(base_url=model_config.provider.base_url or None)
|
||||
return VolcanoArkProvider(base_url=model_config.provider.base_url or None)
|
||||
|
||||
|
||||
def estimate_cost(model_config: ModelConfig) -> Decimal:
|
||||
return model_config.unit_price if model_config.unit_price > 0 else Decimal("1.0000")
|
||||
|
||||
@@ -48,7 +56,9 @@ def build_script_prompt(*, project, user_prompt: str, selling_point_ids: list[st
|
||||
selling_text = "\n".join(f"- {item.title}: {item.detail}" for item in selling_points)
|
||||
system = (
|
||||
"你是电商短视频脚本导演。请为 9:16 竖屏带货短视频生成 60 秒脚本,"
|
||||
"拆成 4 个 15 秒段落。每段包含旁白、画面描述、商品露出方式和转场建议。"
|
||||
"拆成 4 个 15 秒段落。严格按以下格式输出,段落之间空一行,不要输出其他内容:\n"
|
||||
"镜头1\n旁白:这一镜要念出来的口播文案(一两句话)\n画面:这一镜的画面描述、商品露出方式和转场建议\n\n"
|
||||
"镜头2\n旁白:…\n画面:…(依此类推到镜头4)"
|
||||
)
|
||||
user = f"""
|
||||
商品标题:{product.title}
|
||||
@@ -65,6 +75,42 @@ def build_script_prompt(*, project, user_prompt: str, selling_point_ids: list[st
|
||||
return [{"role": "system", "content": system}, {"role": "user", "content": user}]
|
||||
|
||||
|
||||
def parse_segment_fields(block: str) -> tuple[str, str]:
|
||||
"""从一镜文本里拆出(旁白, 画面)。
|
||||
|
||||
模型按 build_script_prompt 的格式输出「旁白:…/画面:…」标签行时精确拆分;
|
||||
自带脚本/旧格式没有标签则两个字段都用整段(保持旧行为),字幕/故事板各自兜底。
|
||||
"""
|
||||
narration_lines: list[str] = []
|
||||
visual_lines: list[str] = []
|
||||
current: list[str] | None = None
|
||||
for raw in (block or "").splitlines():
|
||||
line = raw.strip()
|
||||
if not line:
|
||||
continue
|
||||
matched = re.match(r"^(旁白|口播|台词|文案)\s*[::]\s*(.*)$", line)
|
||||
if matched:
|
||||
current = narration_lines
|
||||
if matched.group(2):
|
||||
current.append(matched.group(2))
|
||||
continue
|
||||
matched = re.match(r"^(画面|镜头描述|视觉|画面描述)\s*[::]\s*(.*)$", line)
|
||||
if matched:
|
||||
current = visual_lines
|
||||
if matched.group(2):
|
||||
current.append(matched.group(2))
|
||||
continue
|
||||
if re.match(r"^(镜头|分镜|场)\s*\d+", line):
|
||||
continue # 「镜头N」标题行不计入任何字段
|
||||
if current is not None:
|
||||
current.append(line)
|
||||
narration = " ".join(narration_lines).strip()
|
||||
visual = " ".join(visual_lines).strip()
|
||||
if not narration and not visual:
|
||||
return block.strip(), block.strip()
|
||||
return narration or visual, visual or narration
|
||||
|
||||
|
||||
def split_script_into_segments(content: str, count: int = 4) -> list[str]:
|
||||
"""把一段脚本稳健地拆成 `count` 个分镜文本,保证每镜都非空、且所有内容都被分配到某一镜。
|
||||
|
||||
@@ -124,7 +170,7 @@ def create_ai_task(*, project, user, task_type: str, model_config: ModelConfig,
|
||||
return task
|
||||
|
||||
|
||||
def generate_project_script(*, project, user, user_prompt: str, selling_point_ids: list[str] | None = None) -> ScriptVersion:
|
||||
def generate_project_script(*, project, user, user_prompt: str, selling_point_ids: list[str] | None = None, source: str = "ai") -> ScriptVersion:
|
||||
model_config = get_default_model(ModelConfig.Capability.TEXT)
|
||||
if model_config is None:
|
||||
raise ValueError("no active text model configured")
|
||||
@@ -162,16 +208,17 @@ def generate_project_script(*, project, user, user_prompt: str, selling_point_id
|
||||
task=task,
|
||||
title="AI 脚本",
|
||||
content=content,
|
||||
source="ai",
|
||||
source=source if source in ("ai", "theme", "manual") else "ai",
|
||||
is_adopted=False,
|
||||
)
|
||||
for index, segment_text in enumerate(split_script_into_segments(content)):
|
||||
narration, visual = parse_segment_fields(segment_text)
|
||||
ScriptSegment.objects.create(
|
||||
script_version=script,
|
||||
sort_order=index,
|
||||
duration_seconds=15,
|
||||
narration=segment_text,
|
||||
visual_prompt=segment_text,
|
||||
narration=narration,
|
||||
visual_prompt=visual,
|
||||
)
|
||||
|
||||
stage, _ = ProjectStage.objects.get_or_create(project=project, stage=ProjectStage.Stage.SCRIPT)
|
||||
@@ -283,7 +330,7 @@ def generate_base_asset(*, project, user, kind: str, prompt: str) -> BaseAssetGr
|
||||
)
|
||||
reservation = task.credit_reservation
|
||||
try:
|
||||
provider = VolcanoArkProvider(base_url=model_config.provider.base_url or None)
|
||||
provider = get_image_provider(model_config)
|
||||
response = provider.image_generation(model=model_config.name, endpoint=model_config.endpoint, prompt=prompt)
|
||||
media = provider.extract_first_media_url(response)
|
||||
with transaction.atomic():
|
||||
@@ -408,7 +455,7 @@ def _storyboard_frame_worker(task_id, version_id, segment_id, user_id) -> None:
|
||||
task.status = AITask.Status.SUBMITTED
|
||||
task.save(update_fields=["status", "updated_at"])
|
||||
try:
|
||||
provider = VolcanoArkProvider(base_url=model_config.provider.base_url or None)
|
||||
provider = get_image_provider(model_config)
|
||||
frame_prompt = task.request_payload.get("prompt") or build_storyboard_frame_prompt(project, version, segment)
|
||||
response = provider.image_generation(
|
||||
model=model_config.name,
|
||||
@@ -416,12 +463,9 @@ def _storyboard_frame_worker(task_id, version_id, segment_id, user_id) -> None:
|
||||
prompt=frame_prompt,
|
||||
)
|
||||
media = provider.extract_first_media_url(response)
|
||||
task.status = AITask.Status.SUCCEEDED
|
||||
task.response_payload = response
|
||||
task.actual_cost = task.estimated_cost
|
||||
task.completed_at = timezone.now()
|
||||
task.save(update_fields=["status", "response_payload", "actual_cost", "completed_at", "updated_at"])
|
||||
charge_reserved_credit(reservation=reservation, actual_amount=task.actual_cost)
|
||||
# 注意顺序:task 是 poll 端的「占位锁」,必须等帧真正落库后才置 SUCCEEDED。
|
||||
# 旧实现先置 SUCCEEDED 再上传 TOS(数秒)最后建帧,中间窗口 poll 会判「无在途且帧缺失」
|
||||
# 为同一镜重复起线程 → 重复帧 + 重复扣费(实测 4 帧出 6 帧)。
|
||||
asset = _store_generated_media(
|
||||
team=project.team,
|
||||
user=user,
|
||||
@@ -432,13 +476,22 @@ def _storyboard_frame_worker(task_id, version_id, segment_id, user_id) -> None:
|
||||
category=Asset.Category.SCENE,
|
||||
asset_type=Asset.Type.IMAGE,
|
||||
)
|
||||
StoryboardFrame.objects.create(
|
||||
storyboard=version,
|
||||
script_segment=segment,
|
||||
asset=asset,
|
||||
sort_order=segment.sort_order,
|
||||
prompt=segment.visual_prompt,
|
||||
)
|
||||
with transaction.atomic():
|
||||
task.status = AITask.Status.SUCCEEDED
|
||||
task.response_payload = response
|
||||
task.actual_cost = task.estimated_cost
|
||||
task.completed_at = timezone.now()
|
||||
task.save(update_fields=["status", "response_payload", "actual_cost", "completed_at", "updated_at"])
|
||||
charge_reserved_credit(reservation=reservation, actual_amount=task.actual_cost)
|
||||
# 幂等守卫:该镜已有帧(任何残余竞态/双 poll)就不再建,保持一镜一帧
|
||||
if not StoryboardFrame.objects.filter(storyboard=version, script_segment=segment).exists():
|
||||
StoryboardFrame.objects.create(
|
||||
storyboard=version,
|
||||
script_segment=segment,
|
||||
asset=asset,
|
||||
sort_order=segment.sort_order,
|
||||
prompt=segment.visual_prompt,
|
||||
)
|
||||
except Exception as exc: # noqa: BLE001 — 失败回滚额度,标记任务失败供 poll 上报
|
||||
task.status = AITask.Status.FAILED
|
||||
task.error_message = str(exc)
|
||||
@@ -470,53 +523,64 @@ def generate_storyboard_frame(*, project, user) -> dict:
|
||||
_finalize_storyboard(project, version)
|
||||
return {"status": "succeeded", "done": total, "total": total, "version_id": str(version.id)}
|
||||
|
||||
# 该版本内是否已有帧在后台生成中(RESERVED/SUBMITTED 的故事板任务即为「占位锁」)。
|
||||
# 仅算「近 3 分钟内」的任务:若进程/线程意外中断留下僵尸任务,超时后不再视为在生成,允许重新发起。
|
||||
# 每镜独立「占位锁」(该镜有 CREATED/RESERVED/SUBMITTED 的任务=在生成中)。
|
||||
# 旧实现是版本级单锁、一次只生成一帧;改为缺哪几镜就同时起哪几镜的线程。
|
||||
# ★ 实测注记(2026-06-10):4 线程并发时当前 Seedream 端点在服务侧排队,单帧 25s→83-105s,
|
||||
# 整版总时长 ≈ 串行(117s)。瓶颈是 ARK 端点并发配额而非本机;并行无额外成本,
|
||||
# 配额提升后自动受益。可用 settings.STORYBOARD_MAX_PARALLEL 调并发(1=回到串行)。
|
||||
# 仅算「近 3 分钟内」的任务:线程意外中断留下的僵尸任务超时后不再占锁,允许重新发起。
|
||||
from django.conf import settings as dj_settings
|
||||
STORYBOARD_MAX_PARALLEL = int(getattr(dj_settings, "STORYBOARD_MAX_PARALLEL", 4))
|
||||
stale_cutoff = timezone.now() - timedelta(minutes=3)
|
||||
inflight = AITask.objects.filter(
|
||||
project=project,
|
||||
task_type=AITask.Type.STORYBOARD,
|
||||
status__in=[AITask.Status.CREATED, AITask.Status.RESERVED, AITask.Status.SUBMITTED],
|
||||
request_payload__storyboard_version=str(version.id),
|
||||
created_at__gte=stale_cutoff,
|
||||
).exists()
|
||||
if inflight:
|
||||
return {"status": "generating", "done": done, "total": total, "version_id": str(version.id)}
|
||||
inflight_segment_ids = {
|
||||
str(v)
|
||||
for v in AITask.objects.filter(
|
||||
project=project,
|
||||
task_type=AITask.Type.STORYBOARD,
|
||||
status__in=[AITask.Status.CREATED, AITask.Status.RESERVED, AITask.Status.SUBMITTED],
|
||||
request_payload__storyboard_version=str(version.id),
|
||||
created_at__gte=stale_cutoff,
|
||||
).values_list("request_payload__storyboard_segment", flat=True)
|
||||
if v
|
||||
}
|
||||
|
||||
pending = [s for s in segments if s.id not in done_segment_ids]
|
||||
segment = pending[0]
|
||||
# 单帧失败次数上限,避免持续失败时无限重试
|
||||
failed_for_segment = AITask.objects.filter(
|
||||
project=project,
|
||||
task_type=AITask.Type.STORYBOARD,
|
||||
status=AITask.Status.FAILED,
|
||||
request_payload__storyboard_segment=str(segment.id),
|
||||
).count()
|
||||
if failed_for_segment >= 2:
|
||||
last = AITask.objects.filter(project=project, task_type=AITask.Type.STORYBOARD, status=AITask.Status.FAILED,
|
||||
request_payload__storyboard_segment=str(segment.id)).order_by("-created_at").first()
|
||||
return {"status": "failed", "done": done, "total": total, "version_id": str(version.id),
|
||||
"error": last.error_message if last else "storyboard frame failed"}
|
||||
# 单帧失败次数上限,避免持续失败时无限重试;任一镜到上限即整版上报失败
|
||||
for segment in pending:
|
||||
failed_for_segment = AITask.objects.filter(
|
||||
project=project,
|
||||
task_type=AITask.Type.STORYBOARD,
|
||||
status=AITask.Status.FAILED,
|
||||
request_payload__storyboard_segment=str(segment.id),
|
||||
).count()
|
||||
if failed_for_segment >= 2:
|
||||
last = AITask.objects.filter(project=project, task_type=AITask.Type.STORYBOARD, status=AITask.Status.FAILED,
|
||||
request_payload__storyboard_segment=str(segment.id)).order_by("-created_at").first()
|
||||
return {"status": "failed", "done": done, "total": total, "version_id": str(version.id),
|
||||
"error": last.error_message if last else "storyboard frame failed"}
|
||||
|
||||
spawnable = [s for s in pending if str(s.id) not in inflight_segment_ids]
|
||||
slots = max(0, STORYBOARD_MAX_PARALLEL - len(inflight_segment_ids))
|
||||
model_config = get_default_model(ModelConfig.Capability.IMAGE)
|
||||
task = create_ai_task(
|
||||
project=project,
|
||||
user=user,
|
||||
task_type=AITask.Type.STORYBOARD,
|
||||
model_config=model_config,
|
||||
request_payload={
|
||||
"model": model_config.name,
|
||||
"endpoint": model_config.endpoint,
|
||||
"prompt": build_storyboard_frame_prompt(project, version, segment),
|
||||
"storyboard_version": str(version.id),
|
||||
"storyboard_segment": str(segment.id),
|
||||
},
|
||||
)
|
||||
threading.Thread(
|
||||
target=_storyboard_frame_worker,
|
||||
args=(str(task.id), str(version.id), str(segment.id), str(user.id)),
|
||||
daemon=True,
|
||||
).start()
|
||||
for segment in spawnable[:slots]:
|
||||
task = create_ai_task(
|
||||
project=project,
|
||||
user=user,
|
||||
task_type=AITask.Type.STORYBOARD,
|
||||
model_config=model_config,
|
||||
request_payload={
|
||||
"model": model_config.name,
|
||||
"endpoint": model_config.endpoint,
|
||||
"prompt": build_storyboard_frame_prompt(project, version, segment),
|
||||
"storyboard_version": str(version.id),
|
||||
"storyboard_segment": str(segment.id),
|
||||
},
|
||||
)
|
||||
threading.Thread(
|
||||
target=_storyboard_frame_worker,
|
||||
args=(str(task.id), str(version.id), str(segment.id), str(user.id)),
|
||||
daemon=True,
|
||||
).start()
|
||||
return {"status": "generating", "done": done, "total": total, "version_id": str(version.id)}
|
||||
|
||||
|
||||
@@ -656,16 +720,17 @@ def poll_video_segment(*, video_segment: VideoSegment, user) -> VideoSegmentVers
|
||||
if video_segment.status == VideoSegment.Status.FAILED:
|
||||
return None
|
||||
|
||||
task = video_segment.versions.order_by("-created_at").first()
|
||||
ai_task = None
|
||||
if task:
|
||||
ai_task = task.task
|
||||
# ★ 先找「在途任务」再回退旧版本的任务。旧实现反过来:重跑时段上已有(旧)版本,
|
||||
# 取到旧版本挂的已成功任务 → 短路返回旧版,在途的新任务永远没人轮询,
|
||||
# 段永远卡「生成中」、新视频取不回来(实测重跑卡 40 分钟,ARK 侧其实早已生成完)。
|
||||
ai_task = video_segment.project.ai_tasks.filter(
|
||||
task_type=AITask.Type.VIDEO_SEGMENT,
|
||||
request_payload__video_segment_id=str(video_segment.id),
|
||||
status__in=[AITask.Status.SUBMITTED, AITask.Status.POLLING],
|
||||
).order_by("-created_at").first()
|
||||
if ai_task is None:
|
||||
ai_task = video_segment.project.ai_tasks.filter(
|
||||
task_type=AITask.Type.VIDEO_SEGMENT,
|
||||
request_payload__video_segment_id=str(video_segment.id),
|
||||
status__in=[AITask.Status.SUBMITTED, AITask.Status.POLLING],
|
||||
).order_by("-created_at").first()
|
||||
latest_version = video_segment.versions.order_by("-created_at").first()
|
||||
ai_task = latest_version.task if latest_version else None
|
||||
if ai_task is None:
|
||||
raise ValueError("no active video generation task")
|
||||
|
||||
@@ -679,9 +744,11 @@ def poll_video_segment(*, video_segment: VideoSegment, user) -> VideoSegmentVers
|
||||
response = provider.poll_video_task(endpoint=ai_task.model_config.endpoint, provider_task_id=ai_task.provider_task_id)
|
||||
remote_status = response.get("status")
|
||||
if remote_status in {"queued", "running", "processing"}:
|
||||
ai_task.status = AITask.Status.POLLING
|
||||
ai_task.response_payload = response
|
||||
ai_task.save(update_fields=["status", "response_payload", "updated_at"])
|
||||
# 仍在生成:只在状态首次进入 POLLING 时落一次库。旧实现每次 poll(5s 一次)都把完整
|
||||
# response JSON 回写远程 MySQL——纯浪费写带宽,终态时反正会存完整 payload。
|
||||
if ai_task.status != AITask.Status.POLLING:
|
||||
ai_task.status = AITask.Status.POLLING
|
||||
ai_task.save(update_fields=["status", "updated_at"])
|
||||
return None
|
||||
if remote_status in {"failed", "expired", "cancelled"}:
|
||||
ai_task.status = AITask.Status.FAILED
|
||||
@@ -706,23 +773,33 @@ def poll_video_segment(*, video_segment: VideoSegment, user) -> VideoSegmentVers
|
||||
category=Asset.Category.VIDEO_CLIP,
|
||||
asset_type=Asset.Type.VIDEO,
|
||||
)
|
||||
ai_task.status = AITask.Status.SUCCEEDED
|
||||
ai_task.response_payload = response
|
||||
ai_task.actual_cost = ai_task.estimated_cost
|
||||
ai_task.completed_at = timezone.now()
|
||||
ai_task.save(update_fields=["status", "response_payload", "actual_cost", "completed_at", "updated_at"])
|
||||
charge_reserved_credit(reservation=ai_task.credit_reservation, actual_amount=ai_task.actual_cost)
|
||||
version = VideoSegmentVersion.objects.create(
|
||||
video_segment=video_segment,
|
||||
task=ai_task,
|
||||
asset=asset,
|
||||
prompt=ai_task.request_payload.get("prompt", ""),
|
||||
is_adopted=True,
|
||||
)
|
||||
video_segment.adopted_version = version
|
||||
video_segment.status = VideoSegment.Status.SUCCEEDED
|
||||
video_segment.error_message = ""
|
||||
video_segment.save(update_fields=["adopted_version", "status", "error_message", "updated_at"])
|
||||
# 终态化必须持锁原子做:两个并发 poll(前端 5s 静默轮询 × 提交后轮询/worker)同时走到这里时,
|
||||
# 旧实现会同 task 建两个版本 + charge_reserved_credit 双扣费(实测 03:08:25 同秒双版本)。
|
||||
# select_for_update 锁 task 行,后到者看到 SUCCEEDED 直接回已有版,不再建版/扣费。
|
||||
with transaction.atomic():
|
||||
locked_task = AITask.objects.select_for_update().get(id=ai_task.id)
|
||||
if locked_task.status == AITask.Status.SUCCEEDED:
|
||||
existing = video_segment.versions.filter(task=locked_task).order_by("-created_at").first()
|
||||
if existing is not None:
|
||||
return existing
|
||||
locked_task.status = AITask.Status.SUCCEEDED
|
||||
locked_task.response_payload = response
|
||||
locked_task.actual_cost = locked_task.estimated_cost
|
||||
locked_task.completed_at = timezone.now()
|
||||
locked_task.save(update_fields=["status", "response_payload", "actual_cost", "completed_at", "updated_at"])
|
||||
charge_reserved_credit(reservation=locked_task.credit_reservation, actual_amount=locked_task.actual_cost)
|
||||
version = VideoSegmentVersion.objects.create(
|
||||
video_segment=video_segment,
|
||||
task=locked_task,
|
||||
asset=asset,
|
||||
prompt=locked_task.request_payload.get("prompt", ""),
|
||||
is_adopted=True,
|
||||
)
|
||||
video_segment.versions.exclude(id=version.id).update(is_adopted=False)
|
||||
video_segment.adopted_version = version
|
||||
video_segment.status = VideoSegment.Status.SUCCEEDED
|
||||
video_segment.error_message = ""
|
||||
video_segment.save(update_fields=["adopted_version", "status", "error_message", "updated_at"])
|
||||
return version
|
||||
|
||||
|
||||
@@ -750,7 +827,7 @@ def generate_standalone_image(*, team, user, prompt: str, mode: str = "image", c
|
||||
category = _STANDALONE_CATEGORY.get(mode, Asset.Category.UNCATEGORIZED)
|
||||
task_type = _STANDALONE_TASK_TYPE.get(mode, AITask.Type.PRODUCT_IMAGE)
|
||||
count = max(1, min(int(count or 1), 4))
|
||||
provider = VolcanoArkProvider(base_url=model_config.provider.base_url or None)
|
||||
provider = get_image_provider(model_config)
|
||||
assets: list[Asset] = []
|
||||
for index in range(count):
|
||||
cost = estimate_cost(model_config)
|
||||
@@ -798,3 +875,112 @@ def generate_standalone_image(*, team, user, prompt: str, mode: str = "image", c
|
||||
release_credit(reservation=reservation, reason=str(exc))
|
||||
raise
|
||||
return assets
|
||||
|
||||
|
||||
# ── 旁白配音(TTS):每镜旁白合成一段语音,导出时作为人声轨混在 BGM 之上 ──
|
||||
|
||||
# 音色按「语音合成(经典版)」试用包实测可用清单配置;大模型音色(*_bigtts)需另开通「语音合成大模型」服务,当前账号 403
|
||||
VOICEOVER_VOICES = [
|
||||
{"key": "BV700_streaming", "label": "灿灿 · 活力女声"},
|
||||
{"key": "BV034_streaming", "label": "知性姐姐 · 沉稳女声"},
|
||||
{"key": "BV001_streaming", "label": "通用女声"},
|
||||
{"key": "BV056_streaming", "label": "阳光男声"},
|
||||
{"key": "BV102_streaming", "label": "儒雅青年 · 解说男声"},
|
||||
{"key": "BV002_streaming", "label": "通用男声"},
|
||||
]
|
||||
DEFAULT_VOICEOVER_VOICE = VOICEOVER_VOICES[0]["key"]
|
||||
|
||||
|
||||
def synthesize_project_voiceover(*, project, user, items: list[dict], voice_type: str, speed_ratio: float = 1.0) -> dict:
|
||||
"""每镜旁白 → **逐句** TTS 配音资产(一句一段音频,带句内起点 offset_ms),映射写入
|
||||
timeline.metadata["voiceover"]。逐句才能支持「拖动字幕块 = 字幕和它的语音一起移动」;
|
||||
一次调用 = 一个 AITask = 计一次费;任何一句失败则整体失败并释放预留(不留半套配音)。"""
|
||||
from apps.projects.services.export import _split_subtitle_text
|
||||
|
||||
texts = [] # (片段 index, 句序 cue, 句文本)
|
||||
for n, item in enumerate(items or []):
|
||||
text = str(item.get("text") or "").strip()
|
||||
if not text:
|
||||
continue
|
||||
idx = int(item.get("index", n))
|
||||
pieces = _split_subtitle_text(text) or [text]
|
||||
for j, piece in enumerate(pieces):
|
||||
texts.append((idx, j, piece))
|
||||
if not texts:
|
||||
raise ValueError("没有可配音的旁白文本")
|
||||
provider = VolcanoTtsProvider()
|
||||
if not provider.configured:
|
||||
raise TtsNotConfigured(
|
||||
"语音合成未配置:请在后端环境变量设置 VOLC_TTS_APPID 和 VOLC_TTS_ACCESS_TOKEN"
|
||||
"(火山引擎控制台 → 语音技术 → 语音合成大模型 → 创建应用)"
|
||||
)
|
||||
voice_type = voice_type or DEFAULT_VOICEOVER_VOICE
|
||||
model_config = get_default_model(ModelConfig.Capability.AUDIO)
|
||||
if model_config is None:
|
||||
raise ValueError("no active audio model configured")
|
||||
task = create_ai_task(
|
||||
project=project,
|
||||
user=user,
|
||||
task_type=AITask.Type.VOICEOVER,
|
||||
model_config=model_config,
|
||||
request_payload={
|
||||
"voice_type": voice_type,
|
||||
"speed_ratio": float(speed_ratio or 1.0),
|
||||
"items": [{"index": idx, "cue": j, "text": text} for idx, j, text in texts],
|
||||
},
|
||||
)
|
||||
reservation = task.credit_reservation
|
||||
try:
|
||||
synthesized = []
|
||||
for idx, j, text in texts:
|
||||
audio, duration_ms = provider.synthesize(text=text, voice_type=voice_type, speed_ratio=speed_ratio, uid=str(user.id))
|
||||
synthesized.append((idx, j, text, audio, duration_ms))
|
||||
with transaction.atomic():
|
||||
task.status = AITask.Status.SUCCEEDED
|
||||
task.actual_cost = task.estimated_cost
|
||||
task.completed_at = timezone.now()
|
||||
task.response_payload = {"segments": len(synthesized)}
|
||||
task.save(update_fields=["status", "actual_cost", "completed_at", "response_payload", "updated_at"])
|
||||
charge_reserved_credit(reservation=reservation, actual_amount=task.actual_cost)
|
||||
vo_items = []
|
||||
offset_acc: dict[int, int] = {} # 同一片段内逐句顺排:句 j 的默认起点 = 前面句时长之和
|
||||
for idx, j, text, audio, duration_ms in synthesized:
|
||||
asset_id = uuid.uuid4()
|
||||
object_key = f"teams/{project.team_id}/projects/{project.id}/voiceover/{asset_id}.mp3"
|
||||
stored = TosStorage().upload_fileobj(fileobj=BytesIO(audio), object_key=object_key, content_type="audio/mpeg")
|
||||
asset = Asset.objects.create(
|
||||
id=asset_id, team=project.team, created_by=user,
|
||||
name=f"配音 · 场 {idx + 1} · 句 {j + 1}", asset_type=Asset.Type.AUDIO,
|
||||
source=Asset.Source.AI_GENERATED, category=Asset.Category.UNCATEGORIZED,
|
||||
origin_task=task, description=text,
|
||||
)
|
||||
AssetFile.objects.create(
|
||||
asset=asset, object_key=stored.object_key, bucket=stored.bucket,
|
||||
content_type=stored.content_type, size_bytes=stored.size_bytes, is_primary=True,
|
||||
)
|
||||
offset_ms = offset_acc.get(idx, 0)
|
||||
offset_acc[idx] = offset_ms + (duration_ms or 0)
|
||||
vo_items.append({
|
||||
"index": idx, "cue": j, "text": text, "asset": str(asset.id),
|
||||
"duration_ms": duration_ms, "offset_ms": offset_ms,
|
||||
})
|
||||
timeline, _ = Timeline.objects.get_or_create(
|
||||
project=project, defaults={"name": f"{project.name} Timeline", "duration_seconds": 60}
|
||||
)
|
||||
metadata = dict(timeline.metadata or {})
|
||||
metadata["voiceover"] = {
|
||||
"enabled": True,
|
||||
"voice_type": voice_type,
|
||||
"speed_ratio": float(speed_ratio or 1.0),
|
||||
"items": vo_items,
|
||||
}
|
||||
timeline.metadata = metadata
|
||||
timeline.save(update_fields=["metadata", "updated_at"])
|
||||
return metadata["voiceover"]
|
||||
except Exception as exc:
|
||||
task.status = AITask.Status.FAILED
|
||||
task.error_message = str(exc)
|
||||
task.completed_at = timezone.now()
|
||||
task.save(update_fields=["status", "error_message", "completed_at", "updated_at"])
|
||||
release_credit(reservation=reservation, reason=str(exc))
|
||||
raise
|
||||
|
||||
@@ -5,6 +5,7 @@ from rest_framework.viewsets import ReadOnlyModelViewSet
|
||||
|
||||
from apps.assets.serializers import AssetSerializer
|
||||
from apps.common.api import TeamScopedViewSetMixin, get_current_team
|
||||
from apps.common.celery_health import require_worker
|
||||
|
||||
from .models import AITask, ModelConfig
|
||||
from .serializers import AITaskSerializer, ModelConfigSerializer
|
||||
@@ -15,6 +16,7 @@ class GenerateImageView(APIView):
|
||||
"""POST /api/ai/generate-image/ — 独立生图(不绑项目)· 图片创作/模特图/平台套图共用。"""
|
||||
|
||||
def post(self, request):
|
||||
require_worker()
|
||||
prompt = str(request.data.get("prompt") or "").strip()
|
||||
if not prompt:
|
||||
return Response({"detail": "prompt 不能为空"}, status=status.HTTP_400_BAD_REQUEST)
|
||||
|
||||
@@ -1,8 +1,11 @@
|
||||
from pathlib import Path
|
||||
import uuid
|
||||
|
||||
import requests
|
||||
from django.db import transaction
|
||||
from django.http import StreamingHttpResponse
|
||||
from rest_framework import status
|
||||
from rest_framework.decorators import action
|
||||
from rest_framework.parsers import FormParser, MultiPartParser
|
||||
from rest_framework.response import Response
|
||||
from rest_framework.views import APIView
|
||||
@@ -21,6 +24,29 @@ class AssetViewSet(TeamScopedViewSetMixin, ModelViewSet):
|
||||
search_fields = ["name", "description"]
|
||||
ordering_fields = ["created_at", "updated_at", "name"]
|
||||
|
||||
@action(detail=True, methods=["get"], url_path="raw")
|
||||
def raw(self, request, pk=None):
|
||||
"""同源流式代理资产主文件。TOS 桶未配 CORS,浏览器 JS 读不到跨域媒体数据——
|
||||
编辑器抽视频缩略图(canvas 会被 taint)/解码音频波形(fetch 被拦)都需要走这里。"""
|
||||
asset = self.get_object()
|
||||
primary = asset.files.filter(is_primary=True).first() or asset.files.first()
|
||||
if primary is None:
|
||||
return Response({"detail": "asset has no file"}, status=status.HTTP_404_NOT_FOUND)
|
||||
url = TosStorage().presigned_get_url(object_key=primary.object_key, expires_in=600)
|
||||
upstream_headers = {}
|
||||
if request.META.get("HTTP_RANGE"):
|
||||
upstream_headers["Range"] = request.META["HTTP_RANGE"]
|
||||
upstream = requests.get(url, headers=upstream_headers, stream=True, timeout=120)
|
||||
response = StreamingHttpResponse(
|
||||
upstream.iter_content(chunk_size=256 * 1024),
|
||||
status=upstream.status_code,
|
||||
content_type=primary.content_type or "application/octet-stream",
|
||||
)
|
||||
for header in ("Content-Length", "Content-Range", "Accept-Ranges"):
|
||||
if header in upstream.headers:
|
||||
response[header] = upstream.headers[header]
|
||||
return response
|
||||
|
||||
|
||||
class AssetUploadView(APIView):
|
||||
parser_classes = [MultiPartParser, FormParser]
|
||||
|
||||
@@ -0,0 +1,63 @@
|
||||
"""释放僵尸预留额度。
|
||||
|
||||
生成流程中断(进程被杀/浏览器关闭/轮询停止)会留下 status=ACTIVE 的 CreditReservation,
|
||||
对应 AITask 永远停在 reserved/submitted/polling——这些预留会一直占用 reserved_balance,
|
||||
吃掉团队可用额度。本命令把「超过 --hours 小时仍未终态」的预留按正规 release 流程释放
|
||||
(走 release_credit,有 RELEASE 流水可审计),并把对应任务标记为 FAILED。
|
||||
|
||||
用法:
|
||||
python manage.py sweep_stale_reservations # 默认释放超过 12 小时的
|
||||
python manage.py sweep_stale_reservations --hours 1 # 收紧窗口
|
||||
python manage.py sweep_stale_reservations --dry-run # 只看不动
|
||||
"""
|
||||
from datetime import timedelta
|
||||
|
||||
from django.core.management.base import BaseCommand
|
||||
from django.utils import timezone
|
||||
|
||||
from apps.ai.models import AITask
|
||||
from apps.billing.models import CreditReservation
|
||||
from apps.billing.services.ledger import release_credit
|
||||
|
||||
# 这些任务状态意味着流程没走到终态(成功会 charge、失败会 release),超时即视为僵尸
|
||||
STUCK_TASK_STATUSES = (
|
||||
AITask.Status.CREATED,
|
||||
AITask.Status.RESERVED,
|
||||
AITask.Status.SUBMITTED,
|
||||
AITask.Status.POLLING,
|
||||
)
|
||||
|
||||
|
||||
class Command(BaseCommand):
|
||||
help = "释放超时未终态的 ACTIVE 预留额度(僵尸预留),并把对应任务标记为失败"
|
||||
|
||||
def add_arguments(self, parser):
|
||||
parser.add_argument("--hours", type=int, default=12, help="超过多少小时视为僵尸(默认 12)")
|
||||
parser.add_argument("--dry-run", action="store_true", help="只列出将要释放的,不实际执行")
|
||||
|
||||
def handle(self, *args, **options):
|
||||
cutoff = timezone.now() - timedelta(hours=options["hours"])
|
||||
stale = (
|
||||
CreditReservation.objects.filter(status=CreditReservation.Status.ACTIVE, created_at__lt=cutoff)
|
||||
.select_related("task")
|
||||
.order_by("created_at")
|
||||
)
|
||||
released = 0
|
||||
for reservation in stale:
|
||||
task = reservation.task
|
||||
if task is not None and task.status not in STUCK_TASK_STATUSES:
|
||||
# 任务已终态却仍 ACTIVE:同样是泄漏,照常释放
|
||||
pass
|
||||
label = f"resv {str(reservation.id)[:8]} ¥{reservation.amount} task={str(task.id)[:8] if task else '-'}({task.status if task else '-'}) created={reservation.created_at:%m-%d %H:%M}"
|
||||
if options["dry_run"]:
|
||||
self.stdout.write(f"[dry-run] would release {label}")
|
||||
continue
|
||||
release_credit(reservation=reservation, reason=f"僵尸预留清扫(>{options['hours']}h 未终态)")
|
||||
if task is not None and task.status in STUCK_TASK_STATUSES:
|
||||
task.status = AITask.Status.FAILED
|
||||
task.error_message = "stale task swept: 流程中断超时,预留已释放"
|
||||
task.completed_at = timezone.now()
|
||||
task.save(update_fields=["status", "error_message", "completed_at", "updated_at"])
|
||||
released += 1
|
||||
self.stdout.write(f"released {label}")
|
||||
self.stdout.write(self.style.SUCCESS(f"done: released {released} stale reservation(s)"))
|
||||
@@ -1,16 +1,46 @@
|
||||
from decimal import Decimal
|
||||
|
||||
from django.db import transaction
|
||||
from django.db.models import Sum
|
||||
from django.utils import timezone
|
||||
|
||||
from apps.billing.models import CreditAccount, CreditLedger, CreditReservation
|
||||
|
||||
|
||||
def _enforce_member_monthly_limit(*, team, user, amount: Decimal) -> None:
|
||||
"""成员月度限额管控。团队页可设 monthly_credit_limit(0=不限),此前只存不管——纯装饰。
|
||||
口径与成员列表的 month_charged 一致:自然月内该成员的 CHARGE 流水合计;
|
||||
再加上其当前 ACTIVE 预留(在途任务),防止并发把限额冲穿。
|
||||
调用方(reserve_credit)已持有 account 行锁,团队内 reserve 串行化,此检查无竞态。"""
|
||||
if user is None:
|
||||
return
|
||||
member = team.members.filter(user=user).first()
|
||||
if member is None:
|
||||
return
|
||||
limit = member.monthly_credit_limit or Decimal("0")
|
||||
if limit <= 0:
|
||||
return # 0 = 不限额
|
||||
now = timezone.now()
|
||||
month_start = now.replace(day=1, hour=0, minute=0, second=0, microsecond=0)
|
||||
charged = (
|
||||
CreditLedger.objects.filter(team=team, user=user, ledger_type=CreditLedger.Type.CHARGE, created_at__gte=month_start)
|
||||
.aggregate(s=Sum("amount"))["s"] or Decimal("0")
|
||||
)
|
||||
reserved = (
|
||||
CreditReservation.objects.filter(team=team, user=user, status=CreditReservation.Status.ACTIVE)
|
||||
.aggregate(s=Sum("amount"))["s"] or Decimal("0")
|
||||
)
|
||||
if charged + reserved + amount > limit:
|
||||
raise ValueError(f"成员本月额度不足:限额 ¥{limit},本月已用 ¥{charged}(另有在途 ¥{reserved})")
|
||||
|
||||
|
||||
@transaction.atomic
|
||||
def reserve_credit(*, team, user, task, amount: Decimal) -> CreditReservation:
|
||||
account, _ = CreditAccount.objects.select_for_update().get_or_create(team=team)
|
||||
available = account.balance - account.reserved_balance
|
||||
if available < amount:
|
||||
raise ValueError("insufficient credit")
|
||||
_enforce_member_monthly_limit(team=team, user=user, amount=amount)
|
||||
|
||||
account.reserved_balance += amount
|
||||
account.save(update_fields=["reserved_balance", "updated_at"])
|
||||
@@ -37,6 +67,9 @@ def reserve_credit(*, team, user, task, amount: Decimal) -> CreditReservation:
|
||||
@transaction.atomic
|
||||
def release_credit(*, reservation: CreditReservation, reason: str = "") -> None:
|
||||
account = CreditAccount.objects.select_for_update().get(team=reservation.team)
|
||||
# 锁内重读:与 charge 同款竞态——并发双 release(或 charge+release 交错)用陈旧快照都看到 ACTIVE,
|
||||
# reserved_balance 会被减两次直接记坏账。已终态(已扣/已释放)→ 幂等返回。
|
||||
reservation = CreditReservation.objects.select_for_update().get(id=reservation.id)
|
||||
if reservation.status != CreditReservation.Status.ACTIVE:
|
||||
return
|
||||
|
||||
@@ -59,6 +92,11 @@ def release_credit(*, reservation: CreditReservation, reason: str = "") -> None:
|
||||
@transaction.atomic
|
||||
def charge_reserved_credit(*, reservation: CreditReservation, actual_amount: Decimal) -> None:
|
||||
account = CreditAccount.objects.select_for_update().get(team=reservation.team)
|
||||
# 锁内重读 reservation:调用方传入的可能是并发竞态下的陈旧快照(双方都看到 ACTIVE),
|
||||
# 旧实现会对同一笔预留扣两次费(实测同秒两条 charge 流水)。已扣费 → 幂等直接返回。
|
||||
reservation = CreditReservation.objects.select_for_update().get(id=reservation.id)
|
||||
if reservation.status == CreditReservation.Status.CHARGED:
|
||||
return
|
||||
if reservation.status != CreditReservation.Status.ACTIVE:
|
||||
raise ValueError("reservation is not active")
|
||||
if actual_amount > reservation.amount:
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
from decimal import Decimal
|
||||
|
||||
from django.test import TestCase
|
||||
from rest_framework.test import APIClient
|
||||
|
||||
from apps.accounts.models import Team, TeamMember, User
|
||||
from apps.ai.models import AITask, ModelConfig, ModelProvider
|
||||
@@ -58,3 +59,31 @@ class CreditLedgerTests(TestCase):
|
||||
self.assertEqual(self.account.reserved_balance, Decimal("0.0000"))
|
||||
self.assertEqual(CreditLedger.objects.filter(ledger_type=CreditLedger.Type.RELEASE).count(), 1)
|
||||
|
||||
|
||||
class RechargePermissionTests(TestCase):
|
||||
"""充值是团队资金操作:仅 owner/admin 可发起,普通成员/访客被 403 拦截。"""
|
||||
|
||||
def setUp(self):
|
||||
self.owner = User.objects.create_user(username="r-owner", password="ownerpass123")
|
||||
self.member = User.objects.create_user(username="r-member", password="memberpass123")
|
||||
self.team = Team.objects.create(name="Recharge Team", owner=self.owner)
|
||||
TeamMember.objects.create(team=self.team, user=self.owner, role=TeamMember.Role.OWNER, status=TeamMember.Status.ACTIVE)
|
||||
TeamMember.objects.create(team=self.team, user=self.member, role=TeamMember.Role.MEMBER, status=TeamMember.Status.ACTIVE)
|
||||
CreditAccount.objects.create(team=self.team, balance=Decimal("0.0000"))
|
||||
|
||||
def test_owner_can_recharge(self):
|
||||
client = APIClient()
|
||||
client.force_authenticate(self.owner)
|
||||
response = client.post("/api/billing/recharge/", {"amount": "100"}, format="json")
|
||||
self.assertEqual(response.status_code, 201)
|
||||
account = CreditAccount.objects.get(team=self.team)
|
||||
self.assertEqual(account.balance, Decimal("100.0000"))
|
||||
|
||||
def test_member_cannot_recharge(self):
|
||||
client = APIClient()
|
||||
client.force_authenticate(self.member)
|
||||
response = client.post("/api/billing/recharge/", {"amount": "100"}, format="json")
|
||||
self.assertEqual(response.status_code, 403)
|
||||
account = CreditAccount.objects.get(team=self.team)
|
||||
self.assertEqual(account.balance, Decimal("0.0000")) # 余额未变,越权被拦
|
||||
|
||||
|
||||
@@ -11,7 +11,7 @@ from rest_framework.permissions import IsAuthenticated
|
||||
from rest_framework.response import Response
|
||||
|
||||
from apps.ai.models import AITask
|
||||
from apps.common.api import get_current_team
|
||||
from apps.common.api import can_manage_team, get_current_team
|
||||
|
||||
from .models import CreditAccount, CreditLedger
|
||||
from .serializers import CreditAccountSerializer, CreditLedgerSerializer
|
||||
@@ -83,6 +83,9 @@ def ledgers(request):
|
||||
@permission_classes([IsAuthenticated])
|
||||
def recharge(request):
|
||||
team = get_current_team(request.user)
|
||||
# 充值是团队资金操作:仅 owner/admin 可发起,普通成员/访客一律拒绝
|
||||
if not can_manage_team(request.user, team):
|
||||
return Response({"detail": "permission denied"}, status=status.HTTP_403_FORBIDDEN)
|
||||
try:
|
||||
amount = Decimal(str(request.data.get("amount", "0")))
|
||||
bonus = Decimal(str(request.data.get("bonus", "0")))
|
||||
|
||||
@@ -8,6 +8,12 @@ def get_current_team(user):
|
||||
return membership.team
|
||||
|
||||
|
||||
def can_manage_team(user, team):
|
||||
"""团队管理权限:仅 owner / admin 可做充值、改限额、增删成员等管理动作。"""
|
||||
membership = user.team_memberships.filter(team=team, status="active").first()
|
||||
return bool(membership and membership.role in {"owner", "admin"})
|
||||
|
||||
|
||||
class TeamScopedViewSetMixin:
|
||||
team_field = "team"
|
||||
|
||||
|
||||
@@ -0,0 +1,57 @@
|
||||
"""Celery worker 在线探测 —— 生成类任务的提交前置闸。
|
||||
|
||||
设计决策(2026-06-10):无 worker 时**禁止提交**图片/视频生成,而非静默降级到前端轮询。
|
||||
旧行为的风险:提交到 ARK 后浏览器一关就无人轮询,生成结果悬在云端无人取回入库、
|
||||
额度预扣长期冻结——等于数据丢失。额度预检挡"钱不够",这道闸挡"环境不齐"。
|
||||
|
||||
探测用 control.ping 广播,结果带缓存:在线缓存 30s(别让每次提交都付 ping 开销),
|
||||
离线只缓存 5s(worker 一启动尽快解封)。CELERY_TASK_ALWAYS_EAGER(测试/同步模式)
|
||||
下任务就地执行不需要 worker,直接放行。
|
||||
"""
|
||||
|
||||
import time
|
||||
|
||||
from django.conf import settings
|
||||
from rest_framework.exceptions import APIException
|
||||
|
||||
|
||||
class WorkerUnavailable(APIException):
|
||||
status_code = 503
|
||||
default_detail = (
|
||||
"后台任务服务(Celery worker)未运行:生成结果将无人取回入库,可能造成数据丢失。"
|
||||
"已暂停生成功能,请启动 worker 后重试(celery -A airshelf worker)。"
|
||||
)
|
||||
default_code = "worker_unavailable"
|
||||
|
||||
|
||||
_OK_TTL = 30.0
|
||||
_FAIL_TTL = 5.0
|
||||
_cache = {"ok": False, "expires": 0.0}
|
||||
|
||||
|
||||
def _ping_workers() -> bool:
|
||||
from airshelf.celery import app as celery_app
|
||||
|
||||
# limit=1:收到第一个 worker 回包立即返回,不傻等满 timeout
|
||||
return bool(celery_app.control.ping(timeout=1.0, limit=1))
|
||||
|
||||
|
||||
def celery_worker_available() -> bool:
|
||||
if getattr(settings, "CELERY_TASK_ALWAYS_EAGER", False):
|
||||
return True
|
||||
now = time.monotonic()
|
||||
if now < _cache["expires"]:
|
||||
return _cache["ok"]
|
||||
try:
|
||||
ok = _ping_workers()
|
||||
except Exception: # noqa: BLE001 — broker 不可达等同于无 worker,一样要拦
|
||||
ok = False
|
||||
_cache["ok"] = ok
|
||||
_cache["expires"] = now + (_OK_TTL if ok else _FAIL_TTL)
|
||||
return ok
|
||||
|
||||
|
||||
def require_worker() -> None:
|
||||
"""生成类提交入口的前置检查:无 worker 直接 503,不让任务出门。"""
|
||||
if not celery_worker_available():
|
||||
raise WorkerUnavailable()
|
||||
@@ -0,0 +1,45 @@
|
||||
from django.core.management.base import BaseCommand
|
||||
|
||||
from apps.ai.catalog import VOLCANO_PROVIDER, YUNQI_MODELS, YUNQI_PROVIDER
|
||||
from apps.ai.models import ModelConfig, ModelProvider
|
||||
|
||||
|
||||
class Command(BaseCommand):
|
||||
help = "Create or update YunQi (New API) provider with gpt-image-2, and disable Volcano image models."
|
||||
|
||||
def handle(self, *args, **options):
|
||||
provider, _ = ModelProvider.objects.update_or_create(
|
||||
name=YUNQI_PROVIDER["name"],
|
||||
defaults={
|
||||
"display_name": YUNQI_PROVIDER["display_name"],
|
||||
"base_url": YUNQI_PROVIDER["base_url"],
|
||||
"status": ModelProvider.Status.ACTIVE,
|
||||
},
|
||||
)
|
||||
|
||||
count = 0
|
||||
for item in YUNQI_MODELS:
|
||||
ModelConfig.objects.update_or_create(
|
||||
provider=provider,
|
||||
name=item["name"],
|
||||
capability=item["capability"],
|
||||
defaults={
|
||||
"display_name": item["display_name"],
|
||||
"endpoint": item["endpoint"],
|
||||
"status": ModelConfig.Status.ACTIVE,
|
||||
"metadata": item["metadata"],
|
||||
},
|
||||
)
|
||||
count += 1
|
||||
|
||||
# 全站生图切到 YunQi:火山 image 模型停用,文本/视频模型保持不动
|
||||
disabled = ModelConfig.objects.filter(
|
||||
provider__name=VOLCANO_PROVIDER["name"],
|
||||
capability=ModelConfig.Capability.IMAGE,
|
||||
).update(status=ModelConfig.Status.DISABLED)
|
||||
|
||||
self.stdout.write(
|
||||
self.style.SUCCESS(
|
||||
f"Bootstrapped {count} YunQi model config(s); disabled {disabled} Volcano image model(s)."
|
||||
)
|
||||
)
|
||||
@@ -1,3 +1,5 @@
|
||||
import uuid
|
||||
|
||||
from rest_framework import serializers
|
||||
|
||||
from apps.assets.serializers import AssetFileSerializer
|
||||
@@ -40,10 +42,11 @@ class ProjectStageSerializer(serializers.ModelSerializer):
|
||||
class VideoSegmentSerializer(serializers.ModelSerializer):
|
||||
adopted_asset = serializers.SerializerMethodField()
|
||||
adopted_asset_url = serializers.SerializerMethodField()
|
||||
versions = serializers.SerializerMethodField()
|
||||
|
||||
class Meta:
|
||||
model = VideoSegment
|
||||
fields = ["id", "sort_order", "target_duration_seconds", "status", "error_message", "adopted_version", "adopted_asset", "adopted_asset_url"]
|
||||
fields = ["id", "sort_order", "target_duration_seconds", "status", "error_message", "adopted_version", "adopted_asset", "adopted_asset_url", "versions"]
|
||||
read_only_fields = ["id", "sort_order", "target_duration_seconds", "status", "error_message", "adopted_version"]
|
||||
|
||||
def get_adopted_asset(self, obj):
|
||||
@@ -55,6 +58,21 @@ class VideoSegmentSerializer(serializers.ModelSerializer):
|
||||
version = obj.adopted_version
|
||||
return _asset_preview_url(version.asset) if version is not None else ""
|
||||
|
||||
def get_versions(self, obj):
|
||||
# 视频详情弹窗:历史版本(倒序,带可播放 URL)。模型无 Meta.ordering,这里按 created_at 排
|
||||
versions = sorted(obj.versions.all(), key=lambda v: (v.created_at is None, v.created_at), reverse=True)
|
||||
return [
|
||||
{
|
||||
"id": str(v.id),
|
||||
"asset": str(v.asset_id) if v.asset_id else None,
|
||||
"asset_url": _asset_preview_url(v.asset) if v.asset_id else "",
|
||||
"prompt": v.prompt,
|
||||
"is_adopted": v.is_adopted,
|
||||
"created_at": v.created_at.isoformat() if v.created_at else "",
|
||||
}
|
||||
for v in versions
|
||||
]
|
||||
|
||||
|
||||
class BaseAssetGroupSerializer(serializers.ModelSerializer):
|
||||
candidate_assets = serializers.PrimaryKeyRelatedField(many=True, read_only=True)
|
||||
@@ -167,11 +185,37 @@ class TimelineSerializer(serializers.ModelSerializer):
|
||||
export_jobs = TimelineExportJobSerializer(many=True, read_only=True)
|
||||
subtitle_tracks = SubtitleTrackSerializer(many=True, read_only=True)
|
||||
bgm_tracks = BgmTrackSerializer(many=True, read_only=True)
|
||||
voiceover = serializers.SerializerMethodField()
|
||||
|
||||
class Meta:
|
||||
model = Timeline
|
||||
fields = ["id", "name", "aspect_ratio", "resolution", "duration_seconds", "metadata", "clips", "export_jobs", "subtitle_tracks", "bgm_tracks"]
|
||||
read_only_fields = ["id", "clips", "export_jobs", "subtitle_tracks", "bgm_tracks"]
|
||||
fields = ["id", "name", "aspect_ratio", "resolution", "duration_seconds", "metadata", "voiceover", "clips", "export_jobs", "subtitle_tracks", "bgm_tracks"]
|
||||
read_only_fields = ["id", "voiceover", "clips", "export_jobs", "subtitle_tracks", "bgm_tracks"]
|
||||
|
||||
def get_voiceover(self, obj):
|
||||
# 旁白配音映射(metadata.voiceover)+ 每段新鲜的可播放 URL(TOS 签名 URL 会过期,不能存死)
|
||||
vo = (obj.metadata or {}).get("voiceover")
|
||||
if not isinstance(vo, dict) or not vo.get("items"):
|
||||
return None
|
||||
from apps.assets.models import Asset
|
||||
|
||||
ids = []
|
||||
for item in vo["items"]:
|
||||
try:
|
||||
ids.append(uuid.UUID(str(item.get("asset"))))
|
||||
except (TypeError, ValueError):
|
||||
continue
|
||||
assets = {str(a.id): a for a in Asset.objects.filter(id__in=ids).prefetch_related("files")}
|
||||
items = []
|
||||
for item in vo["items"]:
|
||||
asset = assets.get(str(item.get("asset")))
|
||||
items.append({**item, "asset_url": _asset_preview_url(asset)})
|
||||
return {
|
||||
"enabled": bool(vo.get("enabled")),
|
||||
"voice_type": str(vo.get("voice_type") or ""),
|
||||
"speed_ratio": vo.get("speed_ratio", 1.0),
|
||||
"items": items,
|
||||
}
|
||||
|
||||
|
||||
class ExportJobSerializer(serializers.ModelSerializer):
|
||||
|
||||
@@ -139,17 +139,49 @@ def _output_starts(specs: list[dict], xfade: float) -> tuple[list[float], float]
|
||||
_AFMT = "aformat=sample_fmts=fltp:sample_rates=44100:channel_layouts=stereo"
|
||||
|
||||
|
||||
def _probe_fps(path: Path) -> float:
|
||||
"""探测视频帧率(avg_frame_rate)。失败返回 0 由调用方兜底。"""
|
||||
try:
|
||||
proc = subprocess.run(
|
||||
["ffprobe", "-v", "error", "-select_streams", "v:0", "-show_entries",
|
||||
"stream=avg_frame_rate", "-of", "csv=p=0", str(path)],
|
||||
capture_output=True, timeout=60,
|
||||
)
|
||||
raw = proc.stdout.decode("utf-8", "ignore").strip().splitlines()[0] if proc.stdout else ""
|
||||
num, _, den = raw.partition("/")
|
||||
fps = float(num) / float(den or 1)
|
||||
return fps if 1.0 <= fps <= 120.0 else 0.0
|
||||
except Exception: # noqa: BLE001
|
||||
return 0.0
|
||||
|
||||
|
||||
def _pick_output_fps(fps_list: list[float]) -> float:
|
||||
"""输出帧率取各片段帧率的众数(并列取更高)。
|
||||
源 24fps 被硬转 30fps 会按 4:5 不均匀复帧——成片肉眼可见一卡一顿;帧率跟随源才是帧帧对应。"""
|
||||
valid = [round(f, 3) for f in fps_list if f > 0]
|
||||
if not valid:
|
||||
return 30.0
|
||||
counts: dict[float, int] = {}
|
||||
for f in valid:
|
||||
counts[f] = counts.get(f, 0) + 1
|
||||
best = max(counts.items(), key=lambda kv: (kv[1], kv[0]))
|
||||
return best[0]
|
||||
|
||||
|
||||
def _build_export_command(*, n: int, specs: list[dict], starts: list[float], total: float,
|
||||
transition: str, sub_overlays: list[tuple[str, float, float]],
|
||||
bgm_name: str | None, bgm_volume: float,
|
||||
has_audio: list[bool] | None = None) -> list[str]:
|
||||
has_audio: list[bool] | None = None, fps: float = 30.0,
|
||||
voice_overlays: list[tuple[str, float, float]] | None = None) -> list[str]:
|
||||
has_audio = has_audio or [False] * n
|
||||
voice_overlays = voice_overlays or []
|
||||
fps_expr = f"{fps:.6g}"
|
||||
parts: list[str] = []
|
||||
for i, s in enumerate(specs):
|
||||
parts.append(
|
||||
f"[{i}:v]trim=start={s['ts']:.3f}:end={s['te']:.3f},setpts=PTS-STARTPTS,"
|
||||
"scale=1080:1920:force_original_aspect_ratio=decrease,"
|
||||
"pad=1080:1920:(ow-iw)/2:(oh-ih)/2,setsar=1,fps=30,format=yuv420p[v" + str(i) + "]"
|
||||
f"pad=1080:1920:(ow-iw)/2:(oh-ih)/2,setsar=1,fps={fps_expr},format=yuv420p[v" + str(i) + "]"
|
||||
)
|
||||
xname = XFADE_MAP.get(transition or "none")
|
||||
if xname and n > 1:
|
||||
@@ -175,7 +207,7 @@ def _build_export_command(*, n: int, specs: list[dict], starts: list[float], tot
|
||||
# 音频:片段自带的人声/原声必须保留(有声片段取原音轨,无声片段补等长静音,否则 concat 会缺流);
|
||||
# 若另挂了 BGM,则把 BGM 混到原声之上(amix,normalize=0 不自动衰减原声音量)。
|
||||
# 三种片段全无声且无 BGM 时,保持旧行为=纯视频不带音轨。
|
||||
want_audio = any(has_audio) or bool(bgm_name)
|
||||
want_audio = any(has_audio) or bool(bgm_name) or bool(voice_overlays)
|
||||
audio_label: str | None = None
|
||||
if want_audio:
|
||||
for i, s in enumerate(specs):
|
||||
@@ -191,9 +223,23 @@ def _build_export_command(*, n: int, specs: list[dict], starts: list[float], tot
|
||||
parts.append("".join(f"[a{i}]" for i in range(n)) + f"concat=n={n}:v=0:a=1[avoice0]")
|
||||
parts.append(f"[avoice0]atrim=0:{total:.3f},asetpts=PTS-STARTPTS[avoice]")
|
||||
audio_label = "avoice"
|
||||
# 旁白配音(TTS):每段延迟到所属片段在输出时间轴的起点,裁到片段时长,叠在基础音轨之上
|
||||
if voice_overlays:
|
||||
vo_base = n + (1 if bgm_name else 0) + len(sub_overlays)
|
||||
for j, (_name, vstart, vdur) in enumerate(voice_overlays):
|
||||
delay_ms = max(0, int(round(vstart * 1000)))
|
||||
parts.append(
|
||||
f"[{vo_base + j}:a]atrim=0:{vdur:.3f},asetpts=PTS-STARTPTS,{_AFMT},"
|
||||
f"adelay={delay_ms}|{delay_ms}[vo{j}]"
|
||||
)
|
||||
vo_inputs = "".join(f"[vo{j}]" for j in range(len(voice_overlays)))
|
||||
parts.append(
|
||||
f"[avoice]{vo_inputs}amix=inputs={len(voice_overlays) + 1}:duration=first:dropout_transition=0:normalize=0[anarr]"
|
||||
)
|
||||
audio_label = "anarr"
|
||||
if bgm_name:
|
||||
parts.append(f"[{n}:a]volume={bgm_volume:.3f},atrim=0:{total:.3f},asetpts=PTS-STARTPTS,{_AFMT}[abgm]")
|
||||
parts.append("[avoice][abgm]amix=inputs=2:duration=longest:dropout_transition=0:normalize=0[aout]")
|
||||
parts.append(f"[{audio_label}][abgm]amix=inputs=2:duration=longest:dropout_transition=0:normalize=0[aout]")
|
||||
audio_label = "aout"
|
||||
|
||||
cmd = ["ffmpeg", "-y"]
|
||||
@@ -203,10 +249,12 @@ def _build_export_command(*, n: int, specs: list[dict], starts: list[float], tot
|
||||
cmd += ["-stream_loop", "-1", "-i", bgm_name]
|
||||
for png, _s, _e in sub_overlays:
|
||||
cmd += ["-loop", "1", "-i", png]
|
||||
for name, _vs, _vd in voice_overlays:
|
||||
cmd += ["-i", name]
|
||||
cmd += ["-filter_complex", ";".join(parts), "-map", f"[{vlabel}]"]
|
||||
if audio_label:
|
||||
cmd += ["-map", f"[{audio_label}]"]
|
||||
cmd += ["-c:v", "libx264", "-pix_fmt", "yuv420p", "-r", "30", "-preset", "veryfast"]
|
||||
cmd += ["-c:v", "libx264", "-pix_fmt", "yuv420p", "-r", fps_expr, "-preset", "veryfast"]
|
||||
if audio_label:
|
||||
cmd += ["-c:a", "aac", "-b:a", "192k"]
|
||||
cmd += ["-t", f"{total:.3f}", "-movflags", "+faststart", "output.mp4"]
|
||||
@@ -232,23 +280,113 @@ def run_export_job_in_thread(export_job_id: str) -> None:
|
||||
threading.Thread(target=_worker, daemon=True).start()
|
||||
|
||||
|
||||
def _split_subtitle_text(text: str) -> list[str]:
|
||||
"""整段旁白 → 短句列表(与前端 splitSubtitleCues 同规则):
|
||||
硬标点(。!?;…)必切;长句(≥12 字)在逗号处再切;过短碎句(<5 字)并入前句;去尾部逗号句号。"""
|
||||
clean = " ".join(str(text or "").split())
|
||||
if not clean:
|
||||
return []
|
||||
hard = "\u3002\uff01\uff1f!?\uff1b;\u2026" # 。!?!?;;… 全角+半角
|
||||
soft = "\uff0c,\u3001" # ,,、
|
||||
parts: list[str] = []
|
||||
cur = ""
|
||||
for ch in clean:
|
||||
cur += ch
|
||||
if ch in hard or (ch in soft and len(cur) >= 12):
|
||||
parts.append(cur)
|
||||
cur = ""
|
||||
if cur.strip():
|
||||
parts.append(cur)
|
||||
merged: list[str] = []
|
||||
for raw in parts:
|
||||
p = raw.strip()
|
||||
if not p:
|
||||
continue
|
||||
core = sum(1 for c in p if c not in hard and c not in soft)
|
||||
if merged and core < 5:
|
||||
merged[-1] += p
|
||||
else:
|
||||
merged.append(p)
|
||||
out: list[str] = []
|
||||
for p in merged:
|
||||
p = p.rstrip("\uff0c\u3002\uff1b,;\u3001")
|
||||
if p:
|
||||
out.append(p)
|
||||
return out
|
||||
|
||||
|
||||
def _subtitle_cues(timeline, project, specs, starts, total) -> list[tuple[float, float, str]]:
|
||||
"""字幕条目:文本取 SubtitleTrack.content,空则回退脚本旁白;时间按输出布局(对 xfade 也对齐)。"""
|
||||
"""字幕条目(逐句):优先用 SubtitleTrack.content 里每条 cue 自带的 start_ms——
|
||||
先定位到所属片段(输入时间轴=各片段时长累计),再重映射到输出时间轴(xfade 会压缩起点);
|
||||
无保存字幕时回退脚本旁白,逐段按句切分铺满片段时长。旧行为是整段旁白糊满 15s(切割错误)。"""
|
||||
track = timeline.subtitle_tracks.filter(enabled=True).first() or timeline.subtitle_tracks.first()
|
||||
if track is None or track.enabled is False:
|
||||
return []
|
||||
texts: list[str] = [str((c or {}).get("text", "")) for c in (track.content or [])]
|
||||
if not any(t.strip() for t in texts):
|
||||
script = project.script_versions.filter(is_adopted=True).prefetch_related("segments").first()
|
||||
if script is not None:
|
||||
texts = [seg.narration for seg in script.segments.all().order_by("sort_order")]
|
||||
in_starts: list[float] = []
|
||||
acc = 0.0
|
||||
for s in specs:
|
||||
in_starts.append(acc)
|
||||
acc += s["dur"]
|
||||
|
||||
def clip_index(t_in: float) -> int:
|
||||
for k in range(len(specs)):
|
||||
if t_in < in_starts[k] + specs[k]["dur"]:
|
||||
return k
|
||||
return len(specs) - 1
|
||||
|
||||
# 旁白配音时长(秒)按片段索引取:有配音时字幕窗口必须跟语音走,而不是铺满片段
|
||||
vo_meta = (timeline.metadata or {}).get("voiceover") or {}
|
||||
vo_durs: dict[int, float] = {}
|
||||
if vo_meta.get("enabled"):
|
||||
for item in vo_meta.get("items") or []:
|
||||
try:
|
||||
vo_idx = int(item.get("index", -1))
|
||||
vo_dur = float(item.get("duration_ms") or 0) / 1000.0
|
||||
except (TypeError, ValueError):
|
||||
continue
|
||||
if vo_idx >= 0 and vo_dur > 0:
|
||||
vo_durs[vo_idx] = vo_dur
|
||||
|
||||
cues: list[tuple[float, float, str]] = []
|
||||
content = [c for c in (track.content or []) if str((c or {}).get("text", "")).strip()]
|
||||
if content:
|
||||
entries = sorted(content, key=lambda c: int((c or {}).get("start_ms", 0) or 0))
|
||||
outs: list[tuple[float, int, str, float | None]] = []
|
||||
for c in entries:
|
||||
t_in = int(c.get("start_ms", 0) or 0) / 1000.0
|
||||
i = clip_index(t_in)
|
||||
offset = max(0.0, min(specs[i]["dur"], t_in - in_starts[i]))
|
||||
# 新格式 cue 自带 end_ms(已对齐配音语速),同样重映射到输出时间轴;旧草稿无 end_ms 走相邻推断
|
||||
out_end: float | None = None
|
||||
if c.get("end_ms"):
|
||||
e_off = max(0.0, min(specs[i]["dur"], int(c["end_ms"]) / 1000.0 - in_starts[i]))
|
||||
out_end = starts[i] + e_off
|
||||
outs.append((starts[i] + offset, i, str(c.get("text", "")).strip(), out_end))
|
||||
for j, (start, i, text, out_end) in enumerate(outs):
|
||||
if out_end is None:
|
||||
if j + 1 < len(outs) and outs[j + 1][1] == i:
|
||||
out_end = outs[j + 1][0]
|
||||
else:
|
||||
out_end = min(total, starts[i] + specs[i]["dur"])
|
||||
cues.append((start, max(start + 0.5, min(total, out_end)), text))
|
||||
return cues
|
||||
|
||||
script = project.script_versions.filter(is_adopted=True).prefetch_related("segments").first()
|
||||
texts = [seg.narration for seg in script.segments.all().order_by("sort_order")] if script else []
|
||||
for i in range(len(specs)):
|
||||
text = texts[i] if i < len(texts) else ""
|
||||
start = starts[i]
|
||||
end = starts[i + 1] if i + 1 < len(starts) else total
|
||||
if text and text.strip():
|
||||
cues.append((start, max(start + 0.5, end), text))
|
||||
pieces = _split_subtitle_text(texts[i] if i < len(texts) else "")
|
||||
if not pieces:
|
||||
continue
|
||||
span = (starts[i + 1] if i + 1 < len(starts) else total) - starts[i]
|
||||
if i in vo_durs:
|
||||
span = min(span, vo_durs[i])
|
||||
total_chars = sum(len(p) for p in pieces) or 1
|
||||
acc_chars = 0
|
||||
for p in pieces:
|
||||
start = starts[i] + (acc_chars / total_chars) * span
|
||||
acc_chars += len(p)
|
||||
end = starts[i] + (acc_chars / total_chars) * span
|
||||
cues.append((start, max(start + 0.5, end), p))
|
||||
return cues
|
||||
|
||||
|
||||
@@ -279,6 +417,8 @@ def run_export_job(export_job_id: str) -> ExportJob:
|
||||
_download_asset_primary_file(clip.asset, tmp / f"clip{index}.mp4")
|
||||
# 逐片段探测是否自带音轨:有声→保留原声,无声→补静音(见 _build_export_command)
|
||||
has_audio = [_has_audio_stream(tmp / f"clip{index}.mp4") for index in range(len(clips))]
|
||||
# 输出帧率跟随源帧率众数(Seedance 出 24fps,硬转 30fps 会不均匀复帧=成片一卡一顿)
|
||||
output_fps = _pick_output_fps([_probe_fps(tmp / f"clip{index}.mp4") for index in range(len(clips))])
|
||||
|
||||
bgm_name = None
|
||||
if bgm_track is not None and bgm_track.asset_id:
|
||||
@@ -294,13 +434,36 @@ def run_export_job(export_job_id: str) -> ExportJob:
|
||||
_render_subtitle_png(text, style_key, tmp / png)
|
||||
sub_overlays.append((png, start, end))
|
||||
|
||||
# 旁白配音(TTS 资产):按 timeline.metadata.voiceover 映射下载,人声轨混在 BGM 之上;
|
||||
# 逐句条目带 offset_ms(句内起点,可被拖动调整),输出位置 = 片段输出起点 + 句内偏移
|
||||
vo_meta = (timeline.metadata or {}).get("voiceover") or {}
|
||||
voice_overlays: list[tuple[str, float, float]] = []
|
||||
if vo_meta.get("enabled") and isinstance(vo_meta.get("items"), list):
|
||||
for j, item in enumerate(vo_meta["items"]):
|
||||
try:
|
||||
seg_index = int(item.get("index", -1))
|
||||
offset_s = max(0.0, float(item.get("offset_ms") or 0) / 1000.0)
|
||||
except (TypeError, ValueError):
|
||||
continue
|
||||
if seg_index < 0 or seg_index >= len(clips) or not item.get("asset"):
|
||||
continue
|
||||
remain = specs[seg_index]["dur"] - offset_s
|
||||
if remain <= 0.05:
|
||||
continue # 句子被拖到片段尾外,导出时丢弃(预览同样不响)
|
||||
vo_asset = Asset.objects.filter(team=project.team, id=item["asset"]).first()
|
||||
if vo_asset is None:
|
||||
continue
|
||||
vo_name = f"vo{j}.mp3"
|
||||
_download_asset_primary_file(vo_asset, tmp / vo_name)
|
||||
voice_overlays.append((vo_name, starts[seg_index] + offset_s, remain))
|
||||
|
||||
export_job.progress = 35
|
||||
export_job.save(update_fields=["progress", "updated_at"])
|
||||
|
||||
command = _build_export_command(
|
||||
n=len(clips), specs=specs, starts=starts, total=total, transition=transition,
|
||||
sub_overlays=sub_overlays, bgm_name=bgm_name, bgm_volume=(bgm_track.volume / 100.0) if bgm_track else 1.0,
|
||||
has_audio=has_audio,
|
||||
has_audio=has_audio, fps=output_fps, voice_overlays=voice_overlays,
|
||||
)
|
||||
proc = subprocess.run(command, cwd=str(tmp), capture_output=True)
|
||||
if proc.returncode != 0:
|
||||
|
||||
@@ -5,9 +5,19 @@ from rest_framework.test import APIClient
|
||||
|
||||
from apps.accounts.models import Team, TeamMember, User
|
||||
from apps.ai.models import ModelConfig, ModelProvider
|
||||
from apps.assets.models import Asset
|
||||
from apps.billing.models import CreditAccount, CreditLedger
|
||||
from apps.products.models import Product
|
||||
from apps.projects.models import Project, ProjectStage, ScriptVersion, VideoSegment
|
||||
from apps.projects.models import (
|
||||
Project,
|
||||
ProjectStage,
|
||||
ScriptVersion,
|
||||
SubtitleTrack,
|
||||
Timeline,
|
||||
TimelineClip,
|
||||
VideoSegment,
|
||||
VideoSegmentVersion,
|
||||
)
|
||||
|
||||
|
||||
class ProjectApiTests(TestCase):
|
||||
@@ -100,3 +110,207 @@ class ProjectApiTests(TestCase):
|
||||
self.assertEqual(script.segments.count(), 4)
|
||||
self.assertEqual(CreditLedger.objects.filter(team=self.team, ledger_type=CreditLedger.Type.CHARGE).count(), 1)
|
||||
|
||||
def test_adopt_video_version_remaps_timeline_draft(self):
|
||||
"""切换采用版本后,时间线草稿里引用本场旧版本资产的片段必须跟随换成新资产(剪辑台/导出都读草稿)。"""
|
||||
project = Project.objects.create(team=self.team, created_by=self.user, product=self.product, name="P")
|
||||
make_asset = lambda name: Asset.objects.create(
|
||||
team=self.team, created_by=self.user, name=name,
|
||||
asset_type=Asset.Type.VIDEO, source=Asset.Source.AI_GENERATED, category=Asset.Category.VIDEO_CLIP,
|
||||
)
|
||||
old_asset, new_asset = make_asset("v1"), make_asset("v2")
|
||||
segment = VideoSegment.objects.create(project=project, sort_order=0)
|
||||
old_ver = VideoSegmentVersion.objects.create(video_segment=segment, asset=old_asset, is_adopted=True)
|
||||
new_ver = VideoSegmentVersion.objects.create(video_segment=segment, asset=new_asset)
|
||||
segment.adopted_version = old_ver
|
||||
segment.save(update_fields=["adopted_version"])
|
||||
timeline = Timeline.objects.create(project=project)
|
||||
clip = TimelineClip.objects.create(
|
||||
timeline=timeline, asset=old_asset, sort_order=0,
|
||||
duration_ms=15000, trim_start_ms=1000, trim_end_ms=9000,
|
||||
)
|
||||
|
||||
response = self.client.post(
|
||||
f"/api/projects/{project.id}/adopt-video-version/",
|
||||
{"video_segment_id": str(segment.id), "version_id": str(new_ver.id)},
|
||||
format="json",
|
||||
)
|
||||
|
||||
self.assertEqual(response.status_code, 200)
|
||||
clip.refresh_from_db()
|
||||
self.assertEqual(clip.asset_id, new_asset.id)
|
||||
# 旧素材上的裁剪点对新素材无意义,必须复位
|
||||
self.assertEqual(clip.trim_start_ms, 0)
|
||||
self.assertIsNone(clip.trim_end_ms)
|
||||
|
||||
def _make_audio_model(self):
|
||||
return ModelConfig.objects.create(
|
||||
provider=self.provider, name="doubao-voice-bigtts", display_name="豆包语音合成",
|
||||
capability="audio", endpoint="api/v1/tts", unit_price="2.0000",
|
||||
)
|
||||
|
||||
@patch("apps.ai.services.TosStorage")
|
||||
@patch("apps.ai.services.VolcanoTtsProvider")
|
||||
def test_generate_voiceover_creates_assets_and_charges_once(self, provider_cls, storage_cls):
|
||||
"""每镜旁白 TTS:产出音频资产、timeline.metadata 落映射、一次调用只计一次费。"""
|
||||
self._make_audio_model()
|
||||
provider = provider_cls.return_value
|
||||
provider.configured = True
|
||||
provider.synthesize.return_value = (b"fake-mp3-bytes", 4200)
|
||||
stored = storage_cls.return_value.upload_fileobj.return_value
|
||||
stored.object_key, stored.bucket, stored.content_type, stored.size_bytes = "k.mp3", "b", "audio/mpeg", 14
|
||||
|
||||
project = Project.objects.create(team=self.team, created_by=self.user, product=self.product, name="VO")
|
||||
response = self.client.post(
|
||||
f"/api/projects/{project.id}/generate-voiceover/",
|
||||
{"items": [{"index": 0, "text": "第一镜旁白"}, {"index": 1, "text": "第二镜旁白"}]},
|
||||
format="json",
|
||||
)
|
||||
|
||||
self.assertEqual(response.status_code, 201)
|
||||
voiceover = response.json()["voiceover"]
|
||||
self.assertTrue(voiceover["enabled"])
|
||||
self.assertEqual(len(voiceover["items"]), 2)
|
||||
self.assertEqual(provider.synthesize.call_count, 2)
|
||||
from apps.assets.models import Asset as AssetModel
|
||||
self.assertEqual(AssetModel.objects.filter(team=self.team, asset_type="audio").count(), 2)
|
||||
self.assertEqual(
|
||||
CreditLedger.objects.filter(team=self.team, ledger_type=CreditLedger.Type.CHARGE).count(), 1
|
||||
)
|
||||
project.refresh_from_db()
|
||||
self.assertEqual(len(project.timeline.metadata["voiceover"]["items"]), 2)
|
||||
|
||||
@patch("apps.ai.services.VolcanoTtsProvider")
|
||||
def test_generate_voiceover_unconfigured_returns_400_and_no_charge(self, provider_cls):
|
||||
"""语音凭证未配置:400 + 人话提示,不建任务不扣费。"""
|
||||
self._make_audio_model()
|
||||
provider_cls.return_value.configured = False
|
||||
project = Project.objects.create(team=self.team, created_by=self.user, product=self.product, name="VO2")
|
||||
response = self.client.post(
|
||||
f"/api/projects/{project.id}/generate-voiceover/",
|
||||
{"items": [{"index": 0, "text": "旁白"}]},
|
||||
format="json",
|
||||
)
|
||||
self.assertEqual(response.status_code, 400)
|
||||
self.assertIn("语音合成未配置", response.json()["detail"])
|
||||
self.assertEqual(CreditLedger.objects.filter(team=self.team, ledger_type=CreditLedger.Type.CHARGE).count(), 0)
|
||||
|
||||
def test_export_command_mixes_voiceover_above_bgm(self):
|
||||
"""导出命令:人声按片段起点 adelay 延迟、先与原声 amix、再与 BGM amix。"""
|
||||
from apps.projects.services.export import _build_export_command
|
||||
|
||||
specs = [{"ts": 0.0, "te": 15.0, "dur": 15.0}, {"ts": 0.0, "te": 15.0, "dur": 15.0}]
|
||||
cmd = _build_export_command(
|
||||
n=2, specs=specs, starts=[0.0, 15.0], total=30.0, transition="none",
|
||||
sub_overlays=[], bgm_name="bgm.mp3", bgm_volume=0.6, has_audio=[False, False],
|
||||
voice_overlays=[("vo0.mp3", 0.0, 15.0), ("vo1.mp3", 15.0, 15.0)],
|
||||
)
|
||||
self.assertIn("vo0.mp3", cmd)
|
||||
self.assertIn("vo1.mp3", cmd)
|
||||
graph = cmd[cmd.index("-filter_complex") + 1]
|
||||
self.assertIn("adelay=15000|15000", graph)
|
||||
self.assertIn("[avoice][vo0][vo1]amix=inputs=3", graph)
|
||||
self.assertIn("[anarr][abgm]amix=inputs=2", graph)
|
||||
|
||||
def test_save_timeline_updates_voiceover_offsets(self):
|
||||
"""拖动字幕块后保存:按 asset 回写句内起点 offset_ms,未拖动的句不受影响。"""
|
||||
project = Project.objects.create(team=self.team, created_by=self.user, product=self.product, name="VOFF")
|
||||
Timeline.objects.create(project=project, metadata={"voiceover": {"enabled": True, "items": [
|
||||
{"index": 0, "cue": 0, "asset": "a1", "text": "x", "duration_ms": 2000, "offset_ms": 0},
|
||||
{"index": 0, "cue": 1, "asset": "a2", "text": "y", "duration_ms": 2000, "offset_ms": 2000},
|
||||
]}})
|
||||
response = self.client.post(
|
||||
f"/api/projects/{project.id}/save-timeline/",
|
||||
{"voiceover": {"items": [{"asset": "a2", "offset_ms": 9000}]}},
|
||||
format="json",
|
||||
)
|
||||
self.assertEqual(response.status_code, 200)
|
||||
project.timeline.refresh_from_db()
|
||||
items = project.timeline.metadata["voiceover"]["items"]
|
||||
self.assertEqual(items[0]["offset_ms"], 0)
|
||||
self.assertEqual(items[1]["offset_ms"], 9000)
|
||||
|
||||
def test_subtitle_cues_use_end_ms_and_follow_voiceover(self):
|
||||
"""字幕烧入时序:cue 自带 end_ms 时按它收尾(不再拖到片段尾)——这是配音对齐的关键;
|
||||
无保存字幕的回退路径,有配音时逐句窗口也要压进语音真实时长。"""
|
||||
from apps.projects.services.export import _subtitle_cues
|
||||
|
||||
project = Project.objects.create(team=self.team, created_by=self.user, product=self.product, name="SUB")
|
||||
timeline = Timeline.objects.create(
|
||||
project=project,
|
||||
metadata={"voiceover": {"enabled": True, "items": [{"index": 0, "asset": "x", "text": "t", "duration_ms": 4000}]}},
|
||||
)
|
||||
SubtitleTrack.objects.create(timeline=timeline, enabled=True, content=[
|
||||
{"start_ms": 0, "end_ms": 2000, "text": "第一句"},
|
||||
{"start_ms": 2000, "end_ms": 4000, "text": "第二句"},
|
||||
])
|
||||
specs = [{"ts": 0.0, "te": 15.0, "dur": 15.0}]
|
||||
cues = _subtitle_cues(timeline, project, specs, [0.0], 15.0)
|
||||
self.assertEqual(len(cues), 2)
|
||||
self.assertEqual((cues[0][0], cues[0][1]), (0.0, 2.0))
|
||||
self.assertEqual(cues[1][1], 4.0) # 最后一句在语音念完处收尾,而非拖到 15s
|
||||
|
||||
def test_save_timeline_voiceover_toggle_and_clear(self):
|
||||
"""保存草稿可关/开配音、可移除映射(不动已生成的音频资产)。"""
|
||||
project = Project.objects.create(team=self.team, created_by=self.user, product=self.product, name="VO3")
|
||||
Timeline.objects.create(
|
||||
project=project,
|
||||
metadata={"voiceover": {"enabled": True, "voice_type": "v", "items": [{"index": 0, "asset": "a", "text": "t"}]}},
|
||||
)
|
||||
response = self.client.post(
|
||||
f"/api/projects/{project.id}/save-timeline/", {"voiceover": {"enabled": False}}, format="json"
|
||||
)
|
||||
self.assertEqual(response.status_code, 200)
|
||||
project.timeline.refresh_from_db()
|
||||
self.assertFalse(project.timeline.metadata["voiceover"]["enabled"])
|
||||
|
||||
response = self.client.post(
|
||||
f"/api/projects/{project.id}/save-timeline/", {"voiceover": {"clear": True}}, format="json"
|
||||
)
|
||||
self.assertEqual(response.status_code, 200)
|
||||
project.timeline.refresh_from_db()
|
||||
self.assertNotIn("voiceover", project.timeline.metadata)
|
||||
|
||||
|
||||
class WorkerGateTests(TestCase):
|
||||
"""无 Celery worker 时,图片/视频生成入口必须 503 拒绝(防结果无人取回=数据丢失)。"""
|
||||
|
||||
def setUp(self):
|
||||
self.user = User.objects.create_user(username="gate", password="pass")
|
||||
self.team = Team.objects.create(name="Gate Team", owner=self.user)
|
||||
TeamMember.objects.create(team=self.team, user=self.user, role=TeamMember.Role.OWNER)
|
||||
CreditAccount.objects.create(team=self.team, balance="100.0000")
|
||||
self.product = Product.objects.create(team=self.team, created_by=self.user, title="Gate Product")
|
||||
self.project = Project.objects.create(team=self.team, created_by=self.user, name="Gate", product=self.product)
|
||||
self.client = APIClient()
|
||||
self.client.force_authenticate(self.user)
|
||||
|
||||
def _blocked(self):
|
||||
return patch("apps.common.celery_health.celery_worker_available", return_value=False)
|
||||
|
||||
def test_generation_endpoints_blocked_without_worker(self):
|
||||
cases = [
|
||||
(f"/api/projects/{self.project.id}/generate-base-asset/", {"kind": "person"}),
|
||||
(f"/api/projects/{self.project.id}/generate-storyboard/", {}),
|
||||
(f"/api/projects/{self.project.id}/submit-video-segment/", {"video_segment_id": "x"}),
|
||||
("/api/ai/generate-image/", {"prompt": "测试"}),
|
||||
]
|
||||
with self._blocked():
|
||||
for url, payload in cases:
|
||||
response = self.client.post(url, payload, format="json")
|
||||
self.assertEqual(response.status_code, 503, msg=f"{url} 应被 worker 闸拦截")
|
||||
self.assertIn("worker", response.json()["detail"].lower() if isinstance(response.json().get("detail"), str) else "")
|
||||
|
||||
def test_eager_mode_bypasses_gate(self):
|
||||
from apps.common.celery_health import celery_worker_available
|
||||
|
||||
# 测试配置 CELERY_TASK_ALWAYS_EAGER=True:任务就地执行,无需 worker,闸放行
|
||||
self.assertTrue(celery_worker_available())
|
||||
|
||||
def test_ping_failure_means_unavailable(self):
|
||||
from apps.common import celery_health
|
||||
|
||||
celery_health._cache["expires"] = 0.0
|
||||
with self.settings(CELERY_TASK_ALWAYS_EAGER=False):
|
||||
with patch.object(celery_health, "_ping_workers", side_effect=ConnectionError("broker down")):
|
||||
self.assertFalse(celery_health.celery_worker_available())
|
||||
celery_health._cache["expires"] = 0.0
|
||||
|
||||
@@ -9,7 +9,10 @@ from rest_framework.parsers import FormParser, MultiPartParser
|
||||
from rest_framework.response import Response
|
||||
from rest_framework.viewsets import ModelViewSet
|
||||
|
||||
from apps.ai.providers import TtsNotConfigured
|
||||
from apps.ai.services import (
|
||||
DEFAULT_VOICEOVER_VOICE,
|
||||
VOICEOVER_VOICES,
|
||||
create_export_job,
|
||||
generate_base_asset,
|
||||
generate_project_script,
|
||||
@@ -17,11 +20,13 @@ from apps.ai.services import (
|
||||
poll_video_segment,
|
||||
submit_storyboard,
|
||||
submit_video_segment,
|
||||
synthesize_project_voiceover,
|
||||
)
|
||||
from apps.assets.models import Asset, AssetFile
|
||||
from apps.assets.serializers import AssetFileSerializer
|
||||
from apps.assets.storage import TosStorage
|
||||
from apps.common.api import TeamScopedViewSetMixin
|
||||
from apps.common.celery_health import require_worker
|
||||
|
||||
from .models import (
|
||||
BaseAssetGroup,
|
||||
@@ -29,6 +34,7 @@ from .models import (
|
||||
ExportJob,
|
||||
Project,
|
||||
ProjectStage,
|
||||
ScriptSegment,
|
||||
ScriptVersion,
|
||||
SubtitleTrack,
|
||||
Timeline,
|
||||
@@ -92,6 +98,7 @@ class ProjectViewSet(TeamScopedViewSetMixin, ModelViewSet):
|
||||
"stages",
|
||||
"video_segments",
|
||||
"video_segments__adopted_version__asset__files",
|
||||
"video_segments__versions__asset__files",
|
||||
"script_versions",
|
||||
"script_versions__segments",
|
||||
"base_asset_groups",
|
||||
@@ -121,6 +128,7 @@ class ProjectViewSet(TeamScopedViewSetMixin, ModelViewSet):
|
||||
user=request.user,
|
||||
user_prompt=request.data.get("prompt", ""),
|
||||
selling_point_ids=request.data.get("selling_point_ids") or [],
|
||||
source=request.data.get("source") or "ai",
|
||||
)
|
||||
return Response(ScriptVersionSerializer(script).data, status=status.HTTP_201_CREATED)
|
||||
|
||||
@@ -143,6 +151,7 @@ class ProjectViewSet(TeamScopedViewSetMixin, ModelViewSet):
|
||||
|
||||
@action(detail=True, methods=["post"], url_path="generate-base-asset")
|
||||
def generate_base_asset_action(self, request, pk=None):
|
||||
require_worker()
|
||||
project = self.get_object()
|
||||
kind = request.data.get("kind")
|
||||
if kind not in BaseAssetGroup.Kind.values:
|
||||
@@ -167,9 +176,117 @@ class ProjectViewSet(TeamScopedViewSetMixin, ModelViewSet):
|
||||
promote_base_asset_stage_if_ready(project)
|
||||
return Response(BaseAssetGroupSerializer(group).data)
|
||||
|
||||
# ── Stage 1 · 镜头脚本逐字段编辑 / 增删分镜 ──
|
||||
|
||||
def _sync_video_segments_to_script(self, project: Project, script: ScriptVersion) -> None:
|
||||
"""采用版分镜数变化时,同步 VideoSegment 数量:不足则尾部补 NOT_STARTED,
|
||||
多出且尾部是「从未生成过」的段则裁掉(已生成的段绝不动)。"""
|
||||
if not script.is_adopted:
|
||||
return
|
||||
target = script.segments.count()
|
||||
segments = list(project.video_segments.order_by("sort_order"))
|
||||
while len(segments) > target:
|
||||
tail = segments[-1]
|
||||
if tail.status == VideoSegment.Status.NOT_STARTED and not tail.versions.exists():
|
||||
tail.delete()
|
||||
segments.pop()
|
||||
else:
|
||||
break
|
||||
next_order = (segments[-1].sort_order + 1) if segments else 0
|
||||
for _ in range(target - len(segments)):
|
||||
segments.append(VideoSegment.objects.create(project=project, sort_order=next_order, target_duration_seconds=15))
|
||||
next_order += 1
|
||||
|
||||
@action(detail=True, methods=["post"], url_path="update-script-segment")
|
||||
def update_script_segment(self, request, pk=None):
|
||||
project = self.get_object()
|
||||
segment = ScriptSegment.objects.get(id=request.data.get("segment_id"), script_version__project=project)
|
||||
changed = []
|
||||
for field in ("narration", "visual_prompt"):
|
||||
if field in request.data:
|
||||
setattr(segment, field, str(request.data.get(field) or "").strip())
|
||||
changed.append(field)
|
||||
if "duration_seconds" in request.data:
|
||||
try:
|
||||
segment.duration_seconds = max(1, min(60, int(request.data["duration_seconds"])))
|
||||
changed.append("duration_seconds")
|
||||
except (TypeError, ValueError):
|
||||
pass
|
||||
if changed:
|
||||
segment.save(update_fields=[*changed, "updated_at"])
|
||||
return Response(ScriptVersionSerializer(segment.script_version).data)
|
||||
|
||||
@action(detail=True, methods=["post"], url_path="add-script-segment")
|
||||
@transaction.atomic
|
||||
def add_script_segment(self, request, pk=None):
|
||||
project = self.get_object()
|
||||
after = ScriptSegment.objects.select_related("script_version").get(
|
||||
id=request.data.get("after_segment_id"), script_version__project=project
|
||||
)
|
||||
script = after.script_version
|
||||
segments = list(script.segments.order_by("sort_order"))
|
||||
insert_at = next(i for i, s in enumerate(segments) if s.id == after.id) + 1
|
||||
created = ScriptSegment.objects.create(
|
||||
script_version=script,
|
||||
sort_order=insert_at,
|
||||
duration_seconds=int(request.data.get("duration_seconds") or 15),
|
||||
narration=str(request.data.get("narration") or "").strip(),
|
||||
visual_prompt=str(request.data.get("visual_prompt") or "").strip(),
|
||||
)
|
||||
segments.insert(insert_at, created)
|
||||
for index, seg in enumerate(segments):
|
||||
if seg.sort_order != index:
|
||||
seg.sort_order = index
|
||||
seg.save(update_fields=["sort_order", "updated_at"])
|
||||
self._sync_video_segments_to_script(project, script)
|
||||
return Response(ScriptVersionSerializer(script).data, status=status.HTTP_201_CREATED)
|
||||
|
||||
@action(detail=True, methods=["post"], url_path="delete-script-segment")
|
||||
@transaction.atomic
|
||||
def delete_script_segment(self, request, pk=None):
|
||||
project = self.get_object()
|
||||
segment = ScriptSegment.objects.select_related("script_version").get(
|
||||
id=request.data.get("segment_id"), script_version__project=project
|
||||
)
|
||||
script = segment.script_version
|
||||
if script.segments.count() <= 1:
|
||||
return Response({"detail": "至少保留一个分镜"}, status=status.HTTP_400_BAD_REQUEST)
|
||||
segment.delete()
|
||||
for index, seg in enumerate(script.segments.order_by("sort_order")):
|
||||
if seg.sort_order != index:
|
||||
seg.sort_order = index
|
||||
seg.save(update_fields=["sort_order", "updated_at"])
|
||||
self._sync_video_segments_to_script(project, script)
|
||||
return Response(ScriptVersionSerializer(script).data)
|
||||
|
||||
# ── Stage 4 · 视频版本采用(详情弹窗里切历史版) ──
|
||||
@action(detail=True, methods=["post"], url_path="adopt-video-version")
|
||||
@transaction.atomic
|
||||
def adopt_video_version(self, request, pk=None):
|
||||
project = self.get_object()
|
||||
segment = VideoSegment.objects.get(project=project, id=request.data.get("video_segment_id"))
|
||||
version = VideoSegmentVersion.objects.get(video_segment=segment, id=request.data.get("version_id"))
|
||||
segment.versions.update(is_adopted=False)
|
||||
version.is_adopted = True
|
||||
version.save(update_fields=["is_adopted", "updated_at"])
|
||||
segment.adopted_version = version
|
||||
segment.status = VideoSegment.Status.SUCCEEDED
|
||||
segment.error_message = ""
|
||||
segment.save(update_fields=["adopted_version", "status", "error_message", "updated_at"])
|
||||
# 时间线草稿(剪辑台/拼接导出都读它)若还引用本场旧版本的资产,必须跟随切换,
|
||||
# 否则采用新版本后预览和导出仍是旧片段;裁剪点是对旧素材设的,一并复位
|
||||
timeline = Timeline.objects.filter(project=project).first()
|
||||
if timeline is not None and version.asset_id:
|
||||
segment_asset_ids = list(segment.versions.values_list("asset_id", flat=True))
|
||||
timeline.clips.filter(asset_id__in=segment_asset_ids).exclude(asset_id=version.asset_id).update(
|
||||
asset_id=version.asset_id, trim_start_ms=0, trim_end_ms=None
|
||||
)
|
||||
return Response(ProjectSerializer(project).data)
|
||||
|
||||
@action(detail=True, methods=["post"], url_path="generate-storyboard")
|
||||
def generate_storyboard_action(self, request, pk=None):
|
||||
"""异步故事板·提交:快速创建版本(不在此生图、不推进阶段)。前端随后轮询 poll-storyboard 逐帧生成。"""
|
||||
require_worker()
|
||||
project = self.get_object()
|
||||
storyboard = submit_storyboard(project=project, user=request.user, prompt=request.data.get("prompt", ""))
|
||||
stage, _ = ProjectStage.objects.get_or_create(project=project, stage=ProjectStage.Stage.STORYBOARD)
|
||||
@@ -206,15 +323,17 @@ class ProjectViewSet(TeamScopedViewSetMixin, ModelViewSet):
|
||||
|
||||
@action(detail=True, methods=["post"], url_path="submit-video-segment")
|
||||
def submit_video_segment_action(self, request, pk=None):
|
||||
# 前置闸:无 worker 禁止提交——提交到 ARK 后若无人轮询,结果悬在云端、预扣额度冻结。
|
||||
require_worker()
|
||||
project = self.get_object()
|
||||
segment = VideoSegment.objects.get(project=project, id=request.data.get("video_segment_id"))
|
||||
submit_video_segment(video_segment=segment, user=request.user, prompt=request.data.get("prompt", ""))
|
||||
# 有 Celery worker 时由它自动轮询;无 worker(本机 dev)则前端驱动 poll-video-segment。
|
||||
# 队列不可用不应让提交 500——已提交到 ARK,轮询是次要路径。
|
||||
# worker 在线由上面的闸保证;此处入队失败只剩极小窗口(刚提交完 broker 闪断),
|
||||
# 任务已在 ARK,前端轮询仍可兜底取回,故不让提交 500。
|
||||
try:
|
||||
poll_video_segment_task.apply_async(args=[str(segment.id)], countdown=30)
|
||||
except Exception: # noqa: BLE001
|
||||
logger.warning("poll_video_segment_task enqueue failed; relying on client polling", exc_info=True)
|
||||
logger.error("poll_video_segment_task enqueue failed; relying on client polling", exc_info=True)
|
||||
return Response(ProjectSerializer(project).data, status=status.HTTP_202_ACCEPTED)
|
||||
|
||||
@action(detail=True, methods=["post"], url_path="poll-video-segment")
|
||||
@@ -345,6 +464,24 @@ class ProjectViewSet(TeamScopedViewSetMixin, ModelViewSet):
|
||||
BgmTrack.objects.create(timeline=timeline, asset=asset, volume=max(0, min(100, volume)), start_ms=0)
|
||||
return Response(ProjectSerializer(self.get_object()).data, status=status.HTTP_201_CREATED)
|
||||
|
||||
# ── Stage 5 · 旁白配音(TTS):每镜旁白合成语音资产,导出时作为人声轨混在 BGM 之上 ──
|
||||
@action(detail=True, methods=["get", "post"], url_path="generate-voiceover")
|
||||
def generate_voiceover_action(self, request, pk=None):
|
||||
project = self.get_object()
|
||||
if request.method == "GET":
|
||||
return Response({"voices": VOICEOVER_VOICES, "default_voice": DEFAULT_VOICEOVER_VOICE})
|
||||
try:
|
||||
voiceover = synthesize_project_voiceover(
|
||||
project=project,
|
||||
user=request.user,
|
||||
items=request.data.get("items") or [],
|
||||
voice_type=request.data.get("voice_type") or DEFAULT_VOICEOVER_VOICE,
|
||||
speed_ratio=float(request.data.get("speed_ratio") or 1.0),
|
||||
)
|
||||
except (TtsNotConfigured, ValueError) as exc:
|
||||
return Response({"detail": str(exc)}, status=status.HTTP_400_BAD_REQUEST)
|
||||
return Response({"voiceover": voiceover}, status=status.HTTP_201_CREATED)
|
||||
|
||||
@action(detail=True, methods=["post", "put"], url_path="save-timeline")
|
||||
@transaction.atomic
|
||||
def save_timeline_action(self, request, pk=None):
|
||||
@@ -402,6 +539,27 @@ class ProjectViewSet(TeamScopedViewSetMixin, ModelViewSet):
|
||||
metadata["transition"] = {"type": str(data["transition"].get("type", "none"))}
|
||||
if isinstance(data.get("draft"), dict):
|
||||
metadata["draft"] = data["draft"]
|
||||
if isinstance(data.get("voiceover"), dict):
|
||||
vo_patch = data["voiceover"]
|
||||
existing = metadata.get("voiceover")
|
||||
if vo_patch.get("clear"):
|
||||
metadata.pop("voiceover", None)
|
||||
elif isinstance(existing, dict):
|
||||
if vo_patch.get("enabled") is not None:
|
||||
existing["enabled"] = bool(vo_patch["enabled"])
|
||||
# 拖动字幕块后回写每句语音的句内起点(按 asset id 对位,只动 offset_ms)
|
||||
if isinstance(vo_patch.get("items"), list):
|
||||
offsets = {}
|
||||
for it in vo_patch["items"]:
|
||||
if isinstance(it, dict) and it.get("asset") is not None and it.get("offset_ms") is not None:
|
||||
try:
|
||||
offsets[str(it["asset"])] = max(0, int(it["offset_ms"]))
|
||||
except (TypeError, ValueError):
|
||||
continue
|
||||
for it in existing.get("items") or []:
|
||||
if str(it.get("asset")) in offsets:
|
||||
it["offset_ms"] = offsets[str(it.get("asset"))]
|
||||
metadata["voiceover"] = existing
|
||||
timeline.metadata = metadata
|
||||
timeline.save(update_fields=["metadata", "duration_seconds", "updated_at"])
|
||||
return Response(ProjectSerializer(self.get_object()).data)
|
||||
|
||||
Reference in New Issue
Block a user