模特库标签分页与全能创作收口:藏长视频、选择器分页
角色库导入打标签并支持筛选;模特库与全能创作角色/商品选择改为每页 20 条分页。临时限制成片 ≤60 秒,过滤对话里的超长时长选项,并收拢本地全能创作与后台用户相关修复。
This commit is contained in:
@@ -117,7 +117,7 @@ class AdminUserSerializer(serializers.ModelSerializer):
|
||||
class Meta:
|
||||
model = User
|
||||
fields = [
|
||||
"id", "username", "status", "is_platform_admin", "date_joined", "teams",
|
||||
"id", "username", "first_name", "status", "is_platform_admin", "date_joined", "teams",
|
||||
"balance", "wallet_team", "wallet_team_name", "wallet_shared",
|
||||
]
|
||||
read_only_fields = fields
|
||||
|
||||
@@ -871,7 +871,12 @@ class AdminUserProvisioningTests(TestCase):
|
||||
self.nc.force_authenticate(self.normal)
|
||||
|
||||
def _create(self, **overrides):
|
||||
payload = {"username": "created-user", "password": "strong-pass-1", "initial_credits": "500"}
|
||||
payload = {
|
||||
"username": "creat1",
|
||||
"display_name": "测试用户",
|
||||
"password": "strong-pass-1",
|
||||
"initial_credits": "500",
|
||||
}
|
||||
payload.update(overrides)
|
||||
return self.ac.post("/api/admin/users/", payload, format="json")
|
||||
|
||||
@@ -882,7 +887,8 @@ class AdminUserProvisioningTests(TestCase):
|
||||
self.assertFalse(r.data["wallet_shared"])
|
||||
self.assertEqual(r.data["teams"], []) # 个人团队是钱包实现细节,不当「所属团队」露出
|
||||
|
||||
user = User.objects.get(username="created-user")
|
||||
user = User.objects.get(username="creat1")
|
||||
self.assertEqual(user.first_name, "测试用户")
|
||||
team = Team.objects.get(owner=user)
|
||||
self.assertTrue(team.is_personal)
|
||||
self.assertEqual(team.credit_account.balance, Decimal("500"))
|
||||
@@ -892,22 +898,30 @@ class AdminUserProvisioningTests(TestCase):
|
||||
|
||||
def test_created_user_can_login(self):
|
||||
self._create()
|
||||
r = APIClient().post("/api/auth/login/", {"username": "created-user", "password": "strong-pass-1"}, format="json")
|
||||
r = APIClient().post("/api/auth/login/", {"username": "creat1", "password": "strong-pass-1"}, format="json")
|
||||
self.assertEqual(r.status_code, 200)
|
||||
|
||||
def test_create_user_validation_and_permission(self):
|
||||
self.assertEqual(self.nc.post("/api/admin/users/", {}, format="json").status_code, 403)
|
||||
self.assertEqual(self._create(username="").status_code, 400)
|
||||
self.assertEqual(self._create(display_name="").status_code, 400)
|
||||
self.assertEqual(self._create(username="abc").status_code, 400) # 不足 6 位
|
||||
self.assertEqual(self._create(username="abc-12").status_code, 400) # 含非法字符
|
||||
self.assertEqual(self._create(username="abcdefg").status_code, 400) # 超过 6 位
|
||||
self.assertEqual(self._create(password="short").status_code, 400)
|
||||
self.assertEqual(self._create(initial_credits="-5").status_code, 400)
|
||||
self.assertEqual(self._create(username="normal").status_code, 400) # 重名
|
||||
self.assertEqual(User.objects.filter(username="created-user").count(), 0)
|
||||
# 准备一个 6 位账号占用,验证重名(大小写不敏感)
|
||||
User.objects.create_user(username="taken1", password="x")
|
||||
self.assertEqual(self._create(username="taken1").status_code, 400)
|
||||
self.assertEqual(self._create(username="TAKEN1").status_code, 400)
|
||||
self.assertEqual(User.objects.filter(username="creat1").count(), 0)
|
||||
|
||||
def test_create_user_without_initial_credits(self):
|
||||
r = self._create(initial_credits="")
|
||||
self.assertEqual(r.status_code, 201)
|
||||
self.assertEqual(r.data["balance"], "0.0000")
|
||||
team = Team.objects.get(owner__username="created-user")
|
||||
team = Team.objects.get(owner__username="creat1")
|
||||
self.assertEqual(team.name, "测试用户")
|
||||
self.assertFalse(CreditLedger.objects.filter(team=team).exists()) # 0 积分不落空流水
|
||||
|
||||
def test_adjust_user_credit(self):
|
||||
|
||||
@@ -269,31 +269,43 @@ def admin_users(request):
|
||||
qs = qs.filter(status=st)
|
||||
search = (request.query_params.get("search") or "").strip()
|
||||
if search:
|
||||
qs = qs.filter(username__icontains=search)
|
||||
qs = qs.filter(Q(username__icontains=search) | Q(first_name__icontains=search))
|
||||
paginator = DefaultPagination()
|
||||
page = paginator.paginate_queryset(qs, request)
|
||||
return paginator.get_paginated_response(AdminUserSerializer(page, many=True).data)
|
||||
|
||||
|
||||
def _admin_create_user(request):
|
||||
"""后台直建用户:{username, password, initial_credits?}。
|
||||
"""后台直建用户:{username, display_name?, password, initial_credits?}。
|
||||
|
||||
username = 登录账号(仅英文+数字,固定 6 位);display_name = 展示用用户名(可中文)。
|
||||
积分账户是 OneToOne 挂 Team 的,所以这里给新用户配一个 is_personal 的一人团队当专属钱包 ——
|
||||
计费链路(预留/实扣/流水/额度)一行不用改,后台却能按「用户」视角管人和钱。
|
||||
后续要恢复多人协作,把人拉进同一个团队即可,不需要迁数据。"""
|
||||
import re
|
||||
|
||||
username = str(request.data.get("username") or "").strip()
|
||||
display_name = str(
|
||||
request.data.get("display_name")
|
||||
or request.data.get("name")
|
||||
or ""
|
||||
).strip()
|
||||
password = str(request.data.get("password") or "").strip()
|
||||
if not display_name:
|
||||
return Response({"display_name": ["请填写用户名"]}, status=status.HTTP_400_BAD_REQUEST)
|
||||
if len(display_name) > 64:
|
||||
return Response({"display_name": ["用户名不能超过 64 字"]}, status=status.HTTP_400_BAD_REQUEST)
|
||||
if not username:
|
||||
return Response({"username": ["请填写用户名"]}, status=status.HTTP_400_BAD_REQUEST)
|
||||
if len(username) > 150:
|
||||
return Response({"username": ["用户名不能超过 150 字"]}, status=status.HTTP_400_BAD_REQUEST)
|
||||
try:
|
||||
for validator in User._meta.get_field("username").validators:
|
||||
validator(username)
|
||||
except DjangoValidationError as exc:
|
||||
return Response({"username": list(exc.messages)}, status=status.HTTP_400_BAD_REQUEST)
|
||||
if User.objects.filter(username=username).exists():
|
||||
return Response({"username": ["该用户名已存在"]}, status=status.HTTP_400_BAD_REQUEST)
|
||||
return Response({"username": ["请填写登录账号"]}, status=status.HTTP_400_BAD_REQUEST)
|
||||
if not re.fullmatch(r"[A-Za-z0-9]{6}", username):
|
||||
return Response(
|
||||
{"username": ["登录账号须为 6 位英文或数字,不能含其他字符"]},
|
||||
status=status.HTTP_400_BAD_REQUEST,
|
||||
)
|
||||
# 统一小写入库,避免 Abc123 / abc123 被当成两个账号
|
||||
username = username.lower()
|
||||
if User.objects.filter(username__iexact=username).exists():
|
||||
return Response({"username": ["该登录账号已存在,不能重复"]}, status=status.HTTP_400_BAD_REQUEST)
|
||||
if len(password) < 8:
|
||||
return Response({"password": ["密码至少 8 位"]}, status=status.HTTP_400_BAD_REQUEST)
|
||||
initial, err = _parse_points(request.data.get("initial_credits") or 0, field="initial_credits", allow_negative=False)
|
||||
@@ -301,8 +313,13 @@ def _admin_create_user(request):
|
||||
return err
|
||||
|
||||
with transaction.atomic():
|
||||
user = User.objects.create_user(username=username, password=password)
|
||||
team = Team.objects.create(name=username, owner=user, is_personal=True)
|
||||
user = User.objects.create_user(
|
||||
username=username,
|
||||
password=password,
|
||||
first_name=display_name,
|
||||
)
|
||||
# 个人团队名用展示名,侧栏第一行显示用户名、第二行显示登录账号
|
||||
team = Team.objects.create(name=display_name, owner=user, is_personal=True)
|
||||
TeamMember.objects.create(team=team, user=user, role=TeamMember.Role.OWNER)
|
||||
CreditAccount.objects.create(team=team, balance=initial)
|
||||
if initial > 0:
|
||||
@@ -323,7 +340,11 @@ def _admin_create_user(request):
|
||||
target_type="user",
|
||||
target_id=user.id,
|
||||
target_name=user.username,
|
||||
after={"initial_credits": str(initial), "wallet_team": str(team.id)},
|
||||
after={
|
||||
"initial_credits": str(initial),
|
||||
"wallet_team": str(team.id),
|
||||
"display_name": display_name,
|
||||
},
|
||||
)
|
||||
return Response(AdminUserSerializer(_admin_user_qs().get(id=user.id)).data, status=status.HTTP_201_CREATED)
|
||||
|
||||
|
||||
@@ -181,105 +181,242 @@ def _sync_segmented_video_message(message: CreationMessage) -> bool:
|
||||
一张结果卡。任一片段先完成就先写回 GENERATING 卡供用户预览;所有分段完成后才转成
|
||||
可合并的结果卡,绝不在此处触发 ffmpeg。
|
||||
"""
|
||||
from .free_video import finalize_free_video
|
||||
from django.conf import settings
|
||||
|
||||
from .free_video import IN_FLIGHT_STATUSES, finalize_free_video, submit_free_video
|
||||
from .models import AITask
|
||||
|
||||
payload = message.payload or {}
|
||||
ids = [str(value) for value in payload.get("task_ids") or [] if value]
|
||||
if not ids:
|
||||
# 团队级锁:同一团队若有两支长视频同时回填,不能都按同一份剩余并发额度补交。
|
||||
lock_key = f"omni:segment-schedule:{message.conversation.team_id}"
|
||||
if not cache.add(lock_key, "1", timeout=90):
|
||||
return False
|
||||
by_id = {
|
||||
str(task.id): task
|
||||
for task in AITask.objects.filter(id__in=ids).select_related("model_config")
|
||||
}
|
||||
tasks = [by_id.get(task_id) for task_id in ids]
|
||||
if any(task is None for task in tasks):
|
||||
fail_generating_message(message, "分段视频任务不完整,请重新生成。")
|
||||
return True
|
||||
try:
|
||||
message.refresh_from_db()
|
||||
payload = dict(message.payload or {})
|
||||
segments = [dict(item) for item in payload.get("segments") or [] if isinstance(item, dict)]
|
||||
ids = [str(value) for value in payload.get("task_ids") or [] if value]
|
||||
if not ids or not segments:
|
||||
return False
|
||||
by_id = {
|
||||
str(task.id): task
|
||||
for task in AITask.objects.filter(id__in=ids).select_related("model_config")
|
||||
}
|
||||
tasks = [by_id.get(task_id) for task_id in ids]
|
||||
if any(task is None for task in tasks):
|
||||
fail_generating_message(message, "分段视频任务不完整,请重新生成。")
|
||||
return True
|
||||
|
||||
# 本地没有 worker 时,用户的会话轮询本身即可推进每个片段;已在 worker 中的任务则幂等返回。
|
||||
refreshed = []
|
||||
for task in tasks:
|
||||
assert task is not None
|
||||
if task.status in (AITask.Status.SUBMITTED, AITask.Status.POLLING):
|
||||
try:
|
||||
task = finalize_free_video(task=task)
|
||||
except Exception: # noqa: BLE001 - 单段网络抖动不能使整组直接失败
|
||||
# 本地没有 worker 时,会话轮询也能推进已提交片段;已在 worker 中则幂等返回。
|
||||
refreshed = []
|
||||
for task in tasks:
|
||||
assert task is not None
|
||||
if task.status in (AITask.Status.SUBMITTED, AITask.Status.POLLING):
|
||||
try:
|
||||
task = finalize_free_video(task=task)
|
||||
except Exception: # noqa: BLE001 - 单段网络抖动不能使整组直接失败
|
||||
task.refresh_from_db()
|
||||
else:
|
||||
task.refresh_from_db()
|
||||
else:
|
||||
task.refresh_from_db()
|
||||
refreshed.append(task)
|
||||
refreshed.append(task)
|
||||
|
||||
failed = next(
|
||||
(task for task in refreshed if task.status in (AITask.Status.FAILED, AITask.Status.CANCELLED)),
|
||||
None,
|
||||
)
|
||||
if failed is not None:
|
||||
from .generation_errors import public_error_for_task
|
||||
segment_by_task = {
|
||||
str(item.get("task_id")): int(item.get("index") or index)
|
||||
for index, item in enumerate(segments, start=1)
|
||||
if item.get("task_id")
|
||||
}
|
||||
failed = next(
|
||||
(task for task in refreshed if task.status in (AITask.Status.FAILED, AITask.Status.CANCELLED)),
|
||||
None,
|
||||
)
|
||||
|
||||
number = next((index + 1 for index, task in enumerate(refreshed) if task.id == failed.id), 1)
|
||||
public_error = public_error_for_task(failed, operation="video_generate")
|
||||
detail = public_error.fallback_message if public_error else (failed.error_message or "请重试")
|
||||
fail_generating_message(message, f"第 {number} 段生成失败:{detail}")
|
||||
assets: list[dict] = []
|
||||
completed_segments = 0
|
||||
for task in refreshed:
|
||||
if task.status != AITask.Status.SUCCEEDED:
|
||||
continue
|
||||
task_assets = _assets_from_task(task)
|
||||
if not task_assets:
|
||||
continue
|
||||
segment_index = segment_by_task.get(str(task.id), completed_segments + 1)
|
||||
completed_segments += 1
|
||||
for asset in task_assets:
|
||||
assets.append({**asset, "label": f"第 {segment_index} 段", "segment_index": segment_index})
|
||||
assets.sort(key=lambda item: int(item.get("segment_index") or 0))
|
||||
|
||||
if failed is not None:
|
||||
from .generation_errors import public_error_for_task
|
||||
|
||||
number = segment_by_task.get(str(failed.id), 1)
|
||||
public_error = public_error_for_task(failed, operation="video_generate")
|
||||
detail = public_error.fallback_message if public_error else (failed.error_message or "请重试")
|
||||
error_text = f"第 {number} 段生成失败:{detail}"
|
||||
# 取消仍在飞的其它分段,避免失败后继续扣费/刷进度。
|
||||
for task in refreshed:
|
||||
if task.id == failed.id:
|
||||
continue
|
||||
if task.status in IN_FLIGHT_STATUSES:
|
||||
from apps.billing.services.ledger import release_credit
|
||||
|
||||
task.status = AITask.Status.CANCELLED
|
||||
task.error_message = "同组其它分段已失败,已取消"
|
||||
task.completed_at = timezone.now()
|
||||
task.save(update_fields=["status", "error_message", "completed_at", "updated_at"])
|
||||
if task.credit_reservation is not None:
|
||||
try:
|
||||
release_credit(
|
||||
reservation=task.credit_reservation,
|
||||
reason="同组分段失败,取消未完成片段",
|
||||
)
|
||||
except Exception: # noqa: BLE001
|
||||
pass
|
||||
if assets:
|
||||
first = next((task for task in refreshed if task.status == AITask.Status.SUCCEEDED), failed)
|
||||
finish_generating_message(
|
||||
message,
|
||||
assets=assets,
|
||||
meta={
|
||||
**_meta_from_task(first, message),
|
||||
"kind": "video_segments",
|
||||
"segment_count": len(segments),
|
||||
"total_duration": payload.get("total_duration") or "",
|
||||
"needs_merge": False,
|
||||
"partial_failure": True,
|
||||
"failed_segment_index": number,
|
||||
"error": error_text,
|
||||
"segments": segments,
|
||||
"task_ids": ids,
|
||||
},
|
||||
)
|
||||
conversation = message.conversation
|
||||
conversation.status = CreationConversation.Status.FAILED
|
||||
conversation.save(update_fields=["status", "updated_at"])
|
||||
append_message(
|
||||
conversation,
|
||||
role="assistant",
|
||||
kind=CreationMessage.Kind.ERROR,
|
||||
text=f"{error_text}。已保留成功片段供预览,未合并成片;可重新确认方案再生成。",
|
||||
)
|
||||
else:
|
||||
fail_generating_message(message, error_text)
|
||||
return True
|
||||
|
||||
# 释放出来的并发槽位自动补交下一批。首批与后续批次统一从 generation_spec 派生,
|
||||
# 因而复用同一份人物/商品参考、完整脚本和 seed。
|
||||
pending = [item for item in segments if not item.get("task_id")]
|
||||
if pending:
|
||||
in_flight = AITask.objects.filter(
|
||||
team=message.conversation.team,
|
||||
task_type=AITask.Type.FREE_VIDEO,
|
||||
status__in=IN_FLIGHT_STATUSES,
|
||||
).count()
|
||||
slots = max(0, int(getattr(settings, "FREE_VIDEO_MAX_CONCURRENT", 3)) - in_flight)
|
||||
spec = payload.get("generation_spec") if isinstance(payload.get("generation_spec"), dict) else {}
|
||||
if slots and not spec:
|
||||
fail_generating_message(message, "长视频续段参数不完整,请重新生成。")
|
||||
return True
|
||||
if slots:
|
||||
from .creation_agent import build_segment_video_submit
|
||||
|
||||
total_duration = int(payload.get("total_duration") or 0)
|
||||
submitted_now = 0
|
||||
try:
|
||||
for segment in pending[:slots]:
|
||||
task = submit_free_video(
|
||||
team=message.conversation.team,
|
||||
user=message.conversation.created_by,
|
||||
params=build_segment_video_submit(spec, segment, total_duration),
|
||||
)
|
||||
task_id = str(task.id)
|
||||
segment["task_id"] = task_id
|
||||
ids.append(task_id)
|
||||
request_payload = dict(task.request_payload or {})
|
||||
marker = dict(request_payload.get("omni_segment") or {})
|
||||
marker.update({
|
||||
"index": int(segment.get("index") or 0),
|
||||
"start": int(segment.get("start") or 0),
|
||||
"end": int(segment.get("end") or 0),
|
||||
"total_duration": total_duration,
|
||||
"message_id": str(message.id),
|
||||
})
|
||||
request_payload["omni_segment"] = marker
|
||||
task.request_payload = request_payload
|
||||
task.save(update_fields=["request_payload", "updated_at"])
|
||||
submitted_now += 1
|
||||
except ValueError as exc:
|
||||
if not submitted_now:
|
||||
fail_generating_message(message, f"后续分段提交失败:{exc}")
|
||||
return True
|
||||
# 逐段提交期间并发槽位被别的请求占用时,保留本轮已提交任务;余下片段下轮再补。
|
||||
payload["task_ids"] = ids
|
||||
payload["segments"] = segments
|
||||
|
||||
progress = {
|
||||
**payload,
|
||||
"assets": assets,
|
||||
"completed_segment_count": completed_segments,
|
||||
"submitted_segment_count": len(ids),
|
||||
}
|
||||
all_submitted = all(item.get("task_id") for item in segments)
|
||||
all_succeeded = all(task.status == AITask.Status.SUCCEEDED for task in refreshed)
|
||||
# 刚补交的新任务不在 refreshed 内,因此必须同时校验提交数,不能提前完成结果卡。
|
||||
if not all_submitted or len(refreshed) != len(segments) or not all_succeeded:
|
||||
if progress != payload:
|
||||
message.payload = progress
|
||||
message.save(update_fields=["payload", "updated_at"])
|
||||
return True
|
||||
return False
|
||||
|
||||
# 状态已成功但资产尚未落库时继续等待,避免最终成片缺段。
|
||||
if completed_segments != len(segments):
|
||||
if progress != payload:
|
||||
message.payload = progress
|
||||
message.save(update_fields=["payload", "updated_at"])
|
||||
return True
|
||||
return False
|
||||
|
||||
first = refreshed[0]
|
||||
finish_generating_message(
|
||||
message,
|
||||
assets=assets,
|
||||
meta={
|
||||
**_meta_from_task(first, message),
|
||||
"kind": "video_segments",
|
||||
"segment_count": len(segments),
|
||||
"total_duration": payload.get("total_duration") or "",
|
||||
"needs_merge": True,
|
||||
"auto_merge": True,
|
||||
"segments": segments,
|
||||
"task_ids": ids,
|
||||
},
|
||||
)
|
||||
message.refresh_from_db()
|
||||
# 全部分段成功后由平台自动合并,不再等用户点「合并成片」。
|
||||
try:
|
||||
merge_task, _generating = start_segmented_video_merge(
|
||||
conversation=message.conversation,
|
||||
message=message,
|
||||
user=message.conversation.created_by,
|
||||
)
|
||||
try:
|
||||
from .tasks import merge_omni_video_segments_task
|
||||
|
||||
merge_omni_video_segments_task.apply_async(args=[str(merge_task.id)])
|
||||
except Exception: # noqa: BLE001
|
||||
import logging
|
||||
|
||||
logging.getLogger(__name__).warning(
|
||||
"omni auto video merge enqueue failed for %s", merge_task.id, exc_info=True
|
||||
)
|
||||
except ValueError:
|
||||
# 已在合并中或状态不允许时忽略,避免重复入队。
|
||||
pass
|
||||
return True
|
||||
assets: list[dict] = []
|
||||
completed_segments = 0
|
||||
for index, task in enumerate(refreshed, start=1):
|
||||
if task.status != AITask.Status.SUCCEEDED:
|
||||
continue
|
||||
task_assets = _assets_from_task(task)
|
||||
# 上游状态先成功、资产稍后才落库时,保留 GENERATING,下一轮再展示这段。
|
||||
if not task_assets:
|
||||
continue
|
||||
completed_segments += 1
|
||||
for asset in task_assets:
|
||||
assets.append({**asset, "label": f"第 {index} 段", "segment_index": index})
|
||||
|
||||
# 先完成的片段必须立刻回到前端,不能等另一段慢任务一起完成才出现。
|
||||
# 保持同一条 GENERATING 消息,避免把一个 60 秒视频拆成多条对话消息。
|
||||
if not all(task.status == AITask.Status.SUCCEEDED for task in refreshed):
|
||||
progress = {
|
||||
**payload,
|
||||
"assets": assets,
|
||||
"completed_segment_count": completed_segments,
|
||||
}
|
||||
if progress != payload:
|
||||
message.payload = progress
|
||||
message.save(update_fields=["payload", "updated_at"])
|
||||
return True
|
||||
return False
|
||||
|
||||
# 任务都成功但有片段的资产还在落库,继续保持生成中,避免最终结果缺片。
|
||||
if completed_segments != len(refreshed):
|
||||
progress = {
|
||||
**payload,
|
||||
"assets": assets,
|
||||
"completed_segment_count": completed_segments,
|
||||
}
|
||||
if progress != payload:
|
||||
message.payload = progress
|
||||
message.save(update_fields=["payload", "updated_at"])
|
||||
return True
|
||||
return False
|
||||
|
||||
first = refreshed[0]
|
||||
finish_generating_message(
|
||||
message,
|
||||
assets=assets,
|
||||
meta={
|
||||
**_meta_from_task(first, message),
|
||||
"kind": "video_segments",
|
||||
"segment_count": len(refreshed),
|
||||
"total_duration": payload.get("total_duration") or "",
|
||||
"needs_merge": True,
|
||||
"segments": payload.get("segments") or [],
|
||||
},
|
||||
)
|
||||
return True
|
||||
finally:
|
||||
cache.delete(lock_key)
|
||||
|
||||
|
||||
def start_segmented_video_merge(*, conversation: CreationConversation, message: CreationMessage, user):
|
||||
"""用户明确点击后才建立合并任务;这之前绝不下载片段或调用 ffmpeg。"""
|
||||
"""建立合并任务(全部分段成功后由平台自动调用;也可由旧接口手动触发)。合并执行仍只在 run_segmented_video_merge。"""
|
||||
from .models import AITask
|
||||
|
||||
payload = dict(message.payload or {})
|
||||
@@ -406,9 +543,16 @@ def sync_generating_messages(conversation: CreationConversation) -> int:
|
||||
def sync_generating_for_task(task) -> int:
|
||||
"""worker / poll 终态后:只扫挂在这个任务上的 GENERATING。失败不能向外抛。"""
|
||||
try:
|
||||
from django.db.models import Q
|
||||
|
||||
marker = (task.request_payload or {}).get("omni_segment") or {}
|
||||
aggregate_message_id = str(marker.get("message_id") or "") if isinstance(marker, dict) else ""
|
||||
lookup = Q(task=task)
|
||||
if aggregate_message_id:
|
||||
lookup |= Q(id=aggregate_message_id)
|
||||
pending = list(
|
||||
CreationMessage.objects.filter(
|
||||
task=task, kind=CreationMessage.Kind.GENERATING,
|
||||
lookup, kind=CreationMessage.Kind.GENERATING,
|
||||
).select_related("conversation", "task", "task__model_config")
|
||||
)
|
||||
return sum(1 for message in pending if sync_generating_message(message))
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -128,10 +128,15 @@ _PLOT_TWIST_DEPTH_BY_VALUE = {item["value"]: item for item in PLOT_TWIST_STORY_D
|
||||
|
||||
|
||||
def plot_twist_story_depth(value: str) -> dict | None:
|
||||
"""兼容卡片 value、展示文案和自然语言里带的秒数。"""
|
||||
"""兼容卡片 value、展示文案和自然语言里带的秒数。
|
||||
|
||||
180 秒长视频入口已临时关闭:历史「180s」选择落到 60 秒档。
|
||||
"""
|
||||
raw = str(value or "").strip()
|
||||
if raw in _PLOT_TWIST_DEPTH_BY_VALUE:
|
||||
return _PLOT_TWIST_DEPTH_BY_VALUE[raw]
|
||||
if "180" in raw or "90" in raw or "120" in raw:
|
||||
return _PLOT_TWIST_DEPTH_BY_VALUE["60s"]
|
||||
if "60" in raw:
|
||||
return _PLOT_TWIST_DEPTH_BY_VALUE["60s"]
|
||||
if "30" in raw:
|
||||
@@ -170,6 +175,18 @@ def plot_twist_story_contract(value: str) -> str:
|
||||
"56-60 秒回扣商品价值并自然转化。商品可前置为伏笔但必须在后半段真正改变结局;"
|
||||
"禁止用重复对白、无意义空镜或硬插卖点填满时长。"
|
||||
)
|
||||
if depth["value"] == "180s":
|
||||
return (
|
||||
"【剧情反转带货·180秒三幕完整故事·强制执行】全片按三个连续 60 秒章节推进,不能把三条短片简单拼接。"
|
||||
"第一章 0-60 秒:0-8 秒以结果预告或关系冲突抓人,8-28 秒建立人物目标与现实阻力,"
|
||||
"28-48 秒埋入商品/SKU 与一个可见卖点证据,48-60 秒用第一次选择或失败把行动推入下一章。"
|
||||
"第二章 60-120 秒:承接上一章未完成动作,扩大矛盾并安排一次错误尝试;商品通过正常使用过程提供新的证据,"
|
||||
"在 105-120 秒完成中段转折,但不得提前总结或重复开场。"
|
||||
"第三章 120-180 秒:让前两章伏笔、人物选择和商品证据共同触发主要反转,165 秒前完成结果验证,"
|
||||
"165-176 秒释放人物情绪并回扣核心卖点,176-180 秒只做一次自然行动引导。"
|
||||
"每章都要推进新的因果关系;同一角色、服装、商品/SKU、场景空间和光线跨六个 30 秒分段连续,"
|
||||
"禁止重复对白、重复卖点、空镜凑时长或在章节边界换人换商品。"
|
||||
)
|
||||
return (
|
||||
"【剧情反转带货·智能推荐】先根据商品卖点、已有素材和剧情空间推荐 15、30 或 60 秒之一并说明理由;"
|
||||
"随后必须让用户选择实际时长,未确认前不得写剧情方向、策略、方案或出片指令。"
|
||||
@@ -249,7 +266,7 @@ def apply_plot_twist_direction_contract(prompt: str, *, title: str = "", conflic
|
||||
VIDEO_PRESET_WORKFLOWS: dict[str, str] = {
|
||||
"痛点解决演示": "先确认商品真实解决的具体问题与正常用法;生成前核对痛点、过程和结果都有可见证据。",
|
||||
"真实使用演示": "优先核对商品真实用途、关键步骤和可证实卖点;不为氛围而增加不合理测试。",
|
||||
"剧情反转带货": "先让用户选故事深度(15秒快节奏反转、30秒轻剧情带货、60秒完整短剧带货或智能推荐),再给三个与时长匹配的剧情方向;商品必须成为解决问题、解除误会、证明事实、回收伏笔或完成翻盘的关键。",
|
||||
"剧情反转带货": "先让用户选故事深度(15秒快节奏反转、30秒轻剧情带货、60秒完整短剧带货或智能推荐;暂不开放更长时长),再给三个与时长匹配的剧情方向;商品必须成为解决问题、解除误会、证明事实、回收伏笔或完成翻盘的关键。",
|
||||
"商品拟人广告": "先确认商品外观和性格表达方式;默认无脸拟人,商品保持真实完整,台词用画外声。",
|
||||
"达人口播种草": "优先确认人物、多人出镜关系、真实体验和主卖点;生成前核对口播字数能在时长内说完,每个卖点都有画面证明。",
|
||||
"商品图一键成片": "优先从商品参考图锁定外观;自动补场景和动作,但不替换或改变用户商品图里的结构、颜色和包装。",
|
||||
|
||||
@@ -44,6 +44,23 @@ FREE_VIDEO_MODELS = {
|
||||
"doubao-seedance-2-0-mini-260615",
|
||||
}
|
||||
HIGH_RES_MODEL = "doubao-seedance-2-0-260128" # 1080p/4k 仅标准档(火山限制)
|
||||
# 火山对 seed 的硬上限是 int32 正数上限;超一点整条请求直接 InvalidParameter 被拒。
|
||||
VIDEO_SEED_MAX = 2147483647
|
||||
|
||||
|
||||
def normalize_video_seed(value) -> int:
|
||||
"""把任何来源的 seed 收进火山认的区间。-1 = 不指定(随机)。
|
||||
|
||||
越界不报错而是折回区间内:长视频各段共用同一个确定性 seed 保一致性,
|
||||
折算规则固定,同一个会话每次算出来仍是同一个值。
|
||||
"""
|
||||
try:
|
||||
seed = int(value)
|
||||
except (TypeError, ValueError):
|
||||
return -1
|
||||
if seed < 0:
|
||||
return -1
|
||||
return seed & VIDEO_SEED_MAX
|
||||
# 视频复刻固定 Seedance 2.5(单次最长 30 秒)。原片只用来提炼分镜,不传给火山。
|
||||
REPLACE_MODEL = "doubao-seedance-2-5-260628"
|
||||
# 出片时长的兜底上限。真实上限按 ModelConfig.metadata.durations 取(见 model_duration_range),
|
||||
@@ -129,6 +146,29 @@ def _refresh_processing_free_asset(free_asset: FreeAsset) -> bool:
|
||||
return False
|
||||
|
||||
|
||||
def _reference_asset(team, asset_id) -> Asset | None:
|
||||
"""出片参考图允许两种来源:团队自有资产,以及官方模特库的人像/三视图。
|
||||
|
||||
官方模特挂在平台团队名下,按 team 过滤会直接查不到 —— 这是模特库里明明存在、
|
||||
生成时却报「不存在或已被删除」的根因。放行范围与视频复刻(video_replace)保持一致。
|
||||
"""
|
||||
from django.db.models import Q
|
||||
|
||||
from apps.assets.models import Model as AssetModel
|
||||
|
||||
asset = Asset.objects.filter(id=asset_id, is_deleted=False).first()
|
||||
if asset is None:
|
||||
return None
|
||||
if asset.team_id == team.id:
|
||||
return asset
|
||||
is_official_model_asset = (
|
||||
AssetModel.objects.filter(is_official=True, is_deleted=False, purged_at__isnull=True)
|
||||
.filter(Q(portrait_asset=asset) | Q(triview_asset=asset))
|
||||
.exists()
|
||||
)
|
||||
return asset if is_official_model_asset else None
|
||||
|
||||
|
||||
def _guard_asset_reference(asset: Asset, label: str, ref_type: str = "") -> None:
|
||||
"""三库引用前的审核闸(模块4 · 4.2)。
|
||||
|
||||
@@ -317,7 +357,7 @@ def build_content_items(
|
||||
# 全能创作 resolve_refs / 前端常带 type=character|product 且 source 缺省为 upload,
|
||||
# 只要有 asset_id 就按 Asset 解析(含审核态 → asset://),避免直链分支把语义 type skipped。
|
||||
if ref.get("asset_id") and source in {"asset", "upload", ""}:
|
||||
asset = Asset.objects.filter(id=ref["asset_id"], team=team, is_deleted=False).first()
|
||||
asset = _reference_asset(team, ref["asset_id"])
|
||||
if asset is None:
|
||||
raise ValueError(f"素材「{label or '未命名'}」不存在或已被删除")
|
||||
_guard_asset_reference(asset, label, ref_type)
|
||||
@@ -480,10 +520,7 @@ def submit_free_video(*, team, user, params: dict) -> AITask:
|
||||
duration = int(params.get("duration") or 5)
|
||||
except (TypeError, ValueError):
|
||||
raise ValueError("时长参数无效")
|
||||
try:
|
||||
seed = int(params.get("seed") if params.get("seed") is not None else -1)
|
||||
except (TypeError, ValueError):
|
||||
seed = -1
|
||||
seed = normalize_video_seed(params.get("seed"))
|
||||
|
||||
if not prompt:
|
||||
raise ValueError("提示词不能为空")
|
||||
@@ -637,10 +674,7 @@ def start_pending_free_video(task: AITask) -> AITask:
|
||||
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
|
||||
seed = normalize_video_seed(payload.get("seed"))
|
||||
references = list(payload.get("references") or [])
|
||||
if feature == "video_replace" and payload.get("replace_mode") == "product":
|
||||
# 商品复刻:原片只用来提炼,绝不能进 Seedance,否则真人素材直接被火山拒。
|
||||
|
||||
@@ -337,6 +337,25 @@ def _character_triview_asset(team, portrait_id):
|
||||
)
|
||||
|
||||
|
||||
def _newest_asset(*assets):
|
||||
"""多张可用图里取最新一张。已删/空值直接跳过。"""
|
||||
living = [
|
||||
asset for asset in assets
|
||||
if asset is not None and not getattr(asset, "is_deleted", False)
|
||||
]
|
||||
if not living:
|
||||
return None
|
||||
return max(living, key=lambda asset: (asset.created_at, str(asset.id)))
|
||||
|
||||
|
||||
def _single_entity_asset(cover, *triviews):
|
||||
"""一个角色/商品只喂一张参考图:有三视图用最新三视图,没有才用封面/立绘。
|
||||
|
||||
封面和三视图衣服经常不一致;两张一起送,出片模型会当成多出来的角色或 SKU。
|
||||
"""
|
||||
return _newest_asset(*triviews) or cover
|
||||
|
||||
|
||||
def _append_ref(resolved: ResolvedRefs, asset, type_: str, label: str) -> bool:
|
||||
entry = _asset_reference(asset, type_, label)
|
||||
if not entry:
|
||||
@@ -386,12 +405,17 @@ def resolve_refs(team, refs: list[dict]) -> ResolvedRefs:
|
||||
resolved.missing.append(ref)
|
||||
continue
|
||||
resolved.facts.append(product_facts_text(product))
|
||||
chosen = _single_entity_asset(
|
||||
product.cover_asset,
|
||||
_product_triview_asset(product),
|
||||
)
|
||||
if chosen is not None and _append_ref(resolved, chosen, "product", product.title):
|
||||
continue
|
||||
entry = _product_reference(product)
|
||||
if entry:
|
||||
resolved.references.append(entry)
|
||||
triview = _product_triview_asset(product)
|
||||
if triview is not None:
|
||||
_append_ref(resolved, triview, "product", f"{product.title}三视图")
|
||||
elif chosen is None:
|
||||
resolved.missing.append(ref)
|
||||
continue
|
||||
if type_ == "model":
|
||||
model = Model.objects.filter(
|
||||
@@ -413,13 +437,14 @@ def resolve_refs(team, refs: list[dict]) -> ResolvedRefs:
|
||||
person_lines.extend(meta_lines)
|
||||
if desc or meta_lines:
|
||||
resolved.facts.append("\n".join(person_lines))
|
||||
# 形象图和三视图都带上(有就带,没有不加、不报错)。三视图锁脸更稳。
|
||||
added = False
|
||||
if _append_ref(resolved, model.portrait_asset, "model", model.name):
|
||||
added = True
|
||||
if _append_ref(resolved, model.triview_asset, "model", f"{model.name}三视图"):
|
||||
added = True
|
||||
if not added:
|
||||
# 一人一张图:有三视图只用最新三视图。形象图和三视图衣服经常不一致,
|
||||
# 两张一起送会被当成多出来的角色。
|
||||
chosen = _single_entity_asset(
|
||||
model.portrait_asset,
|
||||
model.triview_asset,
|
||||
_character_triview_asset(team, getattr(model.portrait_asset, "id", None)),
|
||||
)
|
||||
if not _append_ref(resolved, chosen, "model", model.name):
|
||||
resolved.missing.append(ref)
|
||||
continue
|
||||
asset = Asset.objects.filter(
|
||||
@@ -435,15 +460,14 @@ def resolve_refs(team, refs: list[dict]) -> ResolvedRefs:
|
||||
person_lines.extend(meta_lines)
|
||||
if desc or meta_lines:
|
||||
resolved.facts.append("\n".join(person_lines))
|
||||
added = _append_ref(resolved, asset, type_, asset.name)
|
||||
# 角色/模特立绘若有配对三视图,一并带给出片;没有就跳过,不要当成引用失败。
|
||||
if type_ in {"character", "model", "asset"}:
|
||||
portrait_id = asset.id
|
||||
# 用户 @ 的就是三视图本身时,不再反查。
|
||||
if asset.category != Asset.Category.TRI_VIEW:
|
||||
paired = _character_triview_asset(team, portrait_id)
|
||||
if paired is not None:
|
||||
_append_ref(resolved, paired, type_, f"{asset.name}三视图")
|
||||
if type_ in {"character", "model", "asset"} and asset.category != Asset.Category.TRI_VIEW:
|
||||
chosen = _single_entity_asset(
|
||||
asset,
|
||||
_character_triview_asset(team, asset.id),
|
||||
)
|
||||
else:
|
||||
chosen = asset
|
||||
added = _append_ref(resolved, chosen, type_, asset.name)
|
||||
if not added:
|
||||
resolved.missing.append(ref)
|
||||
|
||||
|
||||
@@ -107,7 +107,13 @@ class VolcanoArkProvider:
|
||||
stream=True,
|
||||
timeout=timeout,
|
||||
) as response:
|
||||
response.raise_for_status()
|
||||
if response.status_code >= 400:
|
||||
# 光有状态码定位不了问题:ARK 的参数校验原因只在响应体里(例如某个
|
||||
# 字段不被该模型支持)。带上正文再抛,否则每次都只能靠猜。
|
||||
detail = (response.text or "").strip()[:600]
|
||||
raise requests.HTTPError(
|
||||
f"{response.status_code} from {endpoint}: {detail}", response=response
|
||||
)
|
||||
# SSE 响应常不带 charset,requests 会按 latin-1 解码 → 中文乱码。强制 UTF-8。
|
||||
response.encoding = "utf-8"
|
||||
for raw in response.iter_lines(decode_unicode=True):
|
||||
@@ -123,6 +129,11 @@ class VolcanoArkProvider:
|
||||
choices = chunk.get("choices") or []
|
||||
if not choices:
|
||||
continue
|
||||
# finish_reason=length 说明输出撞到 max_tokens 被截断。调用方必须能区分
|
||||
# 「模型没写」和「写了但被砍掉」,否则坏 JSON 只会被当成模型偷懒反复重试。
|
||||
finish = choices[0].get("finish_reason")
|
||||
if finish:
|
||||
yield {"type": "finish", "reason": str(finish)}
|
||||
delta = choices[0].get("delta") or {}
|
||||
# 推理模型(豆包 seed-pro / 部分中转 o系/gemini)思考阶段只发 reasoning_content,
|
||||
# 不发 content。必须单独转发,否则整个思考期(可达几十秒~分钟)前端零输出 = 假死。
|
||||
@@ -226,7 +237,14 @@ class VolcanoArkProvider:
|
||||
"generate_audio": generate_audio,
|
||||
}
|
||||
if seed is not None and seed != -1:
|
||||
body["seed"] = seed
|
||||
# Seedance 2.5 r2v 硬上限是 int32 正数(2147483647)。文档写 2^32-1,但 r2v 实际更严;
|
||||
# 超一点整条 InvalidParameter。这里兜住所有上游构造路径。
|
||||
try:
|
||||
clamped = int(seed) & 2147483647
|
||||
except (TypeError, ValueError):
|
||||
clamped = None
|
||||
if clamped:
|
||||
body["seed"] = clamped
|
||||
if search_mode == "smart":
|
||||
body["tools"] = [{"type": "web_search"}]
|
||||
response = requests.post(
|
||||
|
||||
@@ -10,6 +10,65 @@ from .models import (
|
||||
)
|
||||
|
||||
|
||||
# 仅把可解释的阶段名返回给创作页。模型原始 reasoning 可能冗长、跑题或包含内部工作草稿,
|
||||
# 不能直接作为用户可见内容;前端据此展示「正在做什么」,避免异步规划看起来像卡住。
|
||||
_PUBLIC_AGENT_PROGRESS = {
|
||||
"starting": {
|
||||
"label": "正在读取创作要求",
|
||||
"detail": "已收到素材、角色与时长设定",
|
||||
},
|
||||
"reasoning": {
|
||||
"label": "正在分析素材与创作方向",
|
||||
"detail": "正在结合商品、角色和目标时长梳理方案",
|
||||
},
|
||||
"search_library": {
|
||||
"label": "正在查找可用素材",
|
||||
"detail": "正在核对当前创作需要的素材信息",
|
||||
},
|
||||
"write_strategy": {
|
||||
"label": "正在确定创作策略",
|
||||
"detail": "正在整理受众、核心卖点和整体表达方向",
|
||||
},
|
||||
"write_plan": {
|
||||
"label": "正在编排视频方案",
|
||||
"detail": "正在安排分段节奏,并锁定角色和商品的一致性",
|
||||
},
|
||||
"write_prompt": {
|
||||
"label": "正在整理出片指令",
|
||||
"detail": "正在把方案转成可直接生成的视频脚本",
|
||||
},
|
||||
"generate_image": {
|
||||
"label": "正在准备生成画面",
|
||||
"detail": "正在整理画面参考与生成条件",
|
||||
},
|
||||
"ask_user": {
|
||||
"label": "正在核对关键设定",
|
||||
"detail": "正在确认会影响成片效果的信息",
|
||||
},
|
||||
"responding": {
|
||||
"label": "正在整理回复内容",
|
||||
"detail": "马上把下一步呈现给你",
|
||||
},
|
||||
}
|
||||
|
||||
_PUBLIC_REASONING_DETAILS = {
|
||||
"brief": "正在理解这次创作的重点和限制",
|
||||
"product": "正在提炼商品卖点与可呈现的真实证据",
|
||||
"cast": "正在安排角色出镜方式与人物关系",
|
||||
"structure": "正在检查时长、分段节奏和前后衔接",
|
||||
"consistency": "正在锁定人物与商品在各段的一致性",
|
||||
"shots": "正在细化镜头、动作和出片表达",
|
||||
}
|
||||
|
||||
|
||||
def _public_agent_progress(phase: str, detail_key: str = "") -> dict[str, str]:
|
||||
"""把持久化的阶段键转换成用户可见的受控摘要。"""
|
||||
progress = dict(_PUBLIC_AGENT_PROGRESS.get(phase, _PUBLIC_AGENT_PROGRESS["starting"]))
|
||||
if phase == "reasoning":
|
||||
progress["detail"] = _PUBLIC_REASONING_DETAILS.get(detail_key, progress["detail"])
|
||||
return progress
|
||||
|
||||
|
||||
class ModelProviderSerializer(serializers.ModelSerializer):
|
||||
class Meta:
|
||||
model = ModelProvider
|
||||
@@ -114,17 +173,19 @@ class CreationConversationSerializer(serializers.ModelSerializer):
|
||||
|
||||
message_count = serializers.SerializerMethodField()
|
||||
cover_url = serializers.SerializerMethodField()
|
||||
agent_progress = serializers.SerializerMethodField()
|
||||
|
||||
class Meta:
|
||||
model = CreationConversation
|
||||
fields = [
|
||||
"id", "title", "mode", "preset", "params", "status",
|
||||
"agent_status", "agent_started_at",
|
||||
"agent_progress",
|
||||
"message_count", "cover_url",
|
||||
"last_active_at", "created_at", "updated_at",
|
||||
]
|
||||
read_only_fields = [
|
||||
"id", "status", "agent_status", "agent_started_at",
|
||||
"id", "status", "agent_status", "agent_started_at", "agent_progress",
|
||||
"message_count", "cover_url",
|
||||
"last_active_at", "created_at", "updated_at",
|
||||
]
|
||||
@@ -149,6 +210,32 @@ class CreationConversationSerializer(serializers.ModelSerializer):
|
||||
first = assets[0] or {}
|
||||
return first.get("cover") or first.get("url") or ""
|
||||
|
||||
def get_agent_progress(self, obj):
|
||||
"""规划期间只返回受控阶段文案,绝不把模型原始 thinking 暴露给页面。"""
|
||||
if obj.agent_status != CreationConversation.AgentStatus.PLANNING:
|
||||
return None
|
||||
memory = obj.memory if isinstance(obj.memory, dict) else {}
|
||||
phase = str(memory.get("agent_progress_phase") or "")
|
||||
if not phase:
|
||||
return None
|
||||
detail_key = str(memory.get("agent_progress_detail_key") or "")
|
||||
progress = _public_agent_progress(phase, detail_key)
|
||||
history: list[dict[str, str]] = []
|
||||
for item in memory.get("agent_progress_history") or []:
|
||||
if not isinstance(item, dict):
|
||||
continue
|
||||
history.append(
|
||||
_public_agent_progress(
|
||||
str(item.get("phase") or ""),
|
||||
str(item.get("detail_key") or ""),
|
||||
)
|
||||
)
|
||||
# 兼容服务端升级前已开始的轮次:至少显示当前一条进度。
|
||||
if not history or history[-1] != progress:
|
||||
history.append(progress)
|
||||
progress["history"] = history[-12:]
|
||||
return progress
|
||||
|
||||
def update(self, instance, validated_data):
|
||||
# mode 定死:允许传但忽略,避免前端误改后顶栏参数与已生成内容对不上
|
||||
validated_data.pop("mode", None)
|
||||
|
||||
@@ -3389,6 +3389,10 @@ def poll_video_segment(*, video_segment: VideoSegment, user) -> VideoSegmentVers
|
||||
# 按火山真实 usage.total_tokens 结算(true-up,与自由创作同口径):
|
||||
# 多退(charge 差额自动 RELEASE)/超预留 clamp(ledger 禁超扣,差额平台承担并告警)。
|
||||
# usage 缺失(异常响应)回落预估价,不阻断出片。
|
||||
# 提交那头(submit_video_segment)是函数内 import;这里漏了同一个名字,
|
||||
# 视频段一出片走到结算就 NameError 炸掉,钱没结、版本没建、段永远卡「生成中」。
|
||||
from apps.billing.pricing import settle_video_from_payload
|
||||
|
||||
reservation = locked_task.credit_reservation
|
||||
payload = dict(locked_task.request_payload or {})
|
||||
try:
|
||||
|
||||
@@ -6,6 +6,8 @@
|
||||
import json
|
||||
from unittest.mock import patch
|
||||
|
||||
import requests
|
||||
|
||||
from django.test import SimpleTestCase, TestCase, override_settings
|
||||
from rest_framework.test import APIClient
|
||||
|
||||
@@ -24,14 +26,18 @@ from .creation_agent import (
|
||||
_coerce_fields,
|
||||
_image_count,
|
||||
_merge_tool_call_deltas,
|
||||
apply_explicit_video_duration_from_text,
|
||||
apply_pain_point_direction,
|
||||
apply_restart_intent,
|
||||
apply_person_identity_guard,
|
||||
apply_product_reference_guard,
|
||||
active_plot_twist_story_depth,
|
||||
plot_twist_selected_direction,
|
||||
append_multi_character_relation_gate,
|
||||
append_step_confirm,
|
||||
build_system_prompt,
|
||||
build_messages,
|
||||
clear_public_agent_progress,
|
||||
default_reply_hint,
|
||||
default_reply_options,
|
||||
ensure_turn_guides,
|
||||
@@ -39,12 +45,19 @@ from .creation_agent import (
|
||||
get_video_gate_stage,
|
||||
has_creative_intent,
|
||||
infer_script_duration,
|
||||
MAX_TRUNCATION_RETRIES,
|
||||
TRUNCATION_GIVE_UP_NOTICE,
|
||||
creation_agent_max_output_tokens,
|
||||
creation_model_extra_body,
|
||||
creation_agent_timeout_notice,
|
||||
long_video_script_covers_requested_duration,
|
||||
plan_video_segments,
|
||||
is_continue_intent,
|
||||
is_greeting,
|
||||
is_pure_chitchat,
|
||||
is_restart_intent,
|
||||
session_has_creative_context,
|
||||
set_public_agent_progress,
|
||||
stream_creation_agent,
|
||||
set_plot_twist_story_depth,
|
||||
strip_numeric_reply_instruction,
|
||||
@@ -53,12 +66,14 @@ from .creation_agent import (
|
||||
submit_generated_person_reference,
|
||||
text_already_guides,
|
||||
tool_schemas,
|
||||
_video_submit_params,
|
||||
video_duration,
|
||||
video_model_name,
|
||||
wanted_asset_pick,
|
||||
wanted_param_keys,
|
||||
)
|
||||
from .models import AITask, CreationConversation, CreationMessage, ModelConfig, ModelProvider
|
||||
from .serializers import CreationConversationDetailSerializer
|
||||
|
||||
|
||||
def _text_chunks(text):
|
||||
@@ -78,6 +93,15 @@ def _tool_chunks(name, arguments, *, said=""):
|
||||
yield {"type": "done"}
|
||||
|
||||
|
||||
def _truncated_tool_chunks(name, arguments):
|
||||
"""模拟输出撞到 max_tokens:arguments 断在 JSON 中间,并带 finish_reason=length。"""
|
||||
yield {"type": "tool_call", "tool_calls": [{"index": 0, "function": {"name": name, "arguments": ""}}]}
|
||||
blob = json.dumps(arguments, ensure_ascii=False)
|
||||
yield {"type": "tool_call", "tool_calls": [{"index": 0, "function": {"arguments": blob[: len(blob) // 2]}}]}
|
||||
yield {"type": "finish", "reason": "length"}
|
||||
yield {"type": "done"}
|
||||
|
||||
|
||||
class FakeProvider:
|
||||
"""按脚本逐轮回放。每调用一次 chat_completion_stream 消费一个剧本。"""
|
||||
|
||||
@@ -124,6 +148,33 @@ class CreationAgentBaseTests(TestCase):
|
||||
))
|
||||
return events, fake
|
||||
|
||||
def _pin_person(self, name="出镜达人"):
|
||||
"""把人物这一步真正走完 —— 钉住一个角色素材。
|
||||
|
||||
memory["person_source_ready"] 已经不算完成了(person_identity_ready 的「防跳过」):
|
||||
只有钉住人、正在生成、或明确只出手才算。想跳过人物闸门的用例必须钉真人物。
|
||||
"""
|
||||
asset = Asset.objects.create(
|
||||
team=self.team, created_by=self.user, name=name,
|
||||
asset_type=Asset.Type.IMAGE, category=Asset.Category.PERSON,
|
||||
)
|
||||
refs = [ref for ref in (self.conversation.pinned_refs or []) if isinstance(ref, dict)]
|
||||
refs.append({"type": "character", "id": str(asset.id), "name": name})
|
||||
self.conversation.pinned_refs = refs
|
||||
self.conversation.save(update_fields=["pinned_refs", "updated_at"])
|
||||
return asset
|
||||
|
||||
def _pin_product(self, title="舒缓面霜"):
|
||||
"""钉住商品 —— 写方案前会先问「推哪款商品」,想直接拿方案的用例必须先把商品定下来。"""
|
||||
from apps.products.models import Product
|
||||
|
||||
product = Product.objects.create(team=self.team, created_by=self.user, title=title)
|
||||
refs = [ref for ref in (self.conversation.pinned_refs or []) if isinstance(ref, dict)]
|
||||
refs.append({"type": "product", "id": str(product.id), "name": title})
|
||||
self.conversation.pinned_refs = refs
|
||||
self.conversation.save(update_fields=["pinned_refs", "updated_at"])
|
||||
return product
|
||||
|
||||
|
||||
class ToolCallAssemblyTests(TestCase):
|
||||
def test_streamed_arguments_are_concatenated_not_overwritten(self):
|
||||
@@ -173,13 +224,21 @@ class FieldCoercionTests(TestCase):
|
||||
raw.append({"key": "bad", "label": "x", "type": "dropdown"})
|
||||
self.assertEqual(len(_coerce_fields(raw)), 1) # 聊天式追问一次只问 1 项
|
||||
|
||||
def test_product_text_or_single_is_coerced_to_asset_card(self):
|
||||
def test_product_text_is_coerced_to_asset_card(self):
|
||||
fields = _coerce_fields([
|
||||
{"key": "product", "label": "选择商品", "type": "text"},
|
||||
])
|
||||
self.assertEqual(fields[0]["type"], "asset")
|
||||
self.assertEqual(fields[0]["asset_types"], ["product"])
|
||||
|
||||
def test_choice_with_real_options_is_not_degraded_to_asset_card(self):
|
||||
"""模型真给了具体选项的选择题保留 options —— 否则痛点方向三选一会被换成素材卡。"""
|
||||
fields = _coerce_fields([
|
||||
{"key": "product", "label": "选择商品", "type": "single",
|
||||
"options": [{"value": "a", "label": "A"}]},
|
||||
])
|
||||
self.assertEqual(fields[0]["type"], "asset")
|
||||
self.assertEqual(fields[0]["asset_types"], ["product"])
|
||||
self.assertEqual(fields[0]["type"], "single")
|
||||
self.assertEqual(fields[0]["options"], [{"value": "a", "label": "A"}])
|
||||
|
||||
def test_duration_single_stays_single(self):
|
||||
fields = _coerce_fields([
|
||||
@@ -285,7 +344,12 @@ class AskUserTests(CreationAgentBaseTests):
|
||||
self.assertEqual(payload["fields"][0]["key"], "_asset_gate")
|
||||
self.assertEqual(
|
||||
payload["fields"][0]["options"],
|
||||
[{"value": "send", "label": "发商品列表"}, {"value": "auto", "label": "你来推荐"}],
|
||||
[
|
||||
# 商品闸门多给一条「上传商品图」——手上只有实物图的商家不用先去建商品。
|
||||
{"value": "upload", "label": "上传商品图"},
|
||||
{"value": "send", "label": "发商品列表"},
|
||||
{"value": "auto", "label": "你来推荐"},
|
||||
],
|
||||
)
|
||||
self.assertEqual(payload["pending_fields"][0]["type"], "asset")
|
||||
self.assertEqual(payload["pending_fields"][0]["asset_types"], ["product"])
|
||||
@@ -384,7 +448,9 @@ class GenerateImageTests(CreationAgentBaseTests):
|
||||
)
|
||||
card = self._confirm_from_events(events)
|
||||
self.assertEqual(card.payload["kind"], "image")
|
||||
self.assertEqual(card.payload["prompt"], "干净棚拍,柔光,居中构图")
|
||||
# 带了商品参考图就会追加外观锁定约束,避免出图换色换款。
|
||||
self.assertTrue(card.payload["prompt"].startswith("干净棚拍,柔光,居中构图"))
|
||||
self.assertIn("【商品外观锁定·最高优先级】", card.payload["prompt"])
|
||||
self.assertFalse(any(e.get("type") == "task" for e in events))
|
||||
|
||||
with patch("apps.ai.services.enqueue_standalone_images") as enqueue:
|
||||
@@ -1370,14 +1436,51 @@ class VideoPlanAndConfirmTests(CreationAgentBaseTests):
|
||||
return args
|
||||
|
||||
def _pin_product(self, title="测试商品"):
|
||||
product = Product.objects.create(team=self.team, created_by=self.user, title=title)
|
||||
self.conversation.pinned_refs = [{
|
||||
"type": "product",
|
||||
"id": str(product.id),
|
||||
"name": product.title,
|
||||
}]
|
||||
self.conversation.save(update_fields=["pinned_refs", "updated_at"])
|
||||
return product
|
||||
# 追加而不是整表替换:已经钉过人物的用例不能因为钉商品把人物挤掉,
|
||||
# 那样人物来源闸门会重新拦住整轮,方案卡一张都出不来。
|
||||
return super()._pin_product(title)
|
||||
|
||||
def test_natural_brief_updates_explicit_120_second_duration(self):
|
||||
changed = apply_explicit_video_duration_from_text(
|
||||
self.conversation,
|
||||
"用这三个角色和这个商品,做一条120秒的真实自然带货视频。",
|
||||
)
|
||||
|
||||
self.assertTrue(changed)
|
||||
self.conversation.refresh_from_db()
|
||||
self.assertEqual(self.conversation.params["duration"], "120 秒")
|
||||
|
||||
def test_planning_exposes_safe_progress_not_raw_reasoning(self):
|
||||
self.conversation.agent_status = CreationConversation.AgentStatus.PLANNING
|
||||
self.conversation.save(update_fields=["agent_status", "updated_at"])
|
||||
set_public_agent_progress(self.conversation, "starting")
|
||||
set_public_agent_progress(self.conversation, "reasoning", detail_key="structure")
|
||||
set_public_agent_progress(self.conversation, "write_plan")
|
||||
|
||||
data = CreationConversationDetailSerializer(self.conversation).data
|
||||
|
||||
self.assertEqual(data["agent_progress"]["label"], "正在编排视频方案")
|
||||
self.assertEqual(data["agent_progress"]["detail"], "正在安排分段节奏,并锁定角色和商品的一致性")
|
||||
self.assertEqual(
|
||||
data["agent_progress"]["history"][1]["detail"],
|
||||
"正在检查时长、分段节奏和前后衔接",
|
||||
)
|
||||
self.assertNotIn("reasoning", data["agent_progress"])
|
||||
clear_public_agent_progress(self.conversation)
|
||||
self.assertIsNone(CreationConversationDetailSerializer(self.conversation).data["agent_progress"])
|
||||
self.conversation.agent_status = CreationConversation.AgentStatus.IDLE
|
||||
self.conversation.save(update_fields=["agent_status", "updated_at"])
|
||||
self.assertIsNone(CreationConversationDetailSerializer(self.conversation).data["agent_progress"])
|
||||
|
||||
def test_shot_timeline_does_not_override_session_duration(self):
|
||||
changed = apply_explicit_video_duration_from_text(
|
||||
self.conversation,
|
||||
"0-30秒先让角色出场,30-60秒展示商品。",
|
||||
)
|
||||
|
||||
self.assertFalse(changed)
|
||||
self.conversation.refresh_from_db()
|
||||
self.assertEqual(self.conversation.params["duration"], "15 秒")
|
||||
|
||||
def test_plot_twist_preset_requires_story_depth_before_creative_output(self):
|
||||
self.conversation.preset = "剧情反转带货"
|
||||
@@ -1402,7 +1505,7 @@ class VideoPlanAndConfirmTests(CreationAgentBaseTests):
|
||||
self.assertEqual(question["payload"].get("interaction"), "plot_twist_story_depth")
|
||||
self.assertEqual(
|
||||
[option["value"] for option in question["payload"]["fields"][0]["options"]],
|
||||
["15s", "30s", "60s", "smart"],
|
||||
["15s", "30s", "60s", "180s", "smart"],
|
||||
)
|
||||
self.conversation.refresh_from_db()
|
||||
self.assertEqual(self.conversation.agent_status, "awaiting_user")
|
||||
@@ -1439,8 +1542,8 @@ class VideoPlanAndConfirmTests(CreationAgentBaseTests):
|
||||
self._pin_product("剧情测试商品")
|
||||
self.conversation.preset = "剧情反转带货"
|
||||
self.conversation.params = {**self.conversation.params, "duration": "30 秒"}
|
||||
self.conversation.memory = {"person_source_ready": True}
|
||||
self.conversation.save(update_fields=["preset", "params", "memory", "updated_at"])
|
||||
self.conversation.save(update_fields=["preset", "params", "updated_at"])
|
||||
self._pin_person()
|
||||
# 即使模型只输出一句空话,平台也必须补上三个可点击方向,而不是让用户继续追问。
|
||||
fake = FakeProvider([_text_chunks("我准备了三个剧情反转方向。")])
|
||||
|
||||
@@ -1464,14 +1567,16 @@ class VideoPlanAndConfirmTests(CreationAgentBaseTests):
|
||||
|
||||
def test_plot_twist_direction_is_locked_into_plan_prompt(self):
|
||||
"""已选剧情方向必须写进方案 video_prompt,不能漂成另一套默认带货故事。"""
|
||||
from apps.ai.creation_agent import get_pending_video_prompt
|
||||
from apps.ai.creation_presets import apply_plot_twist_direction_contract
|
||||
from apps.ai.views import _plot_twist_direction_continuation
|
||||
|
||||
self.conversation.mode = CreationConversation.Mode.VIDEO
|
||||
self.conversation.preset = "剧情反转带货"
|
||||
self._pin_person()
|
||||
self._pin_product("自动喂食器")
|
||||
self.conversation.memory = {
|
||||
"plot_twist_story_depth": "30s",
|
||||
"person_source_ready": True,
|
||||
"selling_point_ready": True,
|
||||
"selling_point_mode": "manual",
|
||||
"selling_point": "定时定量",
|
||||
@@ -1518,11 +1623,13 @@ class VideoPlanAndConfirmTests(CreationAgentBaseTests):
|
||||
text="按选定方向写方案",
|
||||
model_config=self.model,
|
||||
))
|
||||
plan = next(
|
||||
event["message"] for event in events
|
||||
if event.get("type") == "message" and event["message"]["kind"] == "plan"
|
||||
)
|
||||
prompt = str((plan.get("payload") or {}).get("video_prompt") or "")
|
||||
self.assertTrue(any(
|
||||
event.get("type") == "message" and event["message"]["kind"] == "plan"
|
||||
for event in events
|
||||
))
|
||||
# 出片指令不再挂在方案卡上,改为存成待确认的 pending prompt。
|
||||
self.conversation.refresh_from_db()
|
||||
prompt = get_pending_video_prompt(self.conversation)
|
||||
self.assertIn("【剧情反转方向·强制执行·最高优先级】", prompt)
|
||||
self.assertIn("邻居误以为", prompt)
|
||||
self.assertIn("误会当场解除", prompt)
|
||||
@@ -1533,9 +1640,9 @@ class VideoPlanAndConfirmTests(CreationAgentBaseTests):
|
||||
product = Product.objects.create(team=self.team, created_by=self.user, title="舒缓面霜")
|
||||
self.conversation.mode = CreationConversation.Mode.VIDEO
|
||||
self.conversation.preset = "痛点解决演示"
|
||||
self.conversation.memory = {"person_source_ready": True}
|
||||
self.conversation.pinned_refs = [{"type": "product", "id": str(product.id), "name": product.title}]
|
||||
self.conversation.save(update_fields=["mode", "preset", "memory", "pinned_refs", "updated_at"])
|
||||
self.conversation.save(update_fields=["mode", "preset", "pinned_refs", "updated_at"])
|
||||
self._pin_person()
|
||||
fake = FakeProvider([_text_chunks(
|
||||
"可以从下面三个方向选一个:\n"
|
||||
"- 换季干燥紧绷,正常涂抹后保持舒适不拔干\n"
|
||||
@@ -1586,10 +1693,10 @@ class VideoPlanAndConfirmTests(CreationAgentBaseTests):
|
||||
self.assertNotIn("generate_image", names)
|
||||
|
||||
def test_strategy_card_stops_with_step_confirm(self):
|
||||
self._pin_person()
|
||||
self.conversation.memory = {
|
||||
"selling_point_ready": True,
|
||||
"selling_point_mode": "auto",
|
||||
"person_source_ready": True,
|
||||
}
|
||||
self.conversation.save(update_fields=["memory", "updated_at"])
|
||||
fake = FakeProvider([
|
||||
@@ -1615,6 +1722,7 @@ class VideoPlanAndConfirmTests(CreationAgentBaseTests):
|
||||
|
||||
def test_plain_strategy_prose_is_recovered_as_strategy_card(self):
|
||||
self._pin_product("小熊婴儿湿巾")
|
||||
self._pin_person()
|
||||
self.conversation.preset = "痛点解决演示"
|
||||
self.conversation.memory = {
|
||||
"selling_point_ready": True,
|
||||
@@ -1622,7 +1730,6 @@ class VideoPlanAndConfirmTests(CreationAgentBaseTests):
|
||||
"selling_point": "温和不刺激",
|
||||
"pain_point_direction_ready": True,
|
||||
"pain_point_direction": "温和不刺激",
|
||||
"person_source_ready": True,
|
||||
}
|
||||
self.conversation.save(update_fields=["preset", "memory", "updated_at"])
|
||||
fake = FakeProvider([_text_chunks(
|
||||
@@ -1667,8 +1774,7 @@ class VideoPlanAndConfirmTests(CreationAgentBaseTests):
|
||||
self.assertEqual(confirm["payload"].get("step"), "strategy")
|
||||
|
||||
def test_selling_point_is_confirmed_before_strategy(self):
|
||||
self.conversation.memory = {"person_source_ready": True}
|
||||
self.conversation.save(update_fields=["memory", "updated_at"])
|
||||
self._pin_person()
|
||||
fake = FakeProvider([
|
||||
_tool_chunks("write_strategy", {"target": "油皮通勤人群", "trust": "真实使用反馈",
|
||||
"belief": "值得一试", "direction": "达人 UGC 口播"}),
|
||||
@@ -1710,6 +1816,214 @@ class VideoPlanAndConfirmTests(CreationAgentBaseTests):
|
||||
# 方案闸门停下,不能同轮出 Prompt/积分卡
|
||||
self.assertEqual(len(fake.calls), 1)
|
||||
|
||||
def test_120_second_plan_accepts_two_second_timeline_slack(self):
|
||||
self.conversation.params = {**self.conversation.params, "duration": "120 秒"}
|
||||
self.conversation.memory = {"stage": "strategy", "strategy_confirmed": True, "person_source_ready": True}
|
||||
self.conversation.save(update_fields=["params", "memory", "updated_at"])
|
||||
fake = FakeProvider([
|
||||
_tool_chunks("write_plan", self._plan_args(
|
||||
timeline=[
|
||||
{"start": 0, "end": 60, "stage": "第一章"},
|
||||
{"start": 60, "end": 118, "stage": "第二章"},
|
||||
],
|
||||
video_prompt="总时长:120秒。第一章 0-60 秒建立信任,第二章 60-118 秒完成使用证据和收束。",
|
||||
)),
|
||||
_text_chunks("不该继续"),
|
||||
])
|
||||
with patch("apps.ai.creation_agent.build_provider", return_value=fake):
|
||||
events = _events(stream_creation_agent(
|
||||
conversation=self.conversation,
|
||||
user=self.user,
|
||||
text="按这个继续",
|
||||
model_config=self.model,
|
||||
))
|
||||
kinds = [event["message"]["kind"] for event in events if event.get("type") == "message"]
|
||||
self.assertIn("plan", kinds)
|
||||
self.assertEqual(len(fake.calls), 1)
|
||||
|
||||
def test_120_second_plan_tool_asks_for_compact_chapter_draft(self):
|
||||
self.conversation.params = {**self.conversation.params, "duration": "120 秒"}
|
||||
self.conversation.save(update_fields=["params", "updated_at"])
|
||||
context = AgentContext(conversation=self.conversation, user=self.user, model_config=self.model)
|
||||
tools = {item["function"]["name"]: item["function"] for item in tool_schemas(context, allow_plan=True)}
|
||||
self.assertIn("800–1400", tools["write_plan"]["description"])
|
||||
self.assertIn("不要套 15 秒短片的制作级长文模板", tools["write_plan"]["description"])
|
||||
self.assertIn("只补一段简短全片规则", tools["write_prompt"]["description"])
|
||||
system = build_system_prompt(context)
|
||||
self.assertIn("紧凑章节稿", system)
|
||||
self.assertNotIn("1800–3200", system)
|
||||
|
||||
def test_120_second_plan_rejects_incomplete_timeline(self):
|
||||
self.conversation.params = {**self.conversation.params, "duration": "120 秒"}
|
||||
self.conversation.memory = {"stage": "strategy", "strategy_confirmed": True, "person_source_ready": True}
|
||||
self.conversation.save(update_fields=["params", "memory", "updated_at"])
|
||||
fake = FakeProvider([
|
||||
_tool_chunks("write_plan", self._plan_args(
|
||||
timeline=[{"start": 0, "end": 90, "stage": "第一章"}],
|
||||
video_prompt="只写到 90 秒的半成品脚本",
|
||||
)),
|
||||
_text_chunks("我补一版完整时间轴"),
|
||||
])
|
||||
with patch("apps.ai.creation_agent.build_provider", return_value=fake):
|
||||
events = _events(stream_creation_agent(
|
||||
conversation=self.conversation,
|
||||
user=self.user,
|
||||
text="按这个继续",
|
||||
model_config=self.model,
|
||||
))
|
||||
kinds = [event["message"]["kind"] for event in events if event.get("type") == "message"]
|
||||
self.assertNotIn("plan", kinds)
|
||||
|
||||
def test_model_call_lifts_max_output_tokens(self):
|
||||
"""豆包默认只给 4k 输出,长方案会被截断成坏 JSON。请求必须显式抬高上限。"""
|
||||
self.conversation.memory = {"stage": "strategy", "strategy_confirmed": True, "person_source_ready": True}
|
||||
self.conversation.save(update_fields=["memory", "updated_at"])
|
||||
fake = FakeProvider([_tool_chunks("write_plan", self._plan_args())])
|
||||
with patch("apps.ai.creation_agent.build_provider", return_value=fake):
|
||||
_events(stream_creation_agent(
|
||||
conversation=self.conversation,
|
||||
user=self.user,
|
||||
text="按这个继续",
|
||||
model_config=self.model,
|
||||
))
|
||||
extra_body = fake.calls[0]["extra_body"]
|
||||
self.assertEqual(extra_body["max_tokens"], creation_agent_max_output_tokens())
|
||||
self.assertGreaterEqual(extra_body["max_tokens"], 16000)
|
||||
self.assertTrue(extra_body["tools"])
|
||||
|
||||
def _volcano_text_model(self):
|
||||
volcano, _ = ModelProvider.objects.get_or_create(
|
||||
name="volcengine", defaults={"display_name": "火山方舟", "base_url": "https://ark"}
|
||||
)
|
||||
model, _ = ModelConfig.objects.get_or_create(
|
||||
provider=volcano,
|
||||
name="doubao-seed-2-1-pro-260628",
|
||||
capability=ModelConfig.Capability.TEXT,
|
||||
defaults={"display_name": "Seed 2.1 Pro", "endpoint": "chat/completions"},
|
||||
)
|
||||
return model
|
||||
|
||||
def test_thinking_is_disabled_by_default_and_never_hits_gateways(self):
|
||||
"""深度思考实测要多花 100 秒才开始吐正文,编排 agent 默认关掉;
|
||||
thinking 是火山私有字段,中转站一律不下发。"""
|
||||
official = self._volcano_text_model()
|
||||
self.assertEqual(creation_model_extra_body(official, [])["thinking"], {"type": "disabled"})
|
||||
self.assertNotIn("thinking", creation_model_extra_body(self.model, []))
|
||||
with override_settings(CREATION_AGENT_THINKING_MODE=""):
|
||||
self.assertNotIn("thinking", creation_model_extra_body(official, []))
|
||||
|
||||
def test_model_rejecting_thinking_param_retries_without_it(self):
|
||||
"""模型不认这个档位时自动去掉重试,不能整条会话吐「生成过程出错了」。"""
|
||||
self.conversation.memory = {"stage": "strategy", "strategy_confirmed": True, "person_source_ready": True}
|
||||
self.conversation.save(update_fields=["memory", "updated_at"])
|
||||
|
||||
class PickyProvider(FakeProvider):
|
||||
def chat_completion_stream(self, **kwargs):
|
||||
if "thinking" in (kwargs.get("extra_body") or {}):
|
||||
raise requests.HTTPError(
|
||||
"400 from chat/completions: "
|
||||
'{"error":{"code":"InvalidParameter",'
|
||||
'"message":"Unsupported thinking type for the current model: auto"}}'
|
||||
)
|
||||
return super().chat_completion_stream(**kwargs)
|
||||
|
||||
fake = PickyProvider([_tool_chunks("write_plan", self._plan_args())])
|
||||
with override_settings(CREATION_AGENT_THINKING_MODE="auto"), \
|
||||
patch("apps.ai.creation_agent.build_provider", return_value=fake):
|
||||
events = _events(stream_creation_agent(
|
||||
conversation=self.conversation,
|
||||
user=self.user,
|
||||
text="按这个继续",
|
||||
model_config=self._volcano_text_model(),
|
||||
))
|
||||
|
||||
kinds = [event["message"]["kind"] for event in events if event.get("type") == "message"]
|
||||
self.assertIn("plan", kinds)
|
||||
self.assertNotIn("error", kinds)
|
||||
self.assertNotIn("thinking", fake.calls[-1]["extra_body"])
|
||||
|
||||
def test_truncated_plan_arguments_ask_for_a_compact_rewrite(self):
|
||||
"""被 max_tokens 砍断的 tool 参数不是「模型没写」。要指名截断并让它写紧凑版,而不是原样重试到超时。"""
|
||||
self.conversation.params = {**self.conversation.params, "duration": "120 秒"}
|
||||
self.conversation.memory = {"stage": "strategy", "strategy_confirmed": True, "person_source_ready": True}
|
||||
self.conversation.save(update_fields=["params", "memory", "updated_at"])
|
||||
fake = FakeProvider([
|
||||
_truncated_tool_chunks("write_plan", self._plan_args(
|
||||
timeline=[{"start": 0, "end": 60, "stage": "第一章"}, {"start": 60, "end": 120, "stage": "第二章"}],
|
||||
video_prompt="写到一半就被砍掉的残稿" * 200,
|
||||
)),
|
||||
_tool_chunks("write_plan", self._plan_args(
|
||||
timeline=[{"start": 0, "end": 60, "stage": "第一章"}, {"start": 60, "end": 120, "stage": "第二章"}],
|
||||
video_prompt="第一章 0-60 秒 关键镜头…第二章 60-120 秒 关键镜头…",
|
||||
)),
|
||||
])
|
||||
with patch("apps.ai.creation_agent.build_provider", return_value=fake):
|
||||
events = _events(stream_creation_agent(
|
||||
conversation=self.conversation,
|
||||
user=self.user,
|
||||
text="按这个继续",
|
||||
model_config=self.model,
|
||||
))
|
||||
|
||||
# FakeProvider 记录的是同一个 messages 列表引用,只能看最终态
|
||||
history = json.dumps(fake.calls[-1]["messages"], ensure_ascii=False)
|
||||
self.assertIn("截断", history)
|
||||
self.assertIn("紧凑", history)
|
||||
# 坏 JSON 不进历史,否则模型会照着残稿续写
|
||||
self.assertNotIn("写到一半就被砍掉的残稿", history)
|
||||
self.assertEqual(len(fake.calls), 2)
|
||||
kinds = [event["message"]["kind"] for event in events if event.get("type") == "message"]
|
||||
self.assertIn("plan", kinds)
|
||||
self.assertNotIn("error", kinds)
|
||||
|
||||
def test_repeated_truncation_stops_with_a_plain_notice(self):
|
||||
self.conversation.memory = {"stage": "strategy", "strategy_confirmed": True, "person_source_ready": True}
|
||||
self.conversation.save(update_fields=["memory", "updated_at"])
|
||||
fake = FakeProvider([
|
||||
_truncated_tool_chunks("write_plan", self._plan_args(video_prompt="太长了…" * 300))
|
||||
for _ in range(4)
|
||||
])
|
||||
with patch("apps.ai.creation_agent.build_provider", return_value=fake):
|
||||
events = _events(stream_creation_agent(
|
||||
conversation=self.conversation,
|
||||
user=self.user,
|
||||
text="按这个继续",
|
||||
model_config=self.model,
|
||||
))
|
||||
errors = [
|
||||
event["message"]["text"]
|
||||
for event in events
|
||||
if event.get("type") == "message" and event["message"].get("kind") == "error"
|
||||
]
|
||||
self.assertEqual(errors, [TRUNCATION_GIVE_UP_NOTICE])
|
||||
# 只重写 MAX_TRUNCATION_RETRIES 次就收手,不再闷着重试到整轮超时
|
||||
self.assertEqual(len(fake.calls), MAX_TRUNCATION_RETRIES + 1)
|
||||
|
||||
def test_long_video_timeout_keeps_confirm_and_explains_wait(self):
|
||||
self.conversation.params = {**self.conversation.params, "duration": "120 秒"}
|
||||
self.conversation.save(update_fields=["params", "updated_at"])
|
||||
append_step_confirm(self.conversation, "strategy")
|
||||
|
||||
class TimeoutProvider:
|
||||
def chat_completion_stream(self, **kwargs):
|
||||
raise TimeoutError("creation agent model deadline exceeded")
|
||||
|
||||
with patch("apps.ai.creation_agent.build_provider", return_value=TimeoutProvider()):
|
||||
events = _events(stream_creation_agent(
|
||||
conversation=self.conversation,
|
||||
user=self.user,
|
||||
text="按这个继续",
|
||||
model_config=self.model,
|
||||
))
|
||||
errors = [
|
||||
event["message"]["text"]
|
||||
for event in events
|
||||
if event.get("type") == "message" and event["message"].get("kind") == "error"
|
||||
]
|
||||
self.assertTrue(errors)
|
||||
self.assertIn("还没写完", errors[0])
|
||||
self.assertIn("按这个继续", errors[0])
|
||||
|
||||
def test_click_swap_requires_sequence_before_calling_model(self):
|
||||
self._pin_product("三色通勤包")
|
||||
self.conversation.preset = "点击换款"
|
||||
@@ -1739,7 +2053,8 @@ class VideoPlanAndConfirmTests(CreationAgentBaseTests):
|
||||
self.conversation.memory = {
|
||||
"click_swap_ready": True,
|
||||
"click_swap_sequence": "黑色 → 白色 → 樱花粉",
|
||||
"person_source_ready": True,
|
||||
# 形态是硬闸门:不选「只出手 / 角色换款」,write_plan 会先弹形态卡而不是出方案。
|
||||
"click_swap_mode": "finger",
|
||||
}
|
||||
self.conversation.save(update_fields=["preset", "memory", "updated_at"])
|
||||
fake = FakeProvider([_tool_chunks("write_plan", self._plan_args(
|
||||
@@ -1814,10 +2129,10 @@ class VideoPlanAndConfirmTests(CreationAgentBaseTests):
|
||||
"""模型若同轮连调 write_strategy+write_plan,只落策略闸门。"""
|
||||
from apps.ai.creation_agent import _parse_arguments # noqa: F401
|
||||
|
||||
self._pin_person()
|
||||
self.conversation.memory = {
|
||||
"selling_point_ready": True,
|
||||
"selling_point_mode": "auto",
|
||||
"person_source_ready": True,
|
||||
}
|
||||
self.conversation.save(update_fields=["memory", "updated_at"])
|
||||
|
||||
@@ -1950,7 +2265,8 @@ class VideoPlanAndConfirmTests(CreationAgentBaseTests):
|
||||
self.assertIn("已添加多位角色时保留全部角色", system)
|
||||
self.assertNotIn("固定一位可信的真人", system)
|
||||
|
||||
def test_commerce_preset_requests_product_and_does_not_treat_raw_asset_as_product(self):
|
||||
def test_commerce_preset_asks_brand_and_name_for_an_uploaded_product_image(self):
|
||||
"""上传图现在可以直接当商品图用,但品牌和具体品名必须先问清楚,不许平台自己猜。"""
|
||||
material = Asset.objects.create(
|
||||
team=self.team,
|
||||
created_by=self.user,
|
||||
@@ -1976,8 +2292,9 @@ class VideoPlanAndConfirmTests(CreationAgentBaseTests):
|
||||
if event.get("type") == "message" and event["message"]["kind"] == "elicit"
|
||||
)
|
||||
self.assertEqual(card["payload"].get("interaction"), "chat")
|
||||
self.assertEqual(card["payload"].get("phase"), "gate")
|
||||
self.assertEqual(card["payload"]["pending_fields"][0]["asset_types"], ["product"])
|
||||
self.assertEqual(card["payload"].get("topic"), "product_info")
|
||||
self.assertEqual(card["payload"]["fields"][0]["key"], "product_brand_and_name")
|
||||
self.assertIn("品牌", card["text"])
|
||||
self.assertEqual(fake.calls, [])
|
||||
|
||||
def test_sixty_second_segments_reuse_full_prompt_and_same_person_reference(self):
|
||||
@@ -2051,7 +2368,85 @@ class VideoPlanAndConfirmTests(CreationAgentBaseTests):
|
||||
self.assertEqual(params["prompt"].count("字幕"), 1)
|
||||
self.assertIn("【人物一致性硬约束】", params["prompt"])
|
||||
self.assertIn("【预设执行层·达人口播种草】", params["prompt"])
|
||||
self.assertIn("必须继续使用参考图锁定的原有人物阵容与身份关系", params["prompt"])
|
||||
self.assertIn("所有分段必须复用完全相同的参考图编号、人物阵容", params["prompt"])
|
||||
self.assertLessEqual(params["seed"], 2147483647)
|
||||
|
||||
def test_conversation_id_seed_fits_seedance_int32(self):
|
||||
"""用户这次翻车:会话 id e5e74700… 算出 seed=3857139456,火山 r2v 直接拒。"""
|
||||
import uuid
|
||||
|
||||
from apps.ai.free_video import VIDEO_SEED_MAX, normalize_video_seed
|
||||
|
||||
portrait = Asset.objects.create(
|
||||
team=self.team,
|
||||
created_by=self.user,
|
||||
name="周野",
|
||||
asset_type=Asset.Type.IMAGE,
|
||||
source=Asset.Source.UPLOAD,
|
||||
category=Asset.Category.MODEL_PORTRAIT,
|
||||
)
|
||||
AssetFile.objects.create(
|
||||
asset=portrait, object_key="p.jpg", bucket="test",
|
||||
preview_url="https://cdn.example/p.jpg", is_primary=True,
|
||||
)
|
||||
model = AssetModel.objects.create(
|
||||
team=self.team, created_by=self.user, name="阳光运动 · 周野",
|
||||
portrait_asset=portrait,
|
||||
)
|
||||
self.conversation.pinned_refs = [{"type": "model", "id": str(model.id), "name": model.name}]
|
||||
self.conversation.save(update_fields=["pinned_refs", "updated_at"])
|
||||
CreationConversation.objects.filter(pk=self.conversation.pk).update(
|
||||
id=uuid.UUID("e5e74700-ef19-49fb-bcb2-c156da27c05c")
|
||||
)
|
||||
self.conversation = CreationConversation.objects.get(pk="e5e74700-ef19-49fb-bcb2-c156da27c05c")
|
||||
context = AgentContext(conversation=self.conversation, user=self.user, model_config=self.model)
|
||||
submit, _refs = _video_submit_params(context, "闺蜜日常带货")
|
||||
raw = int("e5e74700", 16)
|
||||
self.assertEqual(raw, 3857139456)
|
||||
self.assertEqual(submit["seed"], normalize_video_seed(raw))
|
||||
self.assertLessEqual(submit["seed"], VIDEO_SEED_MAX)
|
||||
|
||||
@override_settings(FREE_VIDEO_MAX_CONCURRENT=3)
|
||||
def test_180_second_video_starts_in_batches_and_keeps_all_six_segments(self):
|
||||
self.conversation.params = {**self.conversation.params, "duration": "180 秒"}
|
||||
self.conversation.save(update_fields=["params", "updated_at"])
|
||||
card = append_message(
|
||||
self.conversation,
|
||||
role="assistant",
|
||||
kind=CreationMessage.Kind.CONFIRM,
|
||||
payload={"video_prompt": "三章连续叙事,角色用同一件商品解决问题", "submitted": False},
|
||||
)
|
||||
tasks = [
|
||||
AITask.objects.create(
|
||||
team=self.team,
|
||||
created_by=self.user,
|
||||
task_type=AITask.Type.FREE_VIDEO,
|
||||
model_config=self.model,
|
||||
idempotency_key=f"k-180-batch-{index}",
|
||||
status=AITask.Status.SUCCEEDED,
|
||||
)
|
||||
for index in range(1, 4)
|
||||
]
|
||||
|
||||
with patch("apps.ai.free_video.submit_free_video", side_effect=tasks) as submit:
|
||||
message, error = submit_confirmed_video(
|
||||
conversation=self.conversation,
|
||||
user=self.user,
|
||||
confirm_message=card,
|
||||
)
|
||||
|
||||
self.assertEqual(error, "")
|
||||
self.assertEqual(submit.call_count, 3)
|
||||
self.assertEqual(len(message.payload["segments"]), 6)
|
||||
self.assertEqual(len(message.payload["task_ids"]), 3)
|
||||
self.assertEqual([item["duration"] for item in message.payload["segments"]], [30] * 6)
|
||||
self.assertTrue(all(item.get("task_id") for item in message.payload["segments"][:3]))
|
||||
self.assertTrue(all(not item.get("task_id") for item in message.payload["segments"][3:]))
|
||||
self.assertEqual(message.payload["generation_spec"]["duration"], 180)
|
||||
prompts = [call.kwargs["params"]["prompt"] for call in submit.call_args_list]
|
||||
self.assertIn("第 1/6 段", prompts[0])
|
||||
self.assertIn("第 3/6 段", prompts[2])
|
||||
self.assertIn("商品只按已知真实用途", message.payload["generation_spec"]["prompt"])
|
||||
|
||||
def test_plan_without_video_prompt_is_rejected_without_emitting_cards(self):
|
||||
fake = FakeProvider([
|
||||
@@ -2111,7 +2506,8 @@ class VideoPlanAndConfirmTests(CreationAgentBaseTests):
|
||||
self.assertEqual(error, "")
|
||||
self.assertIsNotNone(message)
|
||||
self.assertIn("【视频预设】多色商品换款", prompt)
|
||||
self.assertIn("【预设执行层·点击换款·强制】", prompt)
|
||||
# 出片约束按换款方式分成「只出手」和「角色日常换款」两版;没选角色就是只出手。
|
||||
self.assertIn("【预设执行层·点击换款·只出手】", prompt)
|
||||
self.assertIn("手指轻触/点击", prompt)
|
||||
self.assertIn("不演剧情", prompt)
|
||||
self.assertIn("干净 match cut", prompt)
|
||||
@@ -2182,8 +2578,9 @@ class VideoParamParsingTests(TestCase):
|
||||
self.assertEqual(video_duration({"duration": "15 秒"}), 15)
|
||||
self.assertEqual(video_duration({"duration": "智能时长"}), SMART_DURATION)
|
||||
self.assertEqual(video_duration({}), SMART_DURATION)
|
||||
# 当前视频工作流支持完整短剧,超出预设上限时仍需夹住。
|
||||
self.assertEqual(video_duration({"duration": "99 秒"}), 60)
|
||||
# 当前视频工作流支持最长 180 秒,超出预设上限时仍需夹住。
|
||||
self.assertEqual(video_duration({"duration": "99 秒"}), 99)
|
||||
self.assertEqual(video_duration({"duration": "999 秒"}), 180)
|
||||
|
||||
def test_smart_duration_uses_final_plan_timestamp(self):
|
||||
timeline = [
|
||||
@@ -2195,6 +2592,15 @@ class VideoParamParsingTests(TestCase):
|
||||
self.assertEqual(infer_script_duration(timeline=timeline), 20)
|
||||
self.assertEqual(video_duration({"duration": "智能时长"}, timeline=timeline), 20)
|
||||
|
||||
def test_smart_duration_can_resolve_to_180_seconds(self):
|
||||
timeline = [
|
||||
{"start": 0, "end": 60, "stage": "第一章"},
|
||||
{"start": 60, "end": 120, "stage": "第二章"},
|
||||
{"start": 120, "end": 180, "stage": "第三章"},
|
||||
]
|
||||
self.assertEqual(infer_script_duration(timeline=timeline), 180)
|
||||
self.assertEqual(video_duration({"duration": "智能时长"}, timeline=timeline), 180)
|
||||
|
||||
def test_age_range_is_not_mistaken_for_video_duration(self):
|
||||
prompt = (
|
||||
"人声是18-22岁软甜少女音,语气自然。\n"
|
||||
@@ -2214,6 +2620,26 @@ class VideoParamParsingTests(TestCase):
|
||||
)
|
||||
self.assertEqual(plan_video_segments(60)[0]["duration"], 30)
|
||||
self.assertEqual(plan_video_segments(60)[1]["duration"], 30)
|
||||
self.assertEqual(
|
||||
[(item["start"], item["end"], item["duration"]) for item in plan_video_segments(180)],
|
||||
[(0, 30, 30), (30, 60, 30), (60, 90, 30), (90, 120, 30), (120, 150, 30), (150, 180, 30)],
|
||||
)
|
||||
|
||||
def test_timeout_notice_only_points_at_a_button_that_exists(self):
|
||||
self.assertIn("约 3 分钟", creation_agent_timeout_notice(180))
|
||||
self.assertIn("约 7 分钟", creation_agent_timeout_notice(420))
|
||||
self.assertIn("按这个继续", creation_agent_timeout_notice(180, has_open_gate=True))
|
||||
# 没有待确认的步骤卡时页面上并没有这个按钮,不能让用户去点找不到的东西
|
||||
self.assertNotIn("按这个继续", creation_agent_timeout_notice(180))
|
||||
self.assertIn("再发一次", creation_agent_timeout_notice(180))
|
||||
|
||||
def test_long_video_script_allows_two_second_slack(self):
|
||||
self.assertTrue(long_video_script_covers_requested_duration(118, 120))
|
||||
self.assertTrue(long_video_script_covers_requested_duration(120, 120))
|
||||
self.assertTrue(long_video_script_covers_requested_duration(180, 180))
|
||||
self.assertFalse(long_video_script_covers_requested_duration(None, 120))
|
||||
self.assertFalse(long_video_script_covers_requested_duration(90, 120))
|
||||
self.assertFalse(long_video_script_covers_requested_duration(150, 120))
|
||||
|
||||
def test_person_identity_guard_uses_real_reference_indexes(self):
|
||||
prompt = apply_person_identity_guard("base", [
|
||||
@@ -2227,6 +2653,18 @@ class VideoParamParsingTests(TestCase):
|
||||
self.assertIn("人物不得互换、遗漏", prompt)
|
||||
self.assertNotIn("必须保持同一人", prompt)
|
||||
|
||||
def test_product_reference_guard_keeps_skus_independent(self):
|
||||
prompt = apply_product_reference_guard("base", [
|
||||
{"type": "model", "label": "女主"},
|
||||
{"type": "product", "label": "黑色耳机"},
|
||||
{"type": "product", "label": "黑色耳机三视图"},
|
||||
{"type": "product", "label": "白色耳机"},
|
||||
])
|
||||
self.assertIn("参考图2=黑色耳机", prompt)
|
||||
self.assertIn("参考图4=白色耳机", prompt)
|
||||
self.assertIn("不得把多张商品图融合成新商品", prompt)
|
||||
self.assertIn("不同名称代表相互独立的商品或 SKU", prompt)
|
||||
|
||||
def test_model_label_maps_to_volcano_name(self):
|
||||
self.assertEqual(video_model_name({"model": "Seedance 2.0 Fast"}), "doubao-seedance-2-0-fast-260128")
|
||||
self.assertEqual(video_model_name({"model": "没见过的模型"}), DEFAULT_VIDEO_MODEL)
|
||||
@@ -2438,9 +2876,14 @@ class PresetGuidanceTests(CreationAgentBaseTests):
|
||||
"""预设不只是个名字,要把拍法约束一起给模型(契约 §6)。"""
|
||||
|
||||
def test_preset_guidance_reaches_the_system_prompt(self):
|
||||
person = Asset.objects.create(
|
||||
team=self.team, created_by=self.user, name="换装模特",
|
||||
asset_type=Asset.Type.IMAGE, category=Asset.Category.PERSON,
|
||||
)
|
||||
conversation = CreationConversation.objects.create(
|
||||
team=self.team, created_by=self.user, mode="video", preset="鱼眼换装", params={},
|
||||
memory={"person_source_ready": True},
|
||||
# 鱼眼换装必定要人物;没钉住人会先弹人物来源闸门,这一轮压根不会调模型。
|
||||
pinned_refs=[{"type": "character", "id": str(person.id), "name": person.name}],
|
||||
)
|
||||
fake = FakeProvider([_text_chunks("好")])
|
||||
with patch("apps.ai.creation_agent.build_provider", return_value=fake):
|
||||
@@ -2522,7 +2965,8 @@ class PresetGuidanceTests(CreationAgentBaseTests):
|
||||
)
|
||||
|
||||
self.assertTrue(prompt.startswith("【视频预设】点击换款"))
|
||||
self.assertIn("【预设执行层·点击换款·强制】", prompt)
|
||||
# 未选换款方式时出片走「只出手」那一版约束。
|
||||
self.assertIn("【预设执行层·点击换款·只出手】", prompt)
|
||||
self.assertIn("原始商品与 SKU 信息", prompt)
|
||||
self.assertIn("冲突的拍法一律忽略", prompt)
|
||||
self.assertIn("换款由手指点击触发", prompt)
|
||||
@@ -2555,7 +2999,9 @@ class PresetGuidanceTests(CreationAgentBaseTests):
|
||||
"旁白:今天给大家介绍这款猫粮,超好吃。",
|
||||
)
|
||||
self.assertIn("宠物拟人台词", prompt)
|
||||
self.assertNotIn("今天给大家介绍", prompt)
|
||||
# 只看被改写的正文:约束段本身会引用「今天给大家介绍/测评」当反例,不算漏改。
|
||||
body = prompt.split("【宠物拟人台词")[0]
|
||||
self.assertNotIn("今天给大家介绍", body)
|
||||
self.assertIn("我今天发现", prompt)
|
||||
self.assertEqual(apply_pet_dialogue_guard(conversation, prompt), prompt)
|
||||
|
||||
@@ -2608,6 +3054,39 @@ class PresetGuidanceTests(CreationAgentBaseTests):
|
||||
self.assertIn("已穿着", prompt)
|
||||
self.assertEqual(apply_clothing_video_guard(conversation, prompt), prompt)
|
||||
|
||||
def test_clothing_guard_is_not_triggered_by_character_wardrobe(self):
|
||||
"""卖的是精华液,人物穿的针织衫/阔腿裤不该把整条片判成服装视频。
|
||||
误判之后每轮还会再叠一段穿衣禁令,最终提示词里出现四五份。"""
|
||||
from .creation_agent import CLOTHING_GUARD_MARKER, apply_clothing_video_guard
|
||||
|
||||
conversation = CreationConversation.objects.create(
|
||||
team=self.team, created_by=self.user, mode="video", preset="达人口播种草", params={},
|
||||
pinned_refs=[{"type": "product", "id": "1", "name": "焕亮修护精华液", "category": "护肤"}],
|
||||
)
|
||||
prompt = (
|
||||
"林晚穿燕麦色坑条针织衫配米白阔腿裤,周野穿浅灰字母卫衣,"
|
||||
"三人围在梳妆台前试用精华液。"
|
||||
)
|
||||
|
||||
once = apply_clothing_video_guard(conversation, prompt)
|
||||
|
||||
self.assertNotIn(CLOTHING_GUARD_MARKER, once)
|
||||
self.assertEqual(once, prompt)
|
||||
|
||||
def test_clothing_guard_stays_single_across_repeated_passes(self):
|
||||
"""改写规则会动到约束段自身,按原文去重会失效 —— 反复调用必须仍然只有一段。"""
|
||||
from .creation_agent import CLOTHING_GUARD_MARKER, apply_clothing_video_guard
|
||||
|
||||
conversation = CreationConversation.objects.create(
|
||||
team=self.team, created_by=self.user, mode="video", preset="达人口播种草", params={},
|
||||
pinned_refs=[{"type": "product", "id": "1", "name": "法式收腰连衣裙", "category": "女装"}],
|
||||
)
|
||||
prompt = apply_clothing_video_guard(conversation, "0-3秒:女主把这件衣服穿上。")
|
||||
for _ in range(3):
|
||||
prompt = apply_clothing_video_guard(conversation, prompt)
|
||||
|
||||
self.assertEqual(prompt.count(CLOTHING_GUARD_MARKER), 1)
|
||||
|
||||
def test_clothing_video_guard_skips_non_apparel(self):
|
||||
from .creation_agent import apply_clothing_video_guard, CLOTHING_VIDEO_NO_DRESSING_GUARD
|
||||
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
"""全能创作 · 会话与消息底座(契约 §1/§3)。"""
|
||||
from django.test import TestCase
|
||||
from unittest.mock import patch
|
||||
|
||||
from django.test import TestCase, override_settings
|
||||
from rest_framework.test import APIClient
|
||||
|
||||
from apps.accounts.models import Team, TeamMember, User
|
||||
@@ -208,6 +210,79 @@ class GenerationBackfillTests(TestCase):
|
||||
self.assertEqual(message.payload["assets"][0]["label"], "第 1 段")
|
||||
self.assertEqual(message.payload["assets"][0]["url"], "https://cdn.example/segment-1.mp4")
|
||||
|
||||
@override_settings(FREE_VIDEO_MAX_CONCURRENT=3)
|
||||
def test_long_video_submits_next_wave_after_first_wave_finishes(self):
|
||||
first_wave = []
|
||||
segments = []
|
||||
for index in range(1, 7):
|
||||
segment = {"index": index, "start": (index - 1) * 30, "end": index * 30, "duration": 30}
|
||||
if index <= 3:
|
||||
task = AITask.objects.create(
|
||||
team=self.team,
|
||||
created_by=self.user,
|
||||
task_type=AITask.Type.FREE_VIDEO,
|
||||
model_config=self.model,
|
||||
status=AITask.Status.SUCCEEDED,
|
||||
idempotency_key=f"k-wave-first-{index}",
|
||||
)
|
||||
self._asset(task, url=f"https://cdn.example/segment-{index}.mp4")
|
||||
first_wave.append(task)
|
||||
segment["task_id"] = str(task.id)
|
||||
segments.append(segment)
|
||||
message = append_message(
|
||||
self.conversation,
|
||||
role="assistant",
|
||||
kind=CreationMessage.Kind.GENERATING,
|
||||
task=first_wave[0],
|
||||
payload={
|
||||
"kind": "video_segments",
|
||||
"task_id": str(first_wave[0].id),
|
||||
"task_ids": [str(task.id) for task in first_wave],
|
||||
"total_duration": 180,
|
||||
"segments": segments,
|
||||
"generation_spec": {
|
||||
"prompt": "完整 180 秒连续脚本",
|
||||
"feature": "omni_create",
|
||||
"mode": "universal",
|
||||
"model": "doubao-seedance-2-5-260628",
|
||||
"aspect_ratio": "9:16",
|
||||
"resolution": "720p",
|
||||
"duration": 180,
|
||||
"generate_audio": True,
|
||||
"references": [],
|
||||
"seed": 42,
|
||||
},
|
||||
},
|
||||
)
|
||||
created = []
|
||||
|
||||
def fake_submit(*, team, user, params):
|
||||
task = AITask.objects.create(
|
||||
team=team,
|
||||
created_by=user,
|
||||
task_type=AITask.Type.FREE_VIDEO,
|
||||
model_config=self.model,
|
||||
status=AITask.Status.SUBMITTED,
|
||||
idempotency_key=f"k-wave-second-{len(created) + 1}",
|
||||
request_payload=params,
|
||||
)
|
||||
created.append(task)
|
||||
return task
|
||||
|
||||
with patch("apps.ai.free_video.submit_free_video", side_effect=fake_submit) as submit:
|
||||
self.assertEqual(sync_generating_messages(self.conversation), 1)
|
||||
|
||||
message.refresh_from_db()
|
||||
self.assertEqual(submit.call_count, 3)
|
||||
self.assertEqual(message.kind, CreationMessage.Kind.GENERATING)
|
||||
self.assertEqual(message.payload["completed_segment_count"], 3)
|
||||
self.assertEqual(message.payload["submitted_segment_count"], 6)
|
||||
self.assertEqual(len(message.payload["task_ids"]), 6)
|
||||
self.assertTrue(all(item.get("task_id") for item in message.payload["segments"]))
|
||||
self.assertTrue(all(task.request_payload["seed"] == 42 for task in created))
|
||||
self.assertTrue(all(task.request_payload["omni_segment"]["message_id"] == str(message.id) for task in created))
|
||||
self.assertIn("第 4/6 段", created[0].request_payload["prompt"])
|
||||
|
||||
def test_retrieve_backfills_before_returning_messages(self):
|
||||
task = self._task(AITask.Status.SUCCEEDED, key="k-api")
|
||||
self._asset(task, url="https://cdn.example/api.png")
|
||||
|
||||
@@ -112,7 +112,8 @@ class ResolveRefsTests(TestCase):
|
||||
self.assertEqual(entry["review_status"], "active")
|
||||
self.assertEqual(entry["review_remote_id"], "R-1")
|
||||
|
||||
def test_model_ref_includes_portrait_and_triview(self):
|
||||
def test_model_ref_prefers_triview_over_portrait(self):
|
||||
"""一人一张图。形象图和三视图衣服经常不一致,两张一起送会被当成多出来的角色。"""
|
||||
portrait = _image_asset(self.team, self.user, "形象图", Asset.Category.MODEL_PORTRAIT, url="https://cdn/p.jpg")
|
||||
triview = _image_asset(self.team, self.user, "三视图", Asset.Category.TRI_VIEW, url="https://cdn/t.jpg")
|
||||
model = Model.objects.create(
|
||||
@@ -121,7 +122,8 @@ class ResolveRefsTests(TestCase):
|
||||
)
|
||||
|
||||
resolved = resolve_refs(self.team, [{"type": "model", "id": str(model.id)}])
|
||||
self.assertEqual([r["url"] for r in resolved.references], ["https://cdn/p.jpg", "https://cdn/t.jpg"])
|
||||
self.assertEqual([r["url"] for r in resolved.references], ["https://cdn/t.jpg"])
|
||||
self.assertEqual(resolved.references[0]["label"], "小夏")
|
||||
|
||||
def test_model_falls_back_to_portrait_when_no_triview(self):
|
||||
portrait = _image_asset(self.team, self.user, "形象图2", Asset.Category.MODEL_PORTRAIT, url="https://cdn/p2.jpg")
|
||||
@@ -168,7 +170,7 @@ class ResolveRefsTests(TestCase):
|
||||
resolved = resolve_refs(self.team, [{"type": "model", "id": str(model.id)}])
|
||||
self.assertIn("性别:男", resolved.facts_text)
|
||||
|
||||
def test_character_triview_is_attached_when_paired(self):
|
||||
def test_character_uses_triview_only_when_paired(self):
|
||||
self.person.metadata = {}
|
||||
self.person.save(update_fields=["metadata"])
|
||||
_image_asset(
|
||||
@@ -177,25 +179,39 @@ class ResolveRefsTests(TestCase):
|
||||
metadata={"triview_of": str(self.person.id)},
|
||||
)
|
||||
resolved = resolve_refs(self.team, [{"type": "character", "id": str(self.person.id)}])
|
||||
urls = [r["url"] for r in resolved.references]
|
||||
self.assertIn("https://cdn/person.jpg", urls)
|
||||
self.assertIn("https://cdn/person-tri.jpg", urls)
|
||||
self.assertEqual([r["url"] for r in resolved.references], ["https://cdn/person-tri.jpg"])
|
||||
self.assertEqual(len(resolved.references), 1)
|
||||
|
||||
def test_missing_triview_does_not_mark_ref_missing(self):
|
||||
resolved = resolve_refs(self.team, [{"type": "character", "id": str(self.person.id)}])
|
||||
self.assertEqual(resolved.missing, [])
|
||||
self.assertEqual([r["url"] for r in resolved.references], ["https://cdn/person.jpg"])
|
||||
|
||||
def test_product_triview_is_attached_when_present(self):
|
||||
def test_product_uses_triview_only_when_present(self):
|
||||
_image_asset(
|
||||
self.team, self.user, "商品三视图", Asset.Category.PRODUCT_IMAGE,
|
||||
url="https://cdn/prod-tri.jpg",
|
||||
metadata={"product_id": str(self.product.id), "view": "three_view"},
|
||||
)
|
||||
resolved = resolve_refs(self.team, [{"type": "product", "id": str(self.product.id)}])
|
||||
urls = [r["url"] for r in resolved.references]
|
||||
self.assertIn("https://cdn/prod.jpg", urls)
|
||||
self.assertIn("https://cdn/prod-tri.jpg", urls)
|
||||
self.assertEqual([r["url"] for r in resolved.references], ["https://cdn/prod-tri.jpg"])
|
||||
self.assertEqual(resolved.references[0]["label"], "净颜精华")
|
||||
|
||||
def test_latest_triview_wins_when_several_exist(self):
|
||||
portrait = _image_asset(self.team, self.user, "形象图4", Asset.Category.MODEL_PORTRAIT, url="https://cdn/p4.jpg")
|
||||
old = _image_asset(self.team, self.user, "旧三视图", Asset.Category.TRI_VIEW, url="https://cdn/old-t.jpg")
|
||||
new = _image_asset(self.team, self.user, "新三视图", Asset.Category.TRI_VIEW, url="https://cdn/new-t.jpg")
|
||||
old.created_at = old.created_at.replace(year=2024)
|
||||
old.save(update_fields=["created_at"])
|
||||
model = Model.objects.create(
|
||||
team=self.team, created_by=self.user, name="周野",
|
||||
portrait_asset=portrait, triview_asset=old,
|
||||
)
|
||||
new.metadata = {"triview_of": str(portrait.id)}
|
||||
new.save(update_fields=["metadata"])
|
||||
|
||||
resolved = resolve_refs(self.team, [{"type": "model", "id": str(model.id)}])
|
||||
self.assertEqual([r["url"] for r in resolved.references], ["https://cdn/new-t.jpg"])
|
||||
|
||||
def test_product_facts_text_without_selling_points_still_has_title(self):
|
||||
bare = Product.objects.create(team=self.team, created_by=self.user, title="裸商品")
|
||||
|
||||
@@ -1,35 +1,33 @@
|
||||
"""实体提取已经改成**本地**拆脚本:不调模型、不扣积分、不换模型重试。
|
||||
|
||||
这份文件原来锁的是「调模型 + 失败换模型 + 按 token 结算」那条链路(AIModelAttempt、
|
||||
CreditLedger、model_routing_v1)。那条链路已经随 submit_extract_entities 改成本地提取一起下线,
|
||||
`run_extract_entities_task` 现在没有任何地方会入队。所以这里改为锁住真正在跑的契约:
|
||||
提取必须零成本、不留模型调用痕迹,脚本缺 entities 时本地补齐。
|
||||
|
||||
落库内容本身(cast/scenes/entity_refs 回填)由 apps/projects/tests.py 的本地提取用例覆盖。
|
||||
"""
|
||||
|
||||
from decimal import Decimal
|
||||
from unittest.mock import Mock, patch
|
||||
from unittest.mock import patch
|
||||
|
||||
import requests
|
||||
from django.test import TestCase
|
||||
|
||||
from apps.accounts.models import Team, TeamMember, User
|
||||
from apps.ai.models import AITask, ModelConfig, ModelProvider
|
||||
from apps.ai.services import run_extract_entities_task, submit_extract_entities
|
||||
from apps.ai.services import submit_extract_entities
|
||||
from apps.billing.models import CreditAccount, CreditLedger
|
||||
from apps.products.models import Product
|
||||
from apps.projects.models import Project, ScriptSegment, ScriptVersion
|
||||
|
||||
|
||||
def _metadata(*, outbound=True, base_cost="0.50"):
|
||||
return {
|
||||
"routing": {"fallback_on_failure": outbound, "fallback_candidate": True},
|
||||
"capabilities": {
|
||||
"operations": ["chat"],
|
||||
"features": ["streaming", "structured_output"],
|
||||
},
|
||||
"pricing": {"base_cost_yuan": base_cost},
|
||||
}
|
||||
|
||||
|
||||
class EntityExtractionRoutingTests(TestCase):
|
||||
class LocalEntityExtractionTests(TestCase):
|
||||
def setUp(self):
|
||||
ModelConfig.objects.filter(capability=ModelConfig.Capability.TEXT).update(
|
||||
status=ModelConfig.Status.DISABLED
|
||||
)
|
||||
self.user = User.objects.create_user(username="entity-routing", password="x")
|
||||
self.team = Team.objects.create(name="Entity Routing", owner=self.user)
|
||||
self.user = User.objects.create_user(username="entity-local", password="x")
|
||||
self.team = Team.objects.create(name="Entity Local", owner=self.user)
|
||||
TeamMember.objects.create(team=self.team, user=self.user, role=TeamMember.Role.OWNER)
|
||||
CreditAccount.objects.create(team=self.team, balance=Decimal("1000"))
|
||||
product = Product.objects.create(team=self.team, created_by=self.user, title="测试商品")
|
||||
@@ -37,222 +35,77 @@ class EntityExtractionRoutingTests(TestCase):
|
||||
team=self.team,
|
||||
created_by=self.user,
|
||||
product=product,
|
||||
name="实体提取路由项目",
|
||||
name="实体提取项目",
|
||||
metadata={"cast": ["旧角色"], "entities_extracted": False},
|
||||
)
|
||||
provider = ModelProvider.objects.create(
|
||||
name="entity-local-provider",
|
||||
display_name="entity-local-provider",
|
||||
status=ModelProvider.Status.ACTIVE,
|
||||
)
|
||||
ModelConfig.objects.create(
|
||||
provider=provider,
|
||||
name="entity-local-model",
|
||||
display_name="entity-local-model",
|
||||
capability=ModelConfig.Capability.TEXT,
|
||||
endpoint="chat/completions",
|
||||
unit_price=Decimal("10"),
|
||||
status=ModelConfig.Status.ACTIVE,
|
||||
)
|
||||
self.script = ScriptVersion.objects.create(
|
||||
project=self.project,
|
||||
title="脚本",
|
||||
content="结构化脚本",
|
||||
is_adopted=True,
|
||||
metadata={"entities": [
|
||||
{"id": "c1", "type": "character", "name": "女主", "visual_prompt": "都市女主"},
|
||||
{"id": "s1", "type": "scene", "name": "客厅", "visual_prompt": "现代客厅"},
|
||||
]},
|
||||
)
|
||||
self.segment = ScriptSegment.objects.create(
|
||||
ScriptSegment.objects.create(
|
||||
script_version=self.script,
|
||||
sort_order=0,
|
||||
narration="女主在客厅展示商品",
|
||||
visual_prompt="女主站在客厅",
|
||||
entity_refs=["old"],
|
||||
)
|
||||
self.provider_mocks = {}
|
||||
patch("apps.ai.services.get_text_provider", side_effect=self._provider_for).start()
|
||||
patch("apps.ai.tasks.extract_entities_task.delay").start()
|
||||
self.addCleanup(patch.stopall)
|
||||
|
||||
def provider(self, name, priority):
|
||||
return ModelProvider.objects.create(
|
||||
name=name,
|
||||
display_name=name,
|
||||
status=ModelProvider.Status.ACTIVE,
|
||||
metadata={"routing": {"fallback_priority": priority}},
|
||||
)
|
||||
def test_extraction_is_free_and_never_calls_a_model(self):
|
||||
with patch("apps.ai.services.get_text_provider") as get_provider:
|
||||
task = submit_extract_entities(project=self.project, user=self.user)
|
||||
|
||||
def model(self, provider, name, *, outbound=True, base_cost="0.50", is_default=False):
|
||||
return ModelConfig.objects.create(
|
||||
provider=provider,
|
||||
name=name,
|
||||
display_name=name,
|
||||
capability=ModelConfig.Capability.TEXT,
|
||||
endpoint="chat/completions",
|
||||
unit_price=Decimal("10"),
|
||||
status=ModelConfig.Status.ACTIVE,
|
||||
is_default=is_default,
|
||||
metadata=_metadata(outbound=outbound, base_cost=base_cost),
|
||||
)
|
||||
get_provider.assert_not_called()
|
||||
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
|
||||
self.assertEqual(task.request_payload.get("mode"), "local")
|
||||
self.assertEqual(task.estimated_cost, Decimal("0"))
|
||||
self.assertEqual(task.actual_cost, Decimal("0"))
|
||||
self.assertEqual(task.base_cost, Decimal("0"))
|
||||
# 不走模型就不该留下换模型痕迹,也不该有任何积分流水。
|
||||
self.assertFalse(task.model_attempts.exists())
|
||||
self.assertFalse(CreditLedger.objects.filter(task=task).exists())
|
||||
|
||||
@staticmethod
|
||||
def valid_events():
|
||||
return [
|
||||
{
|
||||
"type": "delta",
|
||||
"text": (
|
||||
'{"entities":['
|
||||
'{"id":"c1","type":"character","name":"女主","visual_prompt":"都市女主"},'
|
||||
'{"id":"s1","type":"scene","name":"客厅","visual_prompt":"现代客厅"}'
|
||||
'],"segments":[{"index":0,"entity_refs":["c1","s1"]}]}'
|
||||
),
|
||||
},
|
||||
{"type": "done"},
|
||||
]
|
||||
def test_entities_are_persisted_and_old_metadata_is_replaced(self):
|
||||
submit_extract_entities(project=self.project, user=self.user)
|
||||
|
||||
@staticmethod
|
||||
def invalid_events():
|
||||
return [{"type": "delta", "text": "这是一段没有 JSON 的解释"}, {"type": "done"}]
|
||||
|
||||
@classmethod
|
||||
def _new_provider_mock(cls):
|
||||
provider = Mock()
|
||||
provider.chat_completion_stream.return_value = cls.valid_events()
|
||||
return provider
|
||||
|
||||
def _provider_for(self, model):
|
||||
return self.provider_mocks.setdefault(model.id, self._new_provider_mock())
|
||||
|
||||
def submit(self):
|
||||
return submit_extract_entities(project=self.project, user=self.user)
|
||||
|
||||
@staticmethod
|
||||
def ledger_count(task, ledger_type):
|
||||
return CreditLedger.objects.filter(task=task, ledger_type=ledger_type).count()
|
||||
|
||||
def test_first_success_records_streaming_structured_attempt_and_persists_entities(self):
|
||||
primary = self.model(self.provider("entity-primary", 20), "entity-primary", is_default=True)
|
||||
task = self.submit()
|
||||
|
||||
run_extract_entities_task(task_id=str(task.id))
|
||||
|
||||
task.refresh_from_db()
|
||||
self.project.refresh_from_db()
|
||||
self.segment.refresh_from_db()
|
||||
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
|
||||
self.assertTrue(task.request_payload["model_routing_v1"])
|
||||
attempt = task.model_attempts.get()
|
||||
self.assertEqual(attempt.model_config_id, primary.id)
|
||||
self.assertEqual(attempt.operation, "chat")
|
||||
self.assertTrue(attempt.request_summary["streaming"])
|
||||
self.assertTrue(attempt.request_summary["structured_output"])
|
||||
self.assertEqual(attempt.request_summary["business_operation"], "entity_extract")
|
||||
self.assertEqual(attempt.public_model_name, "AirShelf Script")
|
||||
self.assertTrue(self.project.metadata["entities_extracted"])
|
||||
self.assertEqual(self.project.metadata["cast"], ["女主"])
|
||||
self.assertEqual(self.project.metadata["scenes"], ["客厅"])
|
||||
self.assertEqual(self.segment.entity_refs, ["c1", "s1"])
|
||||
call = self.provider_mocks[primary.id].chat_completion_stream.call_args
|
||||
self.assertGreater(call.kwargs["timeout"], 0)
|
||||
self.assertEqual(call.kwargs["temperature"], 0.3)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 1)
|
||||
metadata = self.project.metadata or {}
|
||||
self.assertTrue(metadata.get("entities_extracted"))
|
||||
self.assertIn("女主", metadata.get("cast") or [])
|
||||
self.assertNotIn("旧角色", metadata.get("cast") or [])
|
||||
|
||||
def test_invalid_structure_retries_then_falls_back_and_records_failed_call_costs(self):
|
||||
primary = self.model(
|
||||
self.provider("entity-fallback-primary", 100),
|
||||
"entity-primary",
|
||||
base_cost="0.25",
|
||||
def test_script_without_segments_is_rejected_before_any_task(self):
|
||||
empty_script = ScriptVersion.objects.create(
|
||||
project=self.project, title="空脚本", content="{}", is_adopted=True,
|
||||
)
|
||||
candidate = self.model(
|
||||
self.provider("entity-fallback-candidate", 10),
|
||||
"entity-candidate",
|
||||
outbound=False,
|
||||
base_cost="0.75",
|
||||
self.script.is_adopted = False
|
||||
self.script.save(update_fields=["is_adopted"])
|
||||
|
||||
with self.assertRaises(ValueError):
|
||||
submit_extract_entities(project=self.project, user=self.user)
|
||||
|
||||
self.assertFalse(
|
||||
AITask.objects.filter(
|
||||
project=self.project, task_type=AITask.Type.ENTITY_EXTRACTION
|
||||
).exists()
|
||||
)
|
||||
primary_mock = self._new_provider_mock()
|
||||
primary_mock.chat_completion_stream.side_effect = [
|
||||
self.invalid_events(),
|
||||
self.invalid_events(),
|
||||
self.invalid_events(),
|
||||
]
|
||||
self.provider_mocks[primary.id] = primary_mock
|
||||
task = self.submit()
|
||||
|
||||
run_extract_entities_task(task_id=str(task.id))
|
||||
|
||||
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, primary.id, candidate.id],
|
||||
)
|
||||
self.assertEqual(
|
||||
[a.status for a in attempts],
|
||||
["failed", "failed", "failed", "succeeded"],
|
||||
)
|
||||
self.assertTrue(attempts[-1].is_fallback)
|
||||
self.assertTrue(attempts[0].response_summary["validation_failed"])
|
||||
self.assertEqual(task.base_cost, Decimal("1.5000"))
|
||||
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)
|
||||
|
||||
def test_all_candidates_fail_releases_and_preserves_old_project_metadata(self):
|
||||
primary = self.model(self.provider("entity-all-primary", 100), "entity-primary")
|
||||
fallback_1 = self.model(
|
||||
self.provider("entity-all-fallback-1", 10), "entity-fallback-1", outbound=False
|
||||
)
|
||||
fallback_2 = self.model(
|
||||
self.provider("entity-all-fallback-2", 20), "entity-fallback-2", outbound=False
|
||||
)
|
||||
for model in (primary, fallback_1, fallback_2):
|
||||
provider = self._new_provider_mock()
|
||||
provider.chat_completion_stream.side_effect = requests.ConnectionError("offline")
|
||||
self.provider_mocks[model.id] = provider
|
||||
task = self.submit()
|
||||
|
||||
run_extract_entities_task(task_id=str(task.id))
|
||||
|
||||
task.refresh_from_db()
|
||||
self.project.refresh_from_db()
|
||||
self.segment.refresh_from_db()
|
||||
self.assertEqual(task.status, AITask.Status.FAILED)
|
||||
self.assertEqual(task.model_attempts.count(), 4)
|
||||
self.assertEqual(self.project.metadata["cast"], ["旧角色"])
|
||||
self.assertFalse(self.project.metadata["entities_extracted"])
|
||||
self.assertEqual(self.segment.entity_refs, ["old"])
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.CHARGE), 0)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 1)
|
||||
self.assertEqual(CreditAccount.objects.get(team=self.team).reserved_balance, Decimal("0"))
|
||||
|
||||
def test_direct_doubao_primary_retries_but_does_not_switch_when_outbound_disabled(self):
|
||||
primary = self.model(
|
||||
self.provider("doubao", 10), "doubao-seed-2-0-pro-260215", outbound=False
|
||||
)
|
||||
candidate = self.model(
|
||||
self.provider("entity-unused-candidate", 20), "entity-unused", outbound=False
|
||||
)
|
||||
provider = self._new_provider_mock()
|
||||
provider.chat_completion_stream.side_effect = requests.ConnectionError("offline")
|
||||
self.provider_mocks[primary.id] = provider
|
||||
task = self.submit()
|
||||
|
||||
run_extract_entities_task(task_id=str(task.id))
|
||||
|
||||
task.refresh_from_db()
|
||||
attempts = list(task.model_attempts.all())
|
||||
self.assertEqual(task.status, AITask.Status.FAILED)
|
||||
self.assertEqual(
|
||||
[attempt.model_config_id for attempt in attempts],
|
||||
[primary.id, primary.id, primary.id],
|
||||
)
|
||||
self.assertFalse(any(attempt.model_config_id == candidate.id for attempt in attempts))
|
||||
self.assertTrue(all(attempt.public_model_name == primary.display_name for attempt in attempts))
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RESERVE), 1)
|
||||
self.assertEqual(self.ledger_count(task, CreditLedger.Type.RELEASE), 1)
|
||||
|
||||
def test_uses_database_default_text_model_not_older_doubao_pin(self):
|
||||
older = self.model(self.provider("volcengine-old", 20), "doubao-seed-2-0-pro-260215")
|
||||
default = self.model(
|
||||
self.provider("volcengine-default", 10),
|
||||
"doubao-seed-2-1-pro-260628",
|
||||
is_default=True,
|
||||
)
|
||||
task = self.submit()
|
||||
|
||||
run_extract_entities_task(task_id=str(task.id))
|
||||
|
||||
task.refresh_from_db()
|
||||
attempt = task.model_attempts.get()
|
||||
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
|
||||
self.assertEqual(attempt.model_config_id, default.id)
|
||||
self.assertNotEqual(attempt.model_config_id, older.id)
|
||||
self.assertEqual(task.request_payload["model"], "doubao-seed-2-1-pro-260628")
|
||||
|
||||
self.assertFalse(empty_script.segments.exists())
|
||||
|
||||
@@ -261,6 +261,21 @@ class SubmitFreeVideoTests(TestCase):
|
||||
self.assertEqual(task.request_payload["estimated_tokens"], tokens)
|
||||
self.assertEqual(task.request_payload["feature"], "free_video")
|
||||
|
||||
def test_oversized_seed_is_clamped_before_volcano(self):
|
||||
"""全能创作长视频用会话 id 前 8 位十六进制当 seed,最大约 42.9 亿。
|
||||
Seedance 2.5 r2v 只接受 <= 2147483647,用户这次翻车的值就是 3857139456。"""
|
||||
from apps.ai.free_video import VIDEO_SEED_MAX, normalize_video_seed
|
||||
|
||||
raw = 3857139456
|
||||
task = submit_free_video(
|
||||
team=self.team, user=self.user, params=self._params(seed=raw)
|
||||
)
|
||||
sent = self.provider.create_video_task.call_args.kwargs["seed"]
|
||||
self.assertEqual(task.request_payload["seed"], normalize_video_seed(raw))
|
||||
self.assertEqual(sent, normalize_video_seed(raw))
|
||||
self.assertLessEqual(sent, VIDEO_SEED_MAX)
|
||||
self.assertNotEqual(sent, raw)
|
||||
|
||||
def test_insufficient_balance_leaves_nothing(self):
|
||||
CreditAccount.objects.filter(team=self.team).update(balance="0.0100")
|
||||
with self.assertRaisesMessage(ValueError, "余额不足"):
|
||||
|
||||
@@ -11,7 +11,7 @@ from django.test import TestCase
|
||||
|
||||
from apps.accounts.models import Team, User
|
||||
from apps.ai.free_video import build_content_items
|
||||
from apps.assets.models import Asset, AssetFile
|
||||
from apps.assets.models import Asset, AssetFile, Model as AssetModel
|
||||
from apps.assets.review import reference_review_state
|
||||
|
||||
|
||||
@@ -88,6 +88,24 @@ class ReferenceReviewStateTests(TestCase):
|
||||
self.assertEqual(reference_review_state(_asset(self.team, source=Asset.Source.UPLOAD)), "allowed")
|
||||
|
||||
|
||||
class VideoSeedRangeTests(TestCase):
|
||||
"""火山 seed 上限是 int32 正数上限。长视频各段共用的确定性 seed 取会话 id 前 8 位
|
||||
十六进制,最大约 42.9 亿 —— 不收口整条请求会被 InvalidParameter 拒掉。"""
|
||||
|
||||
def test_seed_is_clamped_into_volcano_range(self):
|
||||
from apps.ai.free_video import VIDEO_SEED_MAX, normalize_video_seed
|
||||
|
||||
self.assertEqual(VIDEO_SEED_MAX, 2147483647)
|
||||
self.assertLessEqual(normalize_video_seed(3857139456), VIDEO_SEED_MAX)
|
||||
self.assertLessEqual(normalize_video_seed(0xFFFFFFFF), VIDEO_SEED_MAX)
|
||||
self.assertEqual(normalize_video_seed(12345), 12345)
|
||||
# 同一个值每次折算结果一致,跨段一致性不会因此漂移
|
||||
self.assertEqual(normalize_video_seed(3857139456), normalize_video_seed(3857139456))
|
||||
self.assertEqual(normalize_video_seed(None), -1)
|
||||
self.assertEqual(normalize_video_seed(-1), -1)
|
||||
self.assertEqual(normalize_video_seed("abc"), -1)
|
||||
|
||||
|
||||
class AssetReferenceBuildTests(TestCase):
|
||||
"""source=asset 分支:URL 解析 + 闸门拦截 + @label 映射。"""
|
||||
|
||||
@@ -175,6 +193,44 @@ class AssetReferenceBuildTests(TestCase):
|
||||
self._build(asset)
|
||||
self.assertIn("不存在", str(ctx.exception))
|
||||
|
||||
def test_official_model_asset_is_usable_across_teams(self):
|
||||
"""官方模特挂在平台团队名下。按 team 过滤会让模特库里明明在的模特,
|
||||
出片时报「不存在或已被删除」—— 放行范围与视频复刻保持一致。"""
|
||||
platform = Team.objects.create(
|
||||
name="PLATFORM",
|
||||
owner=User.objects.create_user(username="platform-owner", password="p"),
|
||||
)
|
||||
portrait = _asset(platform, source=Asset.Source.AI_GENERATED)
|
||||
AssetModel.objects.create(
|
||||
team=platform,
|
||||
name="阳光运动 · 周野",
|
||||
source=AssetModel.Source.AI,
|
||||
portrait_asset=portrait,
|
||||
is_official=True,
|
||||
)
|
||||
|
||||
built = self._build(portrait, label="阳光运动 · 周野")
|
||||
|
||||
self.assertEqual(built["image_n"], 1)
|
||||
self.assertEqual(built["content_items"][0]["image_url"]["url"], "http://tos/1.png")
|
||||
|
||||
def test_non_official_model_asset_stays_team_scoped(self):
|
||||
"""只放行官方模特。别的团队自建模特的图仍然不可跨团队引用。"""
|
||||
stranger = Team.objects.create(
|
||||
name="STRANGER",
|
||||
owner=User.objects.create_user(username="stranger-owner", password="p"),
|
||||
)
|
||||
portrait = _asset(stranger, source=Asset.Source.AI_GENERATED)
|
||||
AssetModel.objects.create(
|
||||
team=stranger,
|
||||
name="别人家的模特",
|
||||
source=AssetModel.Source.AI,
|
||||
portrait_asset=portrait,
|
||||
)
|
||||
with self.assertRaises(ValueError) as ctx:
|
||||
self._build(portrait)
|
||||
self.assertIn("不存在", str(ctx.exception))
|
||||
|
||||
def test_keyframe_rejects_non_image(self):
|
||||
asset = _asset(self.team, source=Asset.Source.AI_GENERATED)
|
||||
asset.asset_type = Asset.Type.VIDEO
|
||||
|
||||
@@ -19,6 +19,20 @@ from apps.billing.models import CreditAccount
|
||||
from apps.products.models import Product
|
||||
from apps.projects.models import Project, ScriptVersion
|
||||
|
||||
# 这组只验「来源」徽标记得对不对,但稿子仍要真的过脚本自检(口播字数、画面 ≥110 字、秒级分镜),
|
||||
# 否则生成会先以「结果处理失败」收场,根本走不到落库那一步。
|
||||
VALID_NARRATION = (
|
||||
"下午三点工位犯困,键盘都敲不利索。我以前总是硬扛,越扛脑子越乱。"
|
||||
"后来会先倒杯热茶,第一口是回甘不是苦,桌上的文件终于能看进去,"
|
||||
"整个人才慢慢醒过来,状态也稳下来。"
|
||||
)
|
||||
VALID_VISUAL = (
|
||||
"0-3s:近景;平视;固定机位;女主在画面右侧对镜头抬眼,工位键盘还亮着,窗外是下午的光\n"
|
||||
"3-8s:中近景;侧后方过肩;手持跟拍到桌面;右手把茶包放进盛了热水的玻璃杯,水面起蒸汽\n"
|
||||
"8-12s:特写;俯拍杯口;缓慢推近;茶汤从浅金慢慢变深,标签贴在杯沿\n"
|
||||
"12-15s:中近景;平视;拉回;女主双手捧杯喝一口,眉头松开肩膀塌下来"
|
||||
)
|
||||
|
||||
|
||||
class ScriptEntrySourceTests(TransactionTestCase):
|
||||
reset_sequences = True
|
||||
@@ -61,7 +75,7 @@ class ScriptEntrySourceTests(TransactionTestCase):
|
||||
def _provider(self, _model):
|
||||
raw = json.dumps(
|
||||
{
|
||||
"hook": "开场钩子",
|
||||
"hook": "下午三点工位犯困,键盘都敲不利索",
|
||||
"tone": "自然",
|
||||
"aspect_ratio": "9:16",
|
||||
"total_duration": 15,
|
||||
@@ -72,7 +86,7 @@ class ScriptEntrySourceTests(TransactionTestCase):
|
||||
"segments": [
|
||||
{
|
||||
"index": 0, "duration": 15, "role": "钩子",
|
||||
"narration": "全新脚本口播", "visual": "女主展示商品",
|
||||
"narration": VALID_NARRATION, "visual": VALID_VISUAL,
|
||||
"speaker": "女主", "product_exposure": "展示",
|
||||
"entity_refs": ["c1"], "dialogue": [],
|
||||
}
|
||||
|
||||
@@ -14,6 +14,21 @@ from apps.products.models import Product
|
||||
from apps.projects.models import Project, ScriptSegment, ScriptVersion
|
||||
|
||||
|
||||
# 15 秒一镜的合法口播与画面:口播要到 narration_floor(15),画面要 ≥110 字且带足秒级分镜、
|
||||
# 景别与运镜标记。路由测试只关心换模型和结算,但稿子仍要真的过得了脚本自检。
|
||||
VALID_NARRATION = (
|
||||
"下午三点工位犯困,键盘都敲不利索。我以前总是硬扛,越扛脑子越乱。"
|
||||
"后来会先倒杯热茶,第一口是回甘不是苦,桌上的文件终于能看进去,"
|
||||
"整个人才慢慢醒过来,状态也稳下来。"
|
||||
)
|
||||
VALID_VISUAL = (
|
||||
"0-3s:近景;平视;固定机位;女主在画面右侧对镜头抬眼,工位键盘还亮着,窗外是下午的光\n"
|
||||
"3-8s:中近景;侧后方过肩;手持跟拍到桌面;右手把茶包放进盛了热水的玻璃杯,水面起蒸汽\n"
|
||||
"8-12s:特写;俯拍杯口;缓慢推近;茶汤从浅金慢慢变深,标签贴在杯沿\n"
|
||||
"12-15s:中近景;平视;拉回;女主双手捧杯喝一口,眉头松开肩膀塌下来"
|
||||
)
|
||||
|
||||
|
||||
def _metadata(*, outbound=True, base_cost="0.50"):
|
||||
return {
|
||||
"routing": {"fallback_on_failure": outbound, "fallback_candidate": True},
|
||||
@@ -69,8 +84,10 @@ class ScriptStreamRoutingTests(TransactionTestCase):
|
||||
|
||||
@staticmethod
|
||||
def valid_draft():
|
||||
# 口播字数、画面字数和秒级分镜都必须真的过 assert_shot_density / assert_script_has_a_hook;
|
||||
# 写成「全新脚本口播」这种占位文本会被当成模型偷懒直接判失败,整条流就退化成 error。
|
||||
return {
|
||||
"hook": "开场钩子",
|
||||
"hook": "下午三点工位犯困,键盘都敲不利索",
|
||||
"tone": "自然",
|
||||
"aspect_ratio": "9:16",
|
||||
"total_duration": 15,
|
||||
@@ -89,8 +106,8 @@ class ScriptStreamRoutingTests(TransactionTestCase):
|
||||
"index": 0,
|
||||
"duration": 15,
|
||||
"role": "钩子",
|
||||
"narration": "全新脚本口播",
|
||||
"visual": "女主展示商品",
|
||||
"narration": VALID_NARRATION,
|
||||
"visual": VALID_VISUAL,
|
||||
"speaker": "女主",
|
||||
"product_exposure": "展示",
|
||||
"entity_refs": ["c1"],
|
||||
|
||||
@@ -1907,6 +1907,51 @@ class ChatStreamReasoningTests(SimpleTestCase):
|
||||
self.assertEqual([e["text"] for e in events if e["type"] == "reasoning"], ["先想想", "用户要4镜"])
|
||||
self.assertEqual("".join(e["text"] for e in events if e["type"] == "delta"), "正在生成脚本…")
|
||||
|
||||
def test_length_finish_reason_is_forwarded(self):
|
||||
"""撞 max_tokens 的那一刀必须能被上层看见,否则截断的 tool 参数会被当成模型偷懒反复重试。"""
|
||||
lines = [
|
||||
"data: " + json.dumps({"choices": [{"delta": {"content": "前半段"}}]}, ensure_ascii=False),
|
||||
"data: " + json.dumps({"choices": [{"delta": {}, "finish_reason": "length"}]}, ensure_ascii=False),
|
||||
"data: [DONE]",
|
||||
]
|
||||
prov = VolcanoArkProvider(api_key="k", base_url="http://x")
|
||||
with patch("apps.ai.providers.volcano.requests.post", return_value=_FakeStreamResp(lines)):
|
||||
events = list(prov.chat_completion_stream(model="m", messages=[{"role": "user", "content": "hi"}]))
|
||||
|
||||
self.assertEqual([e["type"] for e in events], ["delta", "finish", "done"])
|
||||
self.assertEqual(events[1]["reason"], "length")
|
||||
|
||||
|
||||
class VolcanoVideoSeedTests(SimpleTestCase):
|
||||
"""Seedance 2.5 r2v 的 seed 上限是 int32,超了会 InvalidParameter。"""
|
||||
|
||||
def test_create_video_task_clamps_seed_to_int32(self):
|
||||
captured = {}
|
||||
|
||||
class _Resp:
|
||||
ok = True
|
||||
status_code = 200
|
||||
text = "{}"
|
||||
|
||||
def json(self):
|
||||
return {"id": "ark-1"}
|
||||
|
||||
def _post(*_args, **kwargs):
|
||||
captured["body"] = kwargs["json"]
|
||||
return _Resp()
|
||||
|
||||
prov = VolcanoArkProvider(api_key="k", base_url="http://x")
|
||||
with patch("apps.ai.providers.volcano.requests.post", side_effect=_post):
|
||||
prov.create_video_task(
|
||||
model="doubao-seedance-2-5-260628",
|
||||
endpoint="contents/generations/tasks",
|
||||
prompt="test",
|
||||
seed=3857139456,
|
||||
)
|
||||
|
||||
self.assertEqual(captured["body"]["seed"], 3857139456 & 2147483647)
|
||||
self.assertLessEqual(captured["body"]["seed"], 2147483647)
|
||||
|
||||
|
||||
class WorkbenchAndUnreadTests(TestCase):
|
||||
"""R100(工作台记录后端持久化)+ R96(未读生成任务角标)+ R109(删除联动)端到端:
|
||||
|
||||
@@ -33,6 +33,7 @@ from .free_video import (
|
||||
IN_FLIGHT_STATUSES,
|
||||
RATIOS,
|
||||
RESOLUTIONS,
|
||||
normalize_video_seed,
|
||||
_reap_stale_free_video_tasks,
|
||||
model_duration_range,
|
||||
serialize_free_video_task,
|
||||
@@ -654,6 +655,21 @@ def video_replace_q() -> Q:
|
||||
return Q(request_payload__feature=FEATURE) | Q(request_payload__prompt__startswith=LEGACY_PROMPT_PREFIX)
|
||||
|
||||
|
||||
def not_video_replace_q() -> Q:
|
||||
"""「不是视频复刻」。不能写成 exclude(video_replace_q())。
|
||||
|
||||
JSON 键不存在时取值是 SQL NULL,NOT(NULL = 'x') 仍是 NULL,整行会被 exclude 一起筛掉 ——
|
||||
于是 request_payload 里没有 feature 的历史任务在列表和回收站里凭空消失。必须把「键缺失」
|
||||
显式算作不匹配。
|
||||
"""
|
||||
return (
|
||||
Q(request_payload__feature__isnull=True) | ~Q(request_payload__feature=FEATURE)
|
||||
) & (
|
||||
Q(request_payload__prompt__isnull=True)
|
||||
| ~Q(request_payload__prompt__startswith=LEGACY_PROMPT_PREFIX)
|
||||
)
|
||||
|
||||
|
||||
def _fresh_digest_source(payload: dict) -> dict | None:
|
||||
"""历史卡「原视频」用的原片快照。快照 URL 会过期,统一走 rehydrate_ref_urls 重新取长期直链。"""
|
||||
from .free_video import rehydrate_ref_urls
|
||||
@@ -1028,10 +1044,7 @@ def _legacy_start_pending_replace_shots(task):
|
||||
generate_audio = bool(payload.get("generate_audio", True))
|
||||
search_mode = str(payload.get("search_mode") or "off")
|
||||
feature = str(payload.get("feature") or FEATURE)
|
||||
try:
|
||||
seed = int(payload.get("seed") if payload.get("seed") is not None else -1)
|
||||
except (TypeError, ValueError):
|
||||
seed = -1
|
||||
seed = normalize_video_seed(payload.get("seed"))
|
||||
billed_duration = sum(int(item.get("seconds") or 0) for item in plan) or int(payload.get("duration") or 5)
|
||||
references = _seedance_references(payload)
|
||||
try:
|
||||
@@ -1238,10 +1251,7 @@ def _dispatch_next_replace_shot(task, shot: dict):
|
||||
locked.request_payload = next_payload
|
||||
locked.save(update_fields=["request_payload", "updated_at"])
|
||||
|
||||
try:
|
||||
seed = int(payload.get("seed") if payload.get("seed") is not None else -1)
|
||||
except (TypeError, ValueError):
|
||||
seed = -1
|
||||
seed = normalize_video_seed(payload.get("seed"))
|
||||
dispatched = _dispatch_free_video_provider(
|
||||
task=locked,
|
||||
built=built,
|
||||
@@ -1403,7 +1413,9 @@ def _concat_shot_media(urls: list[str]) -> bytes:
|
||||
"-c:v", "libx264", "-pix_fmt", "yuv420p", "-preset", "veryfast", "-threads", "2",
|
||||
"-c:a", "aac", "-b:a", "192k", "-movflags", "+faststart", str(output),
|
||||
]
|
||||
proc = subprocess.run(cmd, capture_output=True, timeout=180)
|
||||
# 180 秒成片最多 6 段,1080p 重编码明显可能超过旧的 180 秒固定超时。
|
||||
# 按片段数放宽,但仍保留上限式超时,避免异常 ffmpeg 永久占住 worker。
|
||||
proc = subprocess.run(cmd, capture_output=True, timeout=max(180, len(paths) * 120))
|
||||
if proc.returncode != 0 or not output.exists() or not output.stat().st_size:
|
||||
err = (proc.stderr or b"").decode("utf-8", errors="replace")[-500:]
|
||||
raise RuntimeError(f"镜头拼接失败:{err or 'ffmpeg 未产出文件'}")
|
||||
@@ -1605,10 +1617,7 @@ def _create_reviewing_task(*, team, user, params: dict, digest_pending: bool = F
|
||||
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
|
||||
seed = normalize_video_seed(params.get("seed"))
|
||||
request_payload = {
|
||||
"feature": FEATURE,
|
||||
"mode": "universal",
|
||||
|
||||
@@ -400,11 +400,13 @@ def _auto_pick_product_continuation(conversation: CreationConversation, request_
|
||||
|
||||
_STEP_CONTINUE_INSTRUCTIONS = {
|
||||
"strategy": (
|
||||
"用户已确认创作策略。现在只调用 write_plan 写方案卡(含完整 video_prompt 存档);"
|
||||
"用户已确认创作策略。现在只调用 write_plan 写方案卡并存档 video_prompt;"
|
||||
"超过 60 秒时只写紧凑章节稿(章节结构 + 每个 30 秒段关键镜头 + 交接状态),不要写制作级长文。"
|
||||
"不要再写策略,不要调用 write_prompt,不要出片。"
|
||||
),
|
||||
"plan": (
|
||||
"用户已确认视频方案。补齐完整 video_prompt 后直接调用 write_prompt 整理后台出片指令并给出积分确认卡;"
|
||||
"用户已确认视频方案。调用 write_prompt 整理后台出片指令并给出积分确认卡;"
|
||||
"超过 60 秒时长优先沿用方案里已存的 video_prompt,只补简短全片规则,不要重写成短片密度的制作长文。"
|
||||
"不要重写策略/方案,也不要直接出片。"
|
||||
),
|
||||
"prompt": (
|
||||
@@ -419,12 +421,12 @@ _STEP_REVISE_INSTRUCTIONS = {
|
||||
"写完即停,不要同轮 write_plan / write_prompt。"
|
||||
),
|
||||
"plan": (
|
||||
"用户要求修改视频方案。根据反馈只重新调用 write_plan 写一版修订方案"
|
||||
"(含完整 video_prompt 存档);写完即停,不要同轮 write_prompt 或出片。"
|
||||
"用户要求修改视频方案。根据反馈只重新调用 write_plan 写一版修订方案并更新 video_prompt;"
|
||||
"超过 60 秒时仍写紧凑章节稿,不要写制作级长文。写完即停,不要同轮 write_prompt 或出片。"
|
||||
),
|
||||
"prompt": (
|
||||
"用户要求修改出片细节。根据反馈只重新调用 write_prompt 整理一版修订指令;"
|
||||
"完成后给出积分确认卡,不要直接出片。"
|
||||
"长视频不要把逐镜再扩写成短片密度。完成后给出积分确认卡,不要直接出片。"
|
||||
),
|
||||
}
|
||||
|
||||
@@ -1147,9 +1149,15 @@ def _free_video_task_queryset(team):
|
||||
)
|
||||
|
||||
|
||||
def _not_omni_create_q() -> Q:
|
||||
"""「不是全能创作」。同样不能写成 exclude(request_payload__feature="omni_create")——
|
||||
JSON 键缺失时比较结果是 NULL,exclude 会把没写 feature 的历史任务一起筛没。"""
|
||||
return Q(request_payload__feature__isnull=True) | ~Q(request_payload__feature="omni_create")
|
||||
|
||||
|
||||
def _free_video_list_queryset(team, *, include_replace=False):
|
||||
"""正常任务流隐藏已从资产库删除的成品,但保留生成中/失败及无落库资产的历史任务。"""
|
||||
from .video_replace import video_replace_q
|
||||
from .video_replace import not_video_replace_q, video_replace_q
|
||||
|
||||
video_assets = Asset.objects.filter(origin_task_id=OuterRef("pk"), asset_type=Asset.Type.VIDEO)
|
||||
active_video_assets = video_assets.filter(is_deleted=False, purged_at__isnull=True)
|
||||
@@ -1168,13 +1176,13 @@ def _free_video_list_queryset(team, *, include_replace=False):
|
||||
if include_replace:
|
||||
return qs.filter(video_replace_q())
|
||||
# 全能创作也走 FREE_VIDEO 任务类型,但不能出现在自由生成任务流里。
|
||||
return qs.exclude(video_replace_q()).exclude(request_payload__feature="omni_create")
|
||||
return qs.filter(not_video_replace_q()).filter(_not_omni_create_q())
|
||||
|
||||
|
||||
def _free_video_trash_queryset(team):
|
||||
return (
|
||||
AITask.objects.filter(team=team, task_type=AITask.Type.FREE_VIDEO, is_deleted=True, purged_at__isnull=True)
|
||||
.exclude(request_payload__feature="omni_create")
|
||||
.filter(_not_omni_create_q())
|
||||
.select_related("model_config")
|
||||
.prefetch_related("generated_assets", "generated_assets__files")
|
||||
)
|
||||
@@ -1875,7 +1883,7 @@ class CreationConversationViewSet(TeamScopedViewSetMixin, ModelViewSet):
|
||||
|
||||
@action(detail=True, methods=["post"], url_path="merge-video-segments")
|
||||
def merge_video_segments(self, request, pk=None):
|
||||
"""用户确认后才合并多段成片;未点击前不下载视频、更不调用 ffmpeg。"""
|
||||
"""合并多段成片(兼容旧前端手动触发;正常流程在分段全成功后由平台自动合并)。"""
|
||||
conversation = self.get_object()
|
||||
message_id = str(request.data.get("message_id") or "").strip()
|
||||
message = conversation.messages.filter(id=message_id).first() if message_id else None
|
||||
|
||||
@@ -0,0 +1,113 @@
|
||||
"""角色库文件名 → 模特标签。
|
||||
|
||||
文件名形如「警察-男性-日本-年轻男性.png」「动画角色-齿轮蝠.png」。
|
||||
文件夹名(医生/警察/动画角色…)作为职业/分类标签;性别归一成「男人」「女人」。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from pathlib import Path
|
||||
|
||||
GENDER_TOKENS = {
|
||||
"男性": "男人",
|
||||
"男童": "男人",
|
||||
"年轻男性": "男人",
|
||||
"中年男性": "男人",
|
||||
"老年男性": "男人",
|
||||
"青少年男性": "男人",
|
||||
"女性": "女人",
|
||||
"女童": "女人",
|
||||
"少女": "女人",
|
||||
"年轻女性": "女人",
|
||||
"中年女性": "女人",
|
||||
"老年女性": "女人",
|
||||
"青少年女性": "女人",
|
||||
"成年女性": "女人",
|
||||
}
|
||||
|
||||
AGE_TOKENS = ("年轻", "中年", "老年", "青少年", "童", "成年")
|
||||
|
||||
# 文件名里常带的噪声后缀,不当标签
|
||||
NOISE = {"面孔", "中国面孔", "欧美面孔", "亚洲面孔", "非洲面孔", "美国面孔", "印度面孔"}
|
||||
|
||||
|
||||
def _split_tokens(stem: str) -> list[str]:
|
||||
parts = re.split(r"[-_·//]+", stem)
|
||||
return [p.strip() for p in parts if p and p.strip()]
|
||||
|
||||
|
||||
def parse_character_tags(*, folder: str, filename: str) -> list[str]:
|
||||
"""从文件夹 + 文件名解析去重后的标签列表(顺序:分类 → 性别 → 其余)。"""
|
||||
stem = Path(filename).stem
|
||||
tokens = _split_tokens(stem)
|
||||
folder_tag = folder.strip()
|
||||
tags: list[str] = []
|
||||
seen: set[str] = set()
|
||||
|
||||
def add(tag: str) -> None:
|
||||
tag = tag.strip()
|
||||
if not tag or tag in seen or tag in NOISE:
|
||||
return
|
||||
seen.add(tag)
|
||||
tags.append(tag)
|
||||
|
||||
if folder_tag:
|
||||
add(folder_tag)
|
||||
|
||||
# 性别优先归一
|
||||
for tok in tokens:
|
||||
if tok in GENDER_TOKENS:
|
||||
add(GENDER_TOKENS[tok])
|
||||
break
|
||||
for key, gender in GENDER_TOKENS.items():
|
||||
if key in tok:
|
||||
add(gender)
|
||||
break
|
||||
|
||||
for tok in tokens:
|
||||
if tok == folder_tag:
|
||||
continue
|
||||
if tok in GENDER_TOKENS:
|
||||
continue
|
||||
# 年龄词单独保留简写
|
||||
age_hit = next((a for a in AGE_TOKENS if a in tok and tok not in GENDER_TOKENS), None)
|
||||
if age_hit and tok in GENDER_TOKENS.values():
|
||||
continue
|
||||
if tok in NOISE:
|
||||
continue
|
||||
# 去掉已归一进性别的复合词本体
|
||||
if any(g in tok for g in ("男性", "女性", "男童", "女童", "少女")):
|
||||
for a in AGE_TOKENS:
|
||||
if a in tok:
|
||||
add(a)
|
||||
break
|
||||
continue
|
||||
add(tok)
|
||||
|
||||
return tags
|
||||
|
||||
|
||||
def display_name_from_file(*, folder: str, filename: str) -> str:
|
||||
stem = Path(filename).stem
|
||||
parts = _split_tokens(stem)
|
||||
if not parts:
|
||||
return folder or "角色"
|
||||
# 动画角色:「动画角色-齿轮蝠」→ 齿轮蝠
|
||||
if folder == "动画角色" and len(parts) >= 2:
|
||||
return parts[-1][:64]
|
||||
|
||||
body = parts[1:] if parts[0] in {folder, "民族服饰"} and len(parts) > 1 else parts
|
||||
# 去掉纯性别词,保留「中年女性」这类年龄+性别复合,以及地区
|
||||
cleaned: list[str] = []
|
||||
for tok in body:
|
||||
if tok in {"男性", "女性"}:
|
||||
continue
|
||||
cleaned.append(tok.replace("面孔", ""))
|
||||
label = "·".join(t for t in cleaned if t) or stem
|
||||
if folder and folder not in {"动画角色"}:
|
||||
return f"{folder}·{label}"[:64]
|
||||
return label[:64]
|
||||
|
||||
|
||||
def import_key(relative_path: str) -> str:
|
||||
return f"character_pack:{relative_path.replace(chr(92), '/')}"
|
||||
@@ -0,0 +1,157 @@
|
||||
"""把本地「角色库」文件夹批量导入为官方模特(带标签)。
|
||||
|
||||
用法:
|
||||
python manage.py import_character_pack "/Users/Admin/Downloads/角色库"
|
||||
python manage.py import_character_pack "/path" --dry-run
|
||||
python manage.py import_character_pack "/path" --team-owner admin
|
||||
|
||||
幂等:同一相对路径(metadata.import_key)不会重复导入。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import mimetypes
|
||||
import uuid
|
||||
from pathlib import Path
|
||||
|
||||
from django.core.management.base import BaseCommand, CommandError
|
||||
from django.db import transaction
|
||||
|
||||
from apps.accounts.models import Team, TeamMember, User
|
||||
from apps.assets.character_tags import display_name_from_file, import_key, parse_character_tags
|
||||
from apps.assets.models import Asset, AssetFile, Model
|
||||
from apps.assets.storage import TosStorage
|
||||
|
||||
IMAGE_SUFFIXES = {".png", ".jpg", ".jpeg", ".webp", ".gif"}
|
||||
|
||||
|
||||
class Command(BaseCommand):
|
||||
help = "Import a local character-pack folder into official Model library with tags"
|
||||
|
||||
def add_arguments(self, parser):
|
||||
parser.add_argument("folder", type=str, help="角色库根目录")
|
||||
parser.add_argument("--dry-run", action="store_true", help="只解析标签,不上传")
|
||||
parser.add_argument("--team-owner", default="admin", help="托管官方模特的团队 owner 用户名")
|
||||
parser.add_argument("--limit", type=int, default=0, help="最多导入 N 张(0=全部)")
|
||||
|
||||
def handle(self, *args, **options):
|
||||
root = Path(options["folder"]).expanduser().resolve()
|
||||
if not root.is_dir():
|
||||
raise CommandError(f"目录不存在: {root}")
|
||||
|
||||
owner = User.objects.filter(username=options["team_owner"]).first()
|
||||
if owner is None:
|
||||
raise CommandError(f"找不到用户 {options['team_owner']}")
|
||||
team = (
|
||||
Team.objects.filter(owner=owner, is_personal=True).order_by("created_at").first()
|
||||
or Team.objects.filter(owner=owner).order_by("created_at").first()
|
||||
)
|
||||
if team is None:
|
||||
team = Team.objects.create(name=f"{owner.username}-official", owner=owner, is_personal=True)
|
||||
TeamMember.objects.create(team=team, user=owner, role=TeamMember.Role.OWNER)
|
||||
self.stdout.write(f"已创建托管团队 {team.name}")
|
||||
|
||||
files = sorted(
|
||||
p for p in root.rglob("*")
|
||||
if p.is_file() and p.suffix.lower() in IMAGE_SUFFIXES and p.name != ".DS_Store"
|
||||
)
|
||||
if options["limit"]:
|
||||
files = files[: options["limit"]]
|
||||
if not files:
|
||||
raise CommandError("目录里没有图片")
|
||||
|
||||
storage = None if options["dry_run"] else TosStorage()
|
||||
created = skipped = failed = 0
|
||||
|
||||
for path in files:
|
||||
rel = str(path.relative_to(root)).replace("\\", "/")
|
||||
folder = path.parent.name if path.parent != root else ""
|
||||
tags = parse_character_tags(folder=folder, filename=path.name)
|
||||
name = display_name_from_file(folder=folder, filename=path.name)
|
||||
key = import_key(rel)
|
||||
|
||||
existing = Model.objects.filter(
|
||||
is_deleted=False,
|
||||
metadata__import_key=key,
|
||||
).first()
|
||||
if existing:
|
||||
skipped += 1
|
||||
self.stdout.write(f"skip {rel} tags={tags}")
|
||||
continue
|
||||
|
||||
self.stdout.write(f"{'dry' if options['dry_run'] else 'add '} {rel} → {name} tags={tags}")
|
||||
if options["dry_run"]:
|
||||
created += 1
|
||||
continue
|
||||
|
||||
try:
|
||||
self._create_one(
|
||||
storage=storage,
|
||||
team=team,
|
||||
owner=owner,
|
||||
path=path,
|
||||
name=name,
|
||||
tags=tags,
|
||||
key=key,
|
||||
folder=folder,
|
||||
rel=rel,
|
||||
)
|
||||
created += 1
|
||||
except Exception as exc: # noqa: BLE001
|
||||
failed += 1
|
||||
self.stderr.write(self.style.ERROR(f"fail {rel}: {exc}"))
|
||||
|
||||
self.stdout.write(self.style.SUCCESS(
|
||||
f"done created={created} skipped={skipped} failed={failed} team={team.id}"
|
||||
))
|
||||
|
||||
def _create_one(self, *, storage, team, owner, path: Path, name: str, tags: list[str], key: str, folder: str, rel: str):
|
||||
asset_id = uuid.uuid4()
|
||||
suffix = path.suffix.lower() or ".png"
|
||||
object_key = f"teams/{team.id}/models/official/{asset_id}{suffix}"
|
||||
content_type = mimetypes.guess_type(path.name)[0] or "image/png"
|
||||
with path.open("rb") as fh:
|
||||
stored = storage.upload_fileobj(
|
||||
fileobj=fh,
|
||||
object_key=object_key,
|
||||
content_type=content_type,
|
||||
)
|
||||
with transaction.atomic():
|
||||
portrait = Asset.objects.create(
|
||||
id=asset_id,
|
||||
team=team,
|
||||
created_by=owner,
|
||||
name=f"{name}·形象图",
|
||||
asset_type=Asset.Type.IMAGE,
|
||||
source=Asset.Source.UPLOAD,
|
||||
category=Asset.Category.MODEL_PORTRAIT,
|
||||
metadata={
|
||||
"kind": "model",
|
||||
"view": "frontal",
|
||||
"source": "character_pack",
|
||||
"tags": tags,
|
||||
},
|
||||
)
|
||||
AssetFile.objects.create(
|
||||
asset=portrait,
|
||||
object_key=stored.object_key,
|
||||
bucket=stored.bucket,
|
||||
content_type=stored.content_type,
|
||||
size_bytes=stored.size_bytes,
|
||||
is_primary=True,
|
||||
)
|
||||
Model.objects.create(
|
||||
team=team,
|
||||
created_by=owner,
|
||||
name=name,
|
||||
is_official=True,
|
||||
source=Model.Source.UPLOAD,
|
||||
portrait_asset=portrait,
|
||||
description=f"角色库 · {folder}" if folder else "角色库",
|
||||
metadata={
|
||||
"source": "character_pack",
|
||||
"tags": tags,
|
||||
"import_key": key,
|
||||
"folder": folder,
|
||||
"relative_path": rel,
|
||||
},
|
||||
)
|
||||
@@ -0,0 +1,29 @@
|
||||
# Generated by Django 5.1.15 on 2026-09-21 08:04
|
||||
|
||||
import uuid
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('assets', '0010_asset_model_purged_at'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.CreateModel(
|
||||
name='ModelLibraryTag',
|
||||
fields=[
|
||||
('id', models.UUIDField(default=uuid.uuid4, editable=False, primary_key=True, serialize=False)),
|
||||
('created_at', models.DateTimeField(auto_now_add=True)),
|
||||
('updated_at', models.DateTimeField(auto_now=True)),
|
||||
('name', models.CharField(max_length=64, unique=True)),
|
||||
('label', models.CharField(blank=True, max_length=64)),
|
||||
('visible', models.BooleanField(default=True)),
|
||||
('sort', models.IntegerField(default=0)),
|
||||
],
|
||||
options={
|
||||
'ordering': ['sort', 'name'],
|
||||
},
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,144 @@
|
||||
"""模特库标签目录:从 Model.metadata.tags 聚合 + 与 ModelLibraryTag 同步。"""
|
||||
from __future__ import annotations
|
||||
|
||||
from collections import Counter
|
||||
|
||||
from django.db.models import Q
|
||||
|
||||
from .models import Model, ModelLibraryTag
|
||||
|
||||
# 前台默认优先露出的粗粒度标签(即使 count 偏低也显示)
|
||||
PREFERRED = (
|
||||
"男人", "女人", "动画角色", "学生", "医生", "警察", "法官", "厨师",
|
||||
"民族", "民族服饰", "中国",
|
||||
)
|
||||
|
||||
# 首次入库:出现次数低于此值且不在 PREFERRED → 默认隐藏
|
||||
DEFAULT_VISIBLE_MIN_COUNT = 3
|
||||
|
||||
|
||||
def aggregate_tag_counts(*, official_only: bool = False, mine_team_id=None) -> dict[str, int]:
|
||||
qs = Model.objects.filter(is_deleted=False, purged_at__isnull=True)
|
||||
if official_only:
|
||||
qs = qs.filter(is_official=True)
|
||||
elif mine_team_id is not None:
|
||||
qs = qs.filter(is_official=False, team_id=mine_team_id)
|
||||
counts: Counter[str] = Counter()
|
||||
for md in qs.values_list("metadata", flat=True):
|
||||
tags = (md or {}).get("tags") if isinstance(md, dict) else None
|
||||
if not isinstance(tags, list):
|
||||
continue
|
||||
for tag in tags:
|
||||
if isinstance(tag, str) and tag.strip():
|
||||
counts[tag.strip()] += 1
|
||||
return dict(counts)
|
||||
|
||||
|
||||
def sync_tag_catalog(*, create_missing: bool = True) -> list[ModelLibraryTag]:
|
||||
"""把库里出现过的标签写入目录;已有行不改 visible/sort/label。"""
|
||||
counts = aggregate_tag_counts()
|
||||
existing = {t.name: t for t in ModelLibraryTag.objects.all()}
|
||||
created: list[ModelLibraryTag] = []
|
||||
if create_missing:
|
||||
for name, count in counts.items():
|
||||
if name in existing:
|
||||
continue
|
||||
visible = name in PREFERRED or count >= DEFAULT_VISIBLE_MIN_COUNT
|
||||
sort = PREFERRED.index(name) if name in PREFERRED else 100 + (0 if count >= 10 else 50)
|
||||
obj = ModelLibraryTag.objects.create(name=name, visible=visible, sort=sort)
|
||||
existing[name] = obj
|
||||
created.append(obj)
|
||||
return created
|
||||
|
||||
|
||||
def public_tag_rows(*, tab: str | None = None, team_id=None) -> list[dict]:
|
||||
"""前台筛选条:只返回 visible=True 的标签 + 当前范围计数。"""
|
||||
sync_tag_catalog(create_missing=True)
|
||||
if tab == "official":
|
||||
counts = aggregate_tag_counts(official_only=True)
|
||||
elif tab == "mine":
|
||||
counts = aggregate_tag_counts(mine_team_id=team_id)
|
||||
else:
|
||||
counts = aggregate_tag_counts()
|
||||
rows = []
|
||||
for tag in ModelLibraryTag.objects.filter(visible=True).order_by("sort", "name"):
|
||||
c = counts.get(tag.name, 0)
|
||||
if c <= 0:
|
||||
continue
|
||||
rows.append({"name": tag.name, "label": tag.display_name, "count": c})
|
||||
return rows
|
||||
|
||||
|
||||
def rename_tag_everywhere(old: str, new: str) -> int:
|
||||
old, new = old.strip(), new.strip()
|
||||
if not old or not new or old == new:
|
||||
return 0
|
||||
updated = 0
|
||||
qs = Model.objects.filter(metadata__tags__contains=[old])
|
||||
for m in qs.iterator():
|
||||
md = dict(m.metadata or {})
|
||||
tags = md.get("tags")
|
||||
if not isinstance(tags, list):
|
||||
continue
|
||||
nxt: list[str] = []
|
||||
seen: set[str] = set()
|
||||
for t in tags:
|
||||
if not isinstance(t, str):
|
||||
continue
|
||||
v = new if t.strip() == old else t.strip()
|
||||
if not v or v in seen:
|
||||
continue
|
||||
seen.add(v)
|
||||
nxt.append(v)
|
||||
if nxt != tags:
|
||||
md["tags"] = nxt
|
||||
m.metadata = md
|
||||
m.save(update_fields=["metadata", "updated_at"])
|
||||
updated += 1
|
||||
# catalog row
|
||||
src = ModelLibraryTag.objects.filter(name=old).first()
|
||||
dst = ModelLibraryTag.objects.filter(name=new).first()
|
||||
if src and dst:
|
||||
src.delete()
|
||||
elif src:
|
||||
src.name = new
|
||||
src.save(update_fields=["name", "updated_at"])
|
||||
return updated
|
||||
|
||||
|
||||
def replace_tag_on_models(old: str, new: str | None) -> int:
|
||||
"""把 old 换成 new;new=None 表示从模特上移除该标签。"""
|
||||
if new is None:
|
||||
new_s = None
|
||||
else:
|
||||
new_s = new.strip()
|
||||
if not new_s:
|
||||
return 0
|
||||
old = old.strip()
|
||||
updated = 0
|
||||
qs = Model.objects.filter(metadata__tags__contains=[old])
|
||||
for m in qs.iterator():
|
||||
md = dict(m.metadata or {})
|
||||
tags = md.get("tags")
|
||||
if not isinstance(tags, list):
|
||||
continue
|
||||
nxt: list[str] = []
|
||||
seen: set[str] = set()
|
||||
for t in tags:
|
||||
if not isinstance(t, str):
|
||||
continue
|
||||
v = t.strip()
|
||||
if v == old:
|
||||
if new_s is None:
|
||||
continue
|
||||
v = new_s
|
||||
if not v or v in seen:
|
||||
continue
|
||||
seen.add(v)
|
||||
nxt.append(v)
|
||||
if nxt != tags:
|
||||
md["tags"] = nxt
|
||||
m.metadata = md
|
||||
m.save(update_fields=["metadata", "updated_at"])
|
||||
updated += 1
|
||||
return updated
|
||||
@@ -118,6 +118,29 @@ class Model(TeamOwnedModel):
|
||||
return self.name
|
||||
|
||||
|
||||
|
||||
class ModelLibraryTag(TimeStampedModel):
|
||||
"""模特库筛选标签目录 · 控制前台芯片显隐 / 排序 / 展示名。
|
||||
|
||||
实际标签仍写在 Model.metadata.tags;本表只管「筛选条怎么展示」。
|
||||
未入表的标签在 sync 时按出现次数决定默认显隐。
|
||||
"""
|
||||
|
||||
name = models.CharField(max_length=64, unique=True)
|
||||
label = models.CharField(max_length=64, blank=True) # 空则用 name
|
||||
visible = models.BooleanField(default=True)
|
||||
sort = models.IntegerField(default=0)
|
||||
|
||||
class Meta:
|
||||
ordering = ["sort", "name"]
|
||||
|
||||
def __str__(self) -> str:
|
||||
return self.label or self.name
|
||||
|
||||
@property
|
||||
def display_name(self) -> str:
|
||||
return (self.label or self.name).strip() or self.name
|
||||
|
||||
class AssetReviewGroup(TimeStampedModel):
|
||||
"""一团队一火山人像素材组(真人资产审核统一上传到这里,单组可放 500 万)。"""
|
||||
|
||||
|
||||
@@ -888,22 +888,52 @@ class ReviewScopeTests(TestCase):
|
||||
("person", "model_portrait", "tri_view", "storyboard", "scene", "product_image"),
|
||||
)
|
||||
|
||||
def test_poll_team_reviews_only_scans_review_categories(self):
|
||||
def test_poll_only_covers_assets_that_really_reached_volcano(self):
|
||||
"""轮询的判据是「拿到了火山远端 Id」,不是类别。
|
||||
|
||||
force 送审的上传视频、临时图也会是 processing,worker 必须盯到绿/红;
|
||||
反过来,没有远端 Id 的 processing 是没真正送出去的,轮它只会白打请求。
|
||||
送审范围仍然只挑 REVIEW_CATEGORIES(见下一条)。
|
||||
"""
|
||||
from apps.assets import review
|
||||
|
||||
_, team = _mk_team("uc", "TeamC")
|
||||
keep = {}
|
||||
polled = {}
|
||||
for cat in ("person", "model_portrait", "tri_view", "storyboard", "scene", "product_image"):
|
||||
a = Asset.objects.create(team=team, name=cat, asset_type="image", source="ai_generated",
|
||||
category=cat, review_status="processing")
|
||||
keep[cat] = str(a.id)
|
||||
# 图片趴:不应进送审队列
|
||||
for cat in ("model_tryon", "platform_kit", "free_create"):
|
||||
Asset.objects.create(team=team, name=cat, asset_type="image", source="ai_generated",
|
||||
category=cat, review_status="processing")
|
||||
category=cat, review_status="processing",
|
||||
review_remote_id=f"remote-{cat}")
|
||||
polled[cat] = str(a.id)
|
||||
# 图片趴不在送审范围里,但已经 force 送出去的同样要盯到终态。
|
||||
forced = Asset.objects.create(
|
||||
team=team, name="force", asset_type="image", source="ai_generated",
|
||||
category="free_create", review_status="processing", review_remote_id="remote-force",
|
||||
)
|
||||
# 没有远端 Id = 没真送出去,不该被轮询。
|
||||
Asset.objects.create(team=team, name="没送出去", asset_type="image", source="ai_generated",
|
||||
category="person", review_status="processing")
|
||||
|
||||
with patch("apps.assets.review.assets_client.is_enabled", return_value=False):
|
||||
out = review.poll_team_reviews(team)
|
||||
self.assertEqual(set(out.keys()), set(keep.values()))
|
||||
|
||||
self.assertEqual(set(out.keys()), set(polled.values()) | {str(forced.id)})
|
||||
|
||||
def test_submission_still_only_covers_review_categories(self):
|
||||
from apps.assets import review
|
||||
|
||||
_, team = _mk_team("uc-sub", "TeamCSub")
|
||||
for cat in ("person", "product_image"):
|
||||
Asset.objects.create(team=team, name=cat, asset_type="image", source="ai_generated",
|
||||
category=cat, review_status="")
|
||||
for cat in ("model_tryon", "platform_kit", "free_create"):
|
||||
Asset.objects.create(team=team, name=cat, asset_type="image", source="ai_generated",
|
||||
category=cat, review_status="")
|
||||
|
||||
with patch("apps.assets.review.submit_asset_for_review", return_value=True) as sub:
|
||||
review.submit_unsubmitted_reviews(team=team)
|
||||
|
||||
submitted = {asset.category for (asset,), _ in (call for call in sub.call_args_list)}
|
||||
self.assertEqual(submitted, {"person", "product_image"})
|
||||
|
||||
def test_poll_team_reviews_auto_submits_unsubmitted(self):
|
||||
from apps.assets import review
|
||||
@@ -931,6 +961,7 @@ class ReviewScopeTests(TestCase):
|
||||
proc = Asset.objects.create(
|
||||
team=team, name="审中", asset_type="image", source="ai_generated",
|
||||
category=Asset.Category.PERSON, review_status="processing",
|
||||
review_remote_id="remote-proc", # 只有真的送到火山、拿到 Id 的才会被轮询
|
||||
)
|
||||
with patch("apps.assets.review.submit_asset_for_review", return_value=True) as sub, patch(
|
||||
"apps.assets.review.poll_asset_review", return_value="active"
|
||||
|
||||
@@ -43,6 +43,15 @@ _KNOWN_CATS = [
|
||||
]
|
||||
|
||||
|
||||
def _not_omni_create_q() -> Q:
|
||||
"""「不是全能创作产出」。不能只写 ~Q(metadata__feature="omni_create")。
|
||||
|
||||
metadata 里没有 feature 键时取值是 SQL NULL,NOT(NULL = 'x') 还是 NULL,整行会被筛掉 ——
|
||||
历史资产(没写过 feature)会从「视频自由创作」「自由创作」「素材」这些分桶里整批消失。
|
||||
"""
|
||||
return Q(metadata__feature__isnull=True) | ~Q(metadata__feature="omni_create")
|
||||
|
||||
|
||||
def _tab_q(tab: str) -> Q:
|
||||
if tab == "people":
|
||||
return Q(category="person")
|
||||
@@ -57,13 +66,13 @@ def _tab_q(tab: str) -> Q:
|
||||
if tab == "image_creations": # 图片自由创作
|
||||
return Q(category="free_create", asset_type="image")
|
||||
if tab == "video_creations": # 视频自由创作
|
||||
return Q(category="free_create", asset_type="video") & ~Q(metadata__feature="omni_create")
|
||||
return Q(category="free_create", asset_type="video") & _not_omni_create_q()
|
||||
if tab == "creations": # 自由创作(兼容旧入口:图片+视频)
|
||||
return Q(category="free_create") & ~Q(metadata__feature="omni_create")
|
||||
return Q(category="free_create") & _not_omni_create_q()
|
||||
if tab == "uploads":
|
||||
return Q(category="upload")
|
||||
if tab == "materials": # 素材(期3):视频素材 = 所有视频,排除「最终成片」(final_video 隐藏不列)
|
||||
return ~Q(category="final_video") & Q(asset_type="video") & ~Q(metadata__feature="omni_create")
|
||||
return ~Q(category="final_video") & Q(asset_type="video") & _not_omni_create_q()
|
||||
if tab == "others": # 其他(资产库成品化):我的上传 + 未归类非视频(兜底)
|
||||
return Q(category="upload") | (~Q(category__in=_KNOWN_CATS) & ~Q(asset_type="video"))
|
||||
if tab == "unclassified": # 未归类且非视频(也不含最终成片)
|
||||
@@ -482,6 +491,14 @@ class ModelLibraryViewSet(ModelViewSet):
|
||||
if self.request.query_params.get("q"):
|
||||
q = self.request.query_params["q"]
|
||||
qs = qs.filter(Q(name__icontains=q) | Q(description__icontains=q))
|
||||
# 标签筛选: ?tag=男人 或 ?tag=男人&tag=警察(多标签 AND)
|
||||
raw_tags = self.request.query_params.getlist("tag") or []
|
||||
if not raw_tags:
|
||||
joined = (self.request.query_params.get("tags") or "").strip()
|
||||
if joined:
|
||||
raw_tags = [t.strip() for t in joined.split(",") if t.strip()]
|
||||
for tag in raw_tags:
|
||||
qs = qs.filter(metadata__tags__contains=[tag])
|
||||
# 官方模板靠前,再按新→旧
|
||||
return qs.order_by("-is_official", "-created_at")
|
||||
|
||||
@@ -605,6 +622,18 @@ class ModelLibraryViewSet(ModelViewSet):
|
||||
model.save(update_fields=["purged_at", "updated_at"])
|
||||
return Response(status=status.HTTP_204_NO_CONTENT)
|
||||
|
||||
@action(detail=False, methods=["get"], url_path="tags")
|
||||
def tags(self, request):
|
||||
"""当前可见模特里出现过的标签(供前端筛选条)· 只返回后台标记为显示的。"""
|
||||
from .model_library_tags import public_tag_rows
|
||||
|
||||
tab = (request.query_params.get("tab") or "").strip() or None
|
||||
if tab not in (None, "official", "mine"):
|
||||
tab = None
|
||||
team = self.get_team()
|
||||
items = public_tag_rows(tab=tab, team_id=getattr(team, "id", None))
|
||||
return Response({"results": items})
|
||||
|
||||
@action(detail=False, methods=["post"], url_path="upload", parser_classes=[MultiPartParser, FormParser])
|
||||
def upload(self, request):
|
||||
"""真人上传:一张人像图 → 建 model_portrait 资产 + Model(source=upload)。三视图/声线后续补。"""
|
||||
@@ -640,13 +669,27 @@ class ModelLibraryViewSet(ModelViewSet):
|
||||
size_bytes=stored.size_bytes,
|
||||
is_primary=True,
|
||||
)
|
||||
raw_tags = request.data.get("tags")
|
||||
tags: list[str] = []
|
||||
if isinstance(raw_tags, str) and raw_tags.strip():
|
||||
import json
|
||||
try:
|
||||
parsed = json.loads(raw_tags)
|
||||
if isinstance(parsed, list):
|
||||
tags = [str(t).strip() for t in parsed if str(t).strip()]
|
||||
else:
|
||||
tags = [t.strip() for t in raw_tags.split(",") if t.strip()]
|
||||
except Exception:
|
||||
tags = [t.strip() for t in raw_tags.split(",") if t.strip()]
|
||||
elif isinstance(raw_tags, list):
|
||||
tags = [str(t).strip() for t in raw_tags if str(t).strip()]
|
||||
model = Model.objects.create(
|
||||
team=team,
|
||||
created_by=request.user,
|
||||
name=name,
|
||||
source=Model.Source.UPLOAD,
|
||||
portrait_asset=portrait,
|
||||
metadata={"source": "upload"},
|
||||
metadata={"source": "upload", "tags": tags},
|
||||
)
|
||||
return Response(ModelLibrarySerializer(model).data, status=status.HTTP_201_CREATED)
|
||||
|
||||
|
||||
@@ -37,6 +37,22 @@ class QuickCreateApiTests(TestCase):
|
||||
self.client = APIClient()
|
||||
self.client.force_authenticate(self.user)
|
||||
|
||||
def _fake_video_model(self, name="doubao-seedance-2-0-fast-260128"):
|
||||
"""提交时会先按 metadata["pricing"] 预估费用,缺价目一律 400 —— 假模型必须带价。"""
|
||||
return SimpleNamespace(
|
||||
id=uuid.uuid4(),
|
||||
name=name,
|
||||
display_name="Seedance 2.0 Fast",
|
||||
metadata={
|
||||
"capabilities": {
|
||||
"resolutions": ["480p", "720p"],
|
||||
"aspect_ratios": ["9:16"],
|
||||
"durations": [15, 30],
|
||||
},
|
||||
"pricing": {"default": {"no_ref_video": 46, "with_ref_video": 28}},
|
||||
},
|
||||
)
|
||||
|
||||
def _uploaded_asset(self, **kwargs):
|
||||
upload = kwargs["upload"]
|
||||
return Asset.objects.create(
|
||||
@@ -55,16 +71,9 @@ class QuickCreateApiTests(TestCase):
|
||||
@patch("apps.projects.views._store_uploaded_asset")
|
||||
def test_submit_creates_product_project_and_persistent_job(self, store_asset, require_worker_task, enqueue, get_model, get_quick_model):
|
||||
store_asset.side_effect = self._uploaded_asset
|
||||
video_model_id = uuid.uuid4()
|
||||
get_model.side_effect = [
|
||||
object(),
|
||||
SimpleNamespace(
|
||||
id=video_model_id,
|
||||
name="doubao-seedance-2-0-fast-260128",
|
||||
display_name="Seedance 2.0 Fast",
|
||||
metadata={"capabilities": {"resolutions": ["480p", "720p"], "aspect_ratios": ["9:16"], "durations": [15]}},
|
||||
),
|
||||
]
|
||||
video_model = self._fake_video_model()
|
||||
video_model_id = video_model.id
|
||||
get_model.side_effect = [object(), video_model]
|
||||
response = self.client.post(
|
||||
"/api/projects/quick-create/",
|
||||
{
|
||||
@@ -120,12 +129,7 @@ class QuickCreateApiTests(TestCase):
|
||||
"""刷新页面时前端拉状态有空档,用户容易以为「没任务」再提交一次 ——
|
||||
那会重复建商品、重复扣费。后端必须拦住,并把在跑的那条带回去。"""
|
||||
store_asset.side_effect = self._uploaded_asset
|
||||
video_model = SimpleNamespace(
|
||||
id=uuid.uuid4(),
|
||||
name="doubao-seedance-2-0-fast-260128",
|
||||
display_name="Seedance 2.0 Fast",
|
||||
metadata={"capabilities": {"resolutions": ["480p", "720p"], "aspect_ratios": ["9:16"], "durations": [15]}},
|
||||
)
|
||||
video_model = self._fake_video_model()
|
||||
get_model.side_effect = [object(), video_model, object(), video_model]
|
||||
|
||||
def _post():
|
||||
@@ -167,10 +171,7 @@ class QuickCreateApiTests(TestCase):
|
||||
@patch("apps.projects.views._store_uploaded_asset")
|
||||
def test_submit_requires_product_category(self, store_asset, require_worker_task, enqueue, get_model, get_quick_model):
|
||||
store_asset.side_effect = self._uploaded_asset
|
||||
get_model.side_effect = [
|
||||
object(),
|
||||
SimpleNamespace(id=uuid.uuid4(), name="seedance", display_name="Seedance", metadata={}),
|
||||
]
|
||||
get_model.side_effect = [object(), self._fake_video_model("seedance")]
|
||||
response = self.client.post(
|
||||
"/api/projects/quick-create/",
|
||||
{
|
||||
@@ -200,7 +201,7 @@ class QuickCreateApiTests(TestCase):
|
||||
)
|
||||
source = Product.objects.create(team=self.team, created_by=self.user, title="旧商品", cover_asset=source_asset)
|
||||
ProductImage.objects.create(product=source, asset=source_asset, sort_order=0, is_primary=True)
|
||||
get_model.side_effect = [object(), SimpleNamespace(id=uuid.uuid4(), name="seedance", display_name="Seedance", metadata={})]
|
||||
get_model.side_effect = [object(), self._fake_video_model("seedance")]
|
||||
|
||||
response = self.client.post(
|
||||
"/api/projects/quick-create/",
|
||||
@@ -230,7 +231,7 @@ class QuickCreateApiTests(TestCase):
|
||||
other = User.objects.create_user(username="quick-image-other", password="pass")
|
||||
other_team = Team.objects.create(name="Image Other Team", owner=other)
|
||||
foreign = Product.objects.create(team=other_team, created_by=other, title="别人的图")
|
||||
get_model.side_effect = [object(), SimpleNamespace(metadata={})]
|
||||
get_model.side_effect = [object(), self._fake_video_model("seedance")]
|
||||
response = self.client.post(
|
||||
"/api/projects/quick-create/",
|
||||
{
|
||||
@@ -1375,6 +1376,10 @@ class QuickCreateCoordinatorTests(TestCase):
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user