添加全能创作功能
This commit is contained in:
@@ -0,0 +1,203 @@
|
||||
"""全能创作 · 会话与消息的写入服务(契约 §1/§2)。
|
||||
|
||||
只放「怎么把一条消息安全落库」这类底座能力;agent 循环、工具执行、SSE 在
|
||||
后续的 creation_agent.py 里,别混进来。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from django.db import transaction
|
||||
from django.db.models import Max
|
||||
from django.utils import timezone
|
||||
|
||||
from .models import CreationConversation, CreationMessage
|
||||
|
||||
|
||||
@transaction.atomic
|
||||
def append_message(
|
||||
conversation: CreationConversation,
|
||||
*,
|
||||
role: str,
|
||||
kind: str = CreationMessage.Kind.TEXT,
|
||||
text: str = "",
|
||||
payload: dict | None = None,
|
||||
refs: list | None = None,
|
||||
task=None,
|
||||
) -> CreationMessage:
|
||||
"""往会话尾部追加一条消息,并刷新 last_active_at。
|
||||
|
||||
seq 在事务里 select_for_update 锁住会话行再取 max+1 —— SSE 流式期间可能有
|
||||
并发写(用户抢答 / 轮询回填),不锁会撞 uniq_creation_message_seq。
|
||||
"""
|
||||
locked = CreationConversation.objects.select_for_update().get(pk=conversation.pk)
|
||||
next_seq = (locked.messages.aggregate(m=Max("seq"))["m"] or 0) + 1
|
||||
message = CreationMessage.objects.create(
|
||||
conversation=locked,
|
||||
role=role,
|
||||
kind=kind,
|
||||
text=text,
|
||||
payload=payload or {},
|
||||
refs=refs or [],
|
||||
task=task,
|
||||
seq=next_seq,
|
||||
)
|
||||
locked.last_active_at = timezone.now()
|
||||
locked.save(update_fields=["last_active_at", "updated_at"])
|
||||
conversation.last_active_at = locked.last_active_at
|
||||
return message
|
||||
|
||||
|
||||
def _assets_from_task(task) -> list[dict]:
|
||||
"""任务落库的资产 → 结果卡要用的 {id,url,cover,type}。
|
||||
|
||||
URL 必须走长期直链:预签名链 1 小时过期,写进消息 payload 第二天就打不开。
|
||||
"""
|
||||
from apps.ai.services import asset_stable_url
|
||||
from apps.assets.models import Asset
|
||||
|
||||
items = []
|
||||
assets = Asset.objects.filter(
|
||||
origin_task=task, is_deleted=False, purged_at__isnull=True,
|
||||
).prefetch_related("files")
|
||||
for asset in assets:
|
||||
url, cover = asset_stable_url(asset)
|
||||
if not url:
|
||||
continue
|
||||
items.append({
|
||||
"id": str(asset.id),
|
||||
"url": url,
|
||||
"cover": cover or url,
|
||||
"type": "video" if asset.asset_type == Asset.Type.VIDEO else "image",
|
||||
})
|
||||
return items
|
||||
|
||||
|
||||
def _meta_from_task(task, message: CreationMessage) -> dict:
|
||||
payload = message.payload or {}
|
||||
req = task.request_payload or {}
|
||||
model = ""
|
||||
if task.model_config_id:
|
||||
model = task.model_config.display_name or task.model_config.name
|
||||
return {
|
||||
"model": model or payload.get("model") or req.get("model") or "",
|
||||
"ratio": req.get("ratio") or payload.get("ratio") or "",
|
||||
"prompt": payload.get("prompt") or req.get("prompt") or "",
|
||||
}
|
||||
|
||||
|
||||
def _message_task(message: CreationMessage):
|
||||
"""GENERATING 消息挂的任务:优先 FK,payload.task_id 兜底(旧数据 / 序列化往返)。"""
|
||||
if message.task_id:
|
||||
return message.task
|
||||
task_id = (message.payload or {}).get("task_id")
|
||||
if not task_id:
|
||||
return None
|
||||
from .models import AITask
|
||||
|
||||
return AITask.objects.select_related("model_config").filter(id=task_id).first()
|
||||
|
||||
|
||||
@transaction.atomic
|
||||
def fail_generating_message(message: CreationMessage, error: str) -> CreationMessage:
|
||||
"""生成失败:GENERATING 原地改成 ERROR,不另开一条,避免中间态刷屏。"""
|
||||
message.kind = CreationMessage.Kind.ERROR
|
||||
message.text = (error or "生成失败")[:500]
|
||||
message.save(update_fields=["kind", "text", "updated_at"])
|
||||
conversation = message.conversation
|
||||
conversation.status = CreationConversation.Status.FAILED
|
||||
conversation.last_active_at = timezone.now()
|
||||
conversation.save(update_fields=["status", "last_active_at", "updated_at"])
|
||||
return message
|
||||
|
||||
|
||||
def sync_generating_message(message: CreationMessage) -> bool:
|
||||
"""看挂着的 AITask 是否已经终态,是就把 GENERATING 改成 RESULT / ERROR。
|
||||
|
||||
出图/出片在 worker 里跑,agent 只提交。前端轮询 GET 会话时靠这个回填;
|
||||
worker 结束时也会调一次,不用干等到下一次轮询。
|
||||
返回是否改了这条消息。
|
||||
"""
|
||||
from .models import AITask
|
||||
|
||||
if message.kind != CreationMessage.Kind.GENERATING:
|
||||
return False
|
||||
task = _message_task(message)
|
||||
if task is None:
|
||||
return False
|
||||
if task.status == AITask.Status.SUCCEEDED:
|
||||
assets = _assets_from_task(task)
|
||||
if not assets:
|
||||
return False # 状态已成功但资产还没落(极端竞态),下轮再试
|
||||
finish_generating_message(message, assets=assets, meta=_meta_from_task(task, message))
|
||||
return True
|
||||
if task.status in (AITask.Status.FAILED, AITask.Status.CANCELLED):
|
||||
fail_generating_message(message, task.error_message or "生成失败")
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def sync_generating_messages(conversation: CreationConversation) -> int:
|
||||
"""把一条会话里所有已结束的 GENERATING 回填。返回改了几条。"""
|
||||
pending = list(
|
||||
conversation.messages.filter(kind=CreationMessage.Kind.GENERATING)
|
||||
.select_related("task", "task__model_config")
|
||||
)
|
||||
return sum(1 for message in pending if sync_generating_message(message))
|
||||
|
||||
|
||||
def sync_generating_for_task(task) -> int:
|
||||
"""worker / poll 终态后:只扫挂在这个任务上的 GENERATING。失败不能向外抛。"""
|
||||
try:
|
||||
pending = list(
|
||||
CreationMessage.objects.filter(
|
||||
task=task, kind=CreationMessage.Kind.GENERATING,
|
||||
).select_related("conversation", "task", "task__model_config")
|
||||
)
|
||||
return sum(1 for message in pending if sync_generating_message(message))
|
||||
except Exception: # noqa: BLE001 — 回填失败不能把已经成功的出图任务打成失败
|
||||
import logging
|
||||
|
||||
logging.getLogger(__name__).warning(
|
||||
"omni create: sync generating for task %s failed", getattr(task, "id", "?"),
|
||||
exc_info=True,
|
||||
)
|
||||
return 0
|
||||
|
||||
|
||||
@transaction.atomic
|
||||
def finish_generating_message(message: CreationMessage, *, assets: list[dict], meta: dict) -> CreationMessage:
|
||||
"""把「生成中」原地改成「结果」(契约 §2:不新增消息,避免中间态刷屏)。
|
||||
|
||||
重生成是**新开一条** GENERATING → RESULT,所以对话流仍然是往下叠加;
|
||||
这里改的只是同一次生成自己的中间态。
|
||||
"""
|
||||
message.kind = CreationMessage.Kind.RESULT
|
||||
message.payload = {**(message.payload or {}), **meta, "assets": assets}
|
||||
message.save(update_fields=["kind", "payload", "updated_at"])
|
||||
conversation = message.conversation
|
||||
conversation.status = CreationConversation.Status.COMPLETED
|
||||
conversation.last_active_at = timezone.now()
|
||||
conversation.save(update_fields=["status", "last_active_at", "updated_at"])
|
||||
return message
|
||||
|
||||
|
||||
def pin_refs(conversation: CreationConversation, refs: list[dict]) -> list[dict]:
|
||||
"""把本轮引用的实体并进「实体锁定」。按 (type, id) 去重,保留首次出现的顺序 ——
|
||||
每轮无条件带上它们,这是多次生成之间锁脸/锁商品的唯一手段(契约 §5)。
|
||||
"""
|
||||
merged = list(conversation.pinned_refs or [])
|
||||
seen = {(item.get("type"), str(item.get("id"))) for item in merged}
|
||||
changed = False
|
||||
for ref in refs or []:
|
||||
ref_type, ref_id = ref.get("type"), ref.get("id")
|
||||
if not ref_type or not ref_id:
|
||||
continue # 缺 type/id 的 ref 解析不出实体,直接丢
|
||||
key = (ref_type, str(ref_id))
|
||||
if key in seen:
|
||||
continue
|
||||
merged.append(ref)
|
||||
seen.add(key)
|
||||
changed = True
|
||||
if changed:
|
||||
conversation.pinned_refs = merged
|
||||
conversation.save(update_fields=["pinned_refs", "updated_at"])
|
||||
return merged
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,64 @@
|
||||
"""全能创作 · 预设(契约 §6)。
|
||||
|
||||
首页那 8+6 张预设卡不只是个名字 —— 每个预设代表一套**明确的拍法**。
|
||||
只把「达人口播种草」四个字塞进提示词,模型只能靠猜;这里给它可执行的约束。
|
||||
|
||||
新增预设 = 在下面加一条,前端 `omni-create.tsx` 的 PRESET 列表加一张卡。两边的
|
||||
key 必须是同一个中文名(会话建的时候原样存进 CreationConversation.preset)。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
VIDEO_PRESETS: dict[str, str] = {
|
||||
"剧情反转带货": (
|
||||
"轻剧情短片。必须有一个具体的困境场景 → 意外转折 → 商品成为解决问题的关键道具。"
|
||||
"商品不能在开头硬推,要等冲突立住了再自然介入。禁止把普通不便夸成严重后果。"
|
||||
),
|
||||
"商品拟人广告": (
|
||||
"把商品拟人化成有性格的角色,用它的动作、表情和情绪推进。轻快有趣,不说教。"
|
||||
"拟人化不能牺牲商品真实外观 —— 材质、颜色、结构必须和参考图一致。"
|
||||
),
|
||||
"达人口播种草": (
|
||||
"真人出镜口播,生活化语气,像朋友分享而不是念广告稿。开头一句话建立停留理由,"
|
||||
"中段讲清使用场景和一个核心卖点,结尾给明确行动理由。禁止万能主播腔和空卖点。"
|
||||
),
|
||||
"商品图一键成片": (
|
||||
"从商品参考图出发建立镜头语言:主图定调 → 补场景 → 补使用动作 → 收尾。"
|
||||
"节奏清晰,转场干净。商品外观严格以参考图为准。"
|
||||
),
|
||||
"鱼眼换装": (
|
||||
"鱼眼/广角近距离透视,连续换装节奏。**人物面部和身形必须全程一致**,"
|
||||
"只有服装在变。每次换装用一个明确动作触发。"
|
||||
),
|
||||
"点触换款": (
|
||||
"统一构图和机位,用点击/触碰动作触发商品款式切换,快速展示多个 SKU。"
|
||||
"背景、光线、机位全程不变 —— 变化只发生在商品本身。"
|
||||
),
|
||||
"探店漫游": (
|
||||
"以空间动线串联:入口 → 环境 → 关键细节 → 服务/主推项目。"
|
||||
"镜头连续移动有路线感,不要碎切。"
|
||||
),
|
||||
"品牌质感大片": (
|
||||
"强调光影、材质和镜头节奏,建立品牌识别。慢节奏、精致构图、克制的色彩。"
|
||||
"少即是多,不要堆信息。"
|
||||
),
|
||||
}
|
||||
|
||||
IMAGE_PRESETS: dict[str, str] = {
|
||||
"商品场景套图": (
|
||||
"一组风格统一的电商图:主图(干净突出商品)、场景图(真实使用环境)、细节图(材质/工艺特写)。"
|
||||
"三张的光线、色调、质感必须是同一套。"
|
||||
),
|
||||
"极简棚拍": "干净背景、柔和投影、主体明确居中。大量留白,不加多余道具。适合主图和详情页头图。",
|
||||
"清透自然光人像": "自然光,保留真实肤质和毛孔,不磨皮不过曝。氛围清透,人物状态放松自然。",
|
||||
"生活方式场景": "把商品放进真实生活空间和使用动作里,强调自然可信的生活气息,不要摆拍感。",
|
||||
"高级奢华质感": "深色环境 + 局部高光 + 材质细节特写,强化品牌高级感。对比强但不失细节。",
|
||||
"复古胶片风格": "低饱和、细腻颗粒、柔和对比、偏暖或偏青的胶片色调。怀旧情绪但不脏。",
|
||||
}
|
||||
|
||||
ALL_PRESETS: dict[str, str] = {**VIDEO_PRESETS, **IMAGE_PRESETS}
|
||||
|
||||
|
||||
def preset_guidance(name: str) -> str:
|
||||
"""预设名 → 拍法约束。认不出的名字返回 "" —— 前端加了新卡但这里还没写时,
|
||||
退回「只有名字」的行为,不要报错。"""
|
||||
return ALL_PRESETS.get((name or "").strip(), "")
|
||||
@@ -959,6 +959,9 @@ def finalize_free_video(*, task: AITask) -> AITask:
|
||||
)
|
||||
release_credit(reservation=locked.credit_reservation, reason=raw_message[:200])
|
||||
_notify_failure(locked, raw=f"[{code}] {raw_message}", hint=public_error.fallback_message)
|
||||
from apps.ai.creation import sync_generating_for_task
|
||||
|
||||
sync_generating_for_task(locked)
|
||||
return locked
|
||||
|
||||
# succeeded —— 认领 POSTPROCESSING(并发 finalize 只有一路进入慢活)
|
||||
@@ -1025,6 +1028,9 @@ def finalize_free_video(*, task: AITask) -> AITask:
|
||||
update_fields=["status", "actual_cost", "base_cost", "request_payload", "response_payload", "completed_at", "updated_at"]
|
||||
)
|
||||
charge_reserved_credit(reservation=reservation, actual_amount=actual)
|
||||
from apps.ai.creation import sync_generating_for_task
|
||||
|
||||
sync_generating_for_task(locked)
|
||||
return locked
|
||||
except Exception as exc: # noqa: BLE001 — 后处理失败:标失败退费(release 幂等,已扣则不动)
|
||||
logger.exception("free video finalize failed for task %s", locked.id)
|
||||
@@ -1045,6 +1051,9 @@ def finalize_free_video(*, task: AITask) -> AITask:
|
||||
locked.save(update_fields=["status", "error_code", "error_message", "completed_at", "updated_at"])
|
||||
release_credit(reservation=locked.credit_reservation, reason=str(exc)[:200])
|
||||
_notify_failure(locked, raw=str(exc), hint=public_error.fallback_message)
|
||||
from apps.ai.creation import sync_generating_for_task
|
||||
|
||||
sync_generating_for_task(locked)
|
||||
return locked
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,366 @@
|
||||
"""全能创作 · @引用:实体检索 与 Ref 解析(契约 §1/§3)。
|
||||
|
||||
两件事:
|
||||
1. `search_mentions()` —— 输入框打 @ 时的检索,返回 [Ref] 给前端渲染菜单。
|
||||
2. `resolve_refs()` —— 把消息里的 [Ref] 变成模型真正吃得下的两样东西:
|
||||
**事实文本**(卖点/规格,进提示词)+ **参考图**(进 content_items,锁脸/锁商品/锁场景)。
|
||||
|
||||
铁律:消息里存的是结构化 Ref(type + id),**不是** "@净颜精华" 这串字。
|
||||
后端必须拿 id 回表取事实与图,靠字符串匹配迟早对不上。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import uuid
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from apps.assets.models import Asset, Model
|
||||
from apps.products.models import Product
|
||||
|
||||
from .services import _asset_preview_url, _product_cover_url
|
||||
|
||||
# Ref.type → 前端菜单里的分组名(设计稿 .omni-mention-group 的 small 文案)
|
||||
TYPE_LABELS = {
|
||||
"product": "商品库",
|
||||
"model": "模特库",
|
||||
"character": "角色",
|
||||
"scene": "场景库",
|
||||
"asset": "资产库",
|
||||
}
|
||||
VALID_TYPES = tuple(TYPE_LABELS)
|
||||
DEFAULT_TYPES = VALID_TYPES
|
||||
|
||||
# 参考图顺序固定:角色 → 场景 → 商品。这个顺序是出片模型 @图N 的语义依据,别改。
|
||||
_REF_ORDER = {"model": 0, "character": 0, "scene": 1, "product": 2, "asset": 3}
|
||||
# 火山单次出片最多 9 张图;留足余量,超出的靠优先级截断而不是报错
|
||||
MAX_REFERENCE_IMAGES = 6
|
||||
|
||||
# 「资产库」是兜底分组,不重复列已经有专属分组的资产 ——
|
||||
# 否则同一张定妆照会在「角色」和「资产库」各出现一次,菜单里看着像两个素材。
|
||||
ASSET_EXCLUDED_CATEGORIES = (
|
||||
Asset.Category.PERSON, # → character
|
||||
Asset.Category.SCENE, # → scene
|
||||
Asset.Category.MODEL_PORTRAIT, # → model
|
||||
Asset.Category.TRI_VIEW, # → model
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ResolvedRefs:
|
||||
"""resolve_refs 的产物。facts 进提示词,references 进 content_items。"""
|
||||
|
||||
facts: list[str] = field(default_factory=list)
|
||||
references: list[dict] = field(default_factory=list)
|
||||
missing: list[dict] = field(default_factory=list) # 删掉/不属于本团队的引用,要在对话里告诉用户
|
||||
|
||||
@property
|
||||
def facts_text(self) -> str:
|
||||
return "\n\n".join(self.facts)
|
||||
|
||||
|
||||
def _ref(type_: str, obj_id, name: str, cover: str = "") -> dict:
|
||||
return {"type": type_, "id": str(obj_id), "name": name, "cover": cover}
|
||||
|
||||
|
||||
def _search_products(team, q: str, limit: int) -> list[dict]:
|
||||
queryset = Product.objects.filter(team=team, purged_at__isnull=True)
|
||||
if q:
|
||||
queryset = queryset.filter(title__icontains=q)
|
||||
out = []
|
||||
for product in queryset.order_by("-created_at")[:limit]:
|
||||
out.append(_ref("product", product.id, product.title, _product_cover_url(product)))
|
||||
return out
|
||||
|
||||
|
||||
def _search_models(team, q: str, limit: int) -> list[dict]:
|
||||
queryset = Model.objects.filter(team=team, is_deleted=False, purged_at__isnull=True)
|
||||
if q:
|
||||
queryset = queryset.filter(name__icontains=q)
|
||||
out = []
|
||||
for model in queryset.select_related("portrait_asset")[:limit]:
|
||||
out.append(_ref("model", model.id, model.name, _asset_preview_url(model.portrait_asset)))
|
||||
return out
|
||||
|
||||
|
||||
def _search_assets(team, q: str, limit: int, categories: tuple[str, ...], type_: str) -> list[dict]:
|
||||
queryset = Asset.objects.filter(
|
||||
team=team,
|
||||
is_deleted=False,
|
||||
purged_at__isnull=True,
|
||||
asset_type=Asset.Type.IMAGE,
|
||||
category__in=categories,
|
||||
)
|
||||
if type_ == "asset":
|
||||
# 「资产库」只列用户真正加进库的图,不然工作台的每张试验图都会冒出来
|
||||
queryset = queryset.filter(in_library=True)
|
||||
if q:
|
||||
queryset = queryset.filter(name__icontains=q)
|
||||
out = []
|
||||
for asset in queryset.order_by("-created_at")[:limit]:
|
||||
out.append(_ref(type_, asset.id, asset.name, _asset_preview_url(asset)))
|
||||
return out
|
||||
|
||||
|
||||
def search_mentions(team, q: str = "", types: list[str] | None = None, limit: int = 8) -> list[dict]:
|
||||
"""@ 检索。types 不传则全类型各取 limit 条,按 商品 → 模特 → 角色 → 场景 → 资产 排。"""
|
||||
wanted = [t for t in (types or DEFAULT_TYPES) if t in TYPE_LABELS]
|
||||
q = (q or "").strip()
|
||||
results: list[dict] = []
|
||||
for type_ in wanted:
|
||||
if type_ == "product":
|
||||
results.extend(_search_products(team, q, limit))
|
||||
elif type_ == "model":
|
||||
results.extend(_search_models(team, q, limit))
|
||||
elif type_ == "character":
|
||||
results.extend(_search_assets(team, q, limit, (Asset.Category.PERSON,), "character"))
|
||||
elif type_ == "scene":
|
||||
results.extend(_search_assets(team, q, limit, (Asset.Category.SCENE,), "scene"))
|
||||
elif type_ == "asset":
|
||||
categories = tuple(
|
||||
c for c in Asset.Category.values if c not in ASSET_EXCLUDED_CATEGORIES
|
||||
)
|
||||
results.extend(_search_assets(team, q, limit, categories, "asset"))
|
||||
return results
|
||||
|
||||
|
||||
_KEY_TO_TYPES = {
|
||||
"product": ["product"],
|
||||
"sku": ["product"],
|
||||
"goods": ["product"],
|
||||
"item": ["product"],
|
||||
"model": ["model"],
|
||||
"character": ["character"],
|
||||
"person": ["character"],
|
||||
"scene": ["scene"],
|
||||
"asset": ["asset"],
|
||||
}
|
||||
|
||||
|
||||
def infer_field_types(field: dict) -> list[str]:
|
||||
"""追问卡字段 → 该去哪类库里解析用户的选择。"""
|
||||
typed = [t for t in (field.get("asset_types") or []) if t in TYPE_LABELS]
|
||||
if typed:
|
||||
return typed
|
||||
key = str(field.get("key") or "").strip().lower()
|
||||
if key in _KEY_TO_TYPES:
|
||||
return _KEY_TO_TYPES[key]
|
||||
label = str(field.get("label") or "")
|
||||
if "商品" in label:
|
||||
return ["product"]
|
||||
if "模特" in label:
|
||||
return ["model"]
|
||||
if "角色" in label or "人物" in label:
|
||||
return ["character"]
|
||||
if "场景" in label:
|
||||
return ["scene"]
|
||||
return list(DEFAULT_TYPES)
|
||||
|
||||
|
||||
def lookup_mention(team, value: str, types: list[str] | None = None) -> dict | None:
|
||||
"""把追问卡里的选项值(实体 id 或精确名字)还原成 Ref。对不上就返回 None,绝不瞎配。"""
|
||||
raw = str(value or "").strip()
|
||||
if not raw:
|
||||
return None
|
||||
wanted = [t for t in (types or DEFAULT_TYPES) if t in TYPE_LABELS] or list(DEFAULT_TYPES)
|
||||
uid = None
|
||||
try:
|
||||
uid = str(uuid.UUID(raw))
|
||||
except ValueError:
|
||||
uid = None
|
||||
if uid:
|
||||
if "product" in wanted:
|
||||
product = Product.objects.filter(team=team, id=uid, purged_at__isnull=True).first()
|
||||
if product:
|
||||
return _ref("product", product.id, product.title, _product_cover_url(product))
|
||||
if "model" in wanted:
|
||||
model = (
|
||||
Model.objects.filter(team=team, id=uid, is_deleted=False, purged_at__isnull=True)
|
||||
.select_related("portrait_asset")
|
||||
.first()
|
||||
)
|
||||
if model:
|
||||
return _ref("model", model.id, model.name, _asset_preview_url(model.portrait_asset))
|
||||
if any(t in wanted for t in ("character", "scene", "asset")):
|
||||
asset = Asset.objects.filter(
|
||||
team=team, id=uid, is_deleted=False, purged_at__isnull=True
|
||||
).first()
|
||||
if asset is not None:
|
||||
if asset.category == Asset.Category.PERSON:
|
||||
type_ = "character"
|
||||
elif asset.category == Asset.Category.SCENE:
|
||||
type_ = "scene"
|
||||
else:
|
||||
type_ = "asset"
|
||||
if type_ in wanted:
|
||||
return _ref(type_, asset.id, asset.name, _asset_preview_url(asset))
|
||||
if types:
|
||||
return lookup_mention(team, raw, None)
|
||||
return None
|
||||
hits = search_mentions(team, q=raw, types=wanted, limit=8)
|
||||
for hit in hits:
|
||||
if (hit.get("name") or "") == raw:
|
||||
return hit
|
||||
return None
|
||||
|
||||
|
||||
def refs_from_elicit_answers(team, fields, answers: dict) -> list[dict]:
|
||||
"""用户在追问卡里点选的商品/角色等 → 可 pin 的 Ref。
|
||||
模型常用 type=single + 选项 value=商品id/名字,前端只回 answers 不回 refs,
|
||||
不在这里补上的话出片参考图里就没有这件商品。"""
|
||||
refs: list[dict] = []
|
||||
seen: set[tuple] = set()
|
||||
for field in fields or []:
|
||||
if not isinstance(field, dict):
|
||||
continue
|
||||
key = str(field.get("key") or "")
|
||||
if key in {"duration", "ratio", "resolution", "video_model", "count"}:
|
||||
continue
|
||||
raw = (answers or {}).get(field.get("key"))
|
||||
if raw is None:
|
||||
continue
|
||||
values = raw if isinstance(raw, list) else [raw]
|
||||
inferred = infer_field_types(field)
|
||||
for value in values:
|
||||
ref = lookup_mention(team, str(value or ""), inferred)
|
||||
if ref is None:
|
||||
continue
|
||||
mark = (ref.get("type"), str(ref.get("id")))
|
||||
if mark in seen:
|
||||
continue
|
||||
seen.add(mark)
|
||||
refs.append(ref)
|
||||
return refs
|
||||
|
||||
|
||||
def product_facts_text(product) -> str:
|
||||
"""商品事实块。全能创作没有 project,所以不能复用 script_agent._product_context()。
|
||||
这里只给**客观事实**(标题/品牌/品类/规格/卖点),不带人设和口吻 —— 那些由策略卡决定。"""
|
||||
lines = [f"商品:{product.title}"]
|
||||
if product.brand:
|
||||
lines.append(f"品牌:{product.brand}")
|
||||
if product.category:
|
||||
lines.append(f"品类:{product.category}")
|
||||
if product.target_audience:
|
||||
lines.append(f"目标人群:{product.target_audience}")
|
||||
description = (product.description or "").strip()
|
||||
if description:
|
||||
lines.append(f"商品描述:{description}")
|
||||
specs = product.specs if isinstance(product.specs, dict) else {}
|
||||
spec_text = "、".join(f"{k}:{v}" for k, v in specs.items() if v)
|
||||
if spec_text:
|
||||
lines.append(f"规格:{spec_text}")
|
||||
points = list(product.selling_points.order_by("sort_order", "created_at"))
|
||||
if points:
|
||||
joined = "\n".join(f"- {p.title}:{p.detail or p.title}" for p in points)
|
||||
lines.append(f"卖点:\n{joined}")
|
||||
return "\n".join(lines)
|
||||
|
||||
|
||||
def _asset_reference(asset, type_: str, label: str) -> dict | None:
|
||||
"""Asset → 参考图条目。带上审核态,视频路据此换成火山 asset:// 引用(否则真人图会被判「疑似真人」拒)。"""
|
||||
url = _asset_preview_url(asset)
|
||||
if not url:
|
||||
return None
|
||||
return {
|
||||
"url": url,
|
||||
"type": type_,
|
||||
"label": label,
|
||||
"asset_id": str(asset.id),
|
||||
"review_status": asset.review_status,
|
||||
"review_remote_id": asset.review_remote_id,
|
||||
}
|
||||
|
||||
|
||||
def _product_reference(product) -> dict | None:
|
||||
"""商品参考图:**真实上传图优先,排除 AI 生成图** —— 拿生成图当真相再喂回模型会误差累积。
|
||||
一张真实图都没有才回落封面(可能是 AI 图,但好过纯文生图)。
|
||||
|
||||
这里不复用 services._product_reference_urls():那个只返回 url,而视频路还需要
|
||||
asset_id 和审核态才能把图换成火山 asset:// 引用(商品图也可能出现真人上身)。
|
||||
"""
|
||||
rels = sorted(product.images.select_related("asset").all(), key=lambda im: (not im.is_primary, im.sort_order))
|
||||
for rel in rels:
|
||||
asset = rel.asset
|
||||
if asset is None or asset.source == Asset.Source.AI_GENERATED:
|
||||
continue
|
||||
entry = _asset_reference(asset, "product", product.title)
|
||||
if entry:
|
||||
return entry
|
||||
if product.cover_asset_id:
|
||||
entry = _asset_reference(product.cover_asset, "product", product.title)
|
||||
if entry:
|
||||
return entry
|
||||
cover = _product_cover_url(product)
|
||||
return (
|
||||
{"url": cover, "type": "product", "label": product.title,
|
||||
"asset_id": "", "review_status": "", "review_remote_id": ""}
|
||||
if cover else None
|
||||
)
|
||||
|
||||
|
||||
def resolve_refs(team, refs: list[dict]) -> ResolvedRefs:
|
||||
"""[Ref] → 事实文本 + 参考图。查不到的进 missing,**不抛异常** ——
|
||||
素材被别人删掉不该让整条对话崩掉,该由 agent 在对话里说明。"""
|
||||
resolved = ResolvedRefs()
|
||||
for ref in refs or []:
|
||||
type_, ref_id = ref.get("type"), ref.get("id")
|
||||
if type_ not in TYPE_LABELS or not ref_id:
|
||||
continue
|
||||
if type_ == "product":
|
||||
product = Product.objects.filter(team=team, id=ref_id, purged_at__isnull=True).first()
|
||||
if product is None:
|
||||
resolved.missing.append(ref)
|
||||
continue
|
||||
resolved.facts.append(product_facts_text(product))
|
||||
entry = _product_reference(product)
|
||||
if entry:
|
||||
resolved.references.append(entry)
|
||||
continue
|
||||
if type_ == "model":
|
||||
model = Model.objects.filter(
|
||||
team=team, id=ref_id, is_deleted=False, purged_at__isnull=True
|
||||
).select_related("triview_asset", "portrait_asset").first()
|
||||
if model is None:
|
||||
resolved.missing.append(ref)
|
||||
continue
|
||||
if (model.description or "").strip():
|
||||
resolved.facts.append(f"模特「{model.name}」:{model.description.strip()}")
|
||||
# 锁脸优先用三视图(正/侧/背都在一张 16:9 里,信息量最大),没有才回落形象图
|
||||
entry = _asset_reference(model.triview_asset, "model", model.name) or _asset_reference(
|
||||
model.portrait_asset, "model", model.name
|
||||
)
|
||||
if entry:
|
||||
resolved.references.append(entry)
|
||||
else:
|
||||
resolved.missing.append(ref)
|
||||
continue
|
||||
asset = Asset.objects.filter(
|
||||
team=team, id=ref_id, is_deleted=False, purged_at__isnull=True
|
||||
).first()
|
||||
if asset is None:
|
||||
resolved.missing.append(ref)
|
||||
continue
|
||||
if (asset.description or "").strip():
|
||||
resolved.facts.append(f"{TYPE_LABELS[type_]}「{asset.name}」:{asset.description.strip()}")
|
||||
entry = _asset_reference(asset, type_, asset.name)
|
||||
if entry:
|
||||
resolved.references.append(entry)
|
||||
else:
|
||||
resolved.missing.append(ref)
|
||||
|
||||
# 角色 → 场景 → 商品。同优先级内保持用户 @ 的先后。
|
||||
resolved.references.sort(key=lambda item: _REF_ORDER.get(item["type"], 9))
|
||||
resolved.references = _dedupe_references(resolved.references)[:MAX_REFERENCE_IMAGES]
|
||||
return resolved
|
||||
|
||||
|
||||
def _dedupe_references(references: list[dict]) -> list[dict]:
|
||||
"""同一张图被 @ 两次(比如商品图同时是资产库图)只留一条,否则 @图N 编号会错位。"""
|
||||
out: list[dict] = []
|
||||
seen: set[str] = set()
|
||||
for item in references:
|
||||
key = item.get("url") or ""
|
||||
if not key or key in seen:
|
||||
continue
|
||||
seen.add(key)
|
||||
out.append(item)
|
||||
return out
|
||||
@@ -0,0 +1,77 @@
|
||||
# Generated by Django 5.1.15 on 2026-09-02 09:45
|
||||
|
||||
import django.db.models.deletion
|
||||
import uuid
|
||||
from django.conf import settings
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('accounts', '0009_team_price_multiplier'),
|
||||
('ai', '0033_seedance_25_capabilities'),
|
||||
migrations.swappable_dependency(settings.AUTH_USER_MODEL),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.CreateModel(
|
||||
name='CreationConversation',
|
||||
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)),
|
||||
('title', models.CharField(default='未命名创作', max_length=120)),
|
||||
('mode', models.CharField(choices=[('video', '视频创作'), ('image', '图片创作')], default='video', max_length=16)),
|
||||
('preset', models.CharField(blank=True, default='', max_length=64)),
|
||||
('params', models.JSONField(blank=True, default=dict)),
|
||||
('pinned_refs', models.JSONField(blank=True, default=list)),
|
||||
('memory', models.JSONField(blank=True, default=dict)),
|
||||
('status', models.CharField(choices=[('running', '进行中'), ('completed', '已完成'), ('failed', '失败')], default='running', max_length=16)),
|
||||
('last_active_at', models.DateTimeField(auto_now_add=True)),
|
||||
('is_deleted', models.BooleanField(default=False)),
|
||||
('purged_at', models.DateTimeField(blank=True, null=True)),
|
||||
('created_by', models.ForeignKey(blank=True, null=True, on_delete=django.db.models.deletion.SET_NULL, related_name='created_%(class)s_set', to=settings.AUTH_USER_MODEL)),
|
||||
('team', models.ForeignKey(on_delete=django.db.models.deletion.CASCADE, related_name='%(class)s_set', to='accounts.team')),
|
||||
],
|
||||
),
|
||||
migrations.CreateModel(
|
||||
name='CreationMessage',
|
||||
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)),
|
||||
('role', models.CharField(choices=[('user', '用户'), ('assistant', 'AI'), ('system', '系统')], max_length=16)),
|
||||
('kind', models.CharField(choices=[('text', '文字气泡'), ('elicit', '追问卡'), ('strategy', '创作策略理解卡'), ('plan', '视频最终方案卡'), ('prompt_file', '生成 Prompt 文件卡'), ('confirm', '确认闸门(带预计积分)'), ('generating', '生成中'), ('result', '生成结果'), ('error', '错误')], default='text', max_length=24)),
|
||||
('text', models.TextField(blank=True, default='')),
|
||||
('payload', models.JSONField(blank=True, default=dict)),
|
||||
('refs', models.JSONField(blank=True, default=list)),
|
||||
('seq', models.PositiveIntegerField(default=0)),
|
||||
('conversation', models.ForeignKey(on_delete=django.db.models.deletion.CASCADE, related_name='messages', to='ai.creationconversation')),
|
||||
('task', models.ForeignKey(blank=True, null=True, on_delete=django.db.models.deletion.SET_NULL, related_name='creation_messages', to='ai.aitask')),
|
||||
],
|
||||
options={
|
||||
'ordering': ['seq', 'created_at'],
|
||||
},
|
||||
),
|
||||
migrations.AddIndex(
|
||||
model_name='creationconversation',
|
||||
index=models.Index(fields=['team', '-last_active_at'], name='ai_creation_team_id_82f111_idx'),
|
||||
),
|
||||
migrations.AddIndex(
|
||||
model_name='creationconversation',
|
||||
index=models.Index(fields=['team', 'status', '-last_active_at'], name='ai_creation_team_id_3aa2e5_idx'),
|
||||
),
|
||||
migrations.AddIndex(
|
||||
model_name='creationconversation',
|
||||
index=models.Index(fields=['team', 'is_deleted', 'purged_at'], name='ai_creation_team_id_26e56a_idx'),
|
||||
),
|
||||
migrations.AddIndex(
|
||||
model_name='creationmessage',
|
||||
index=models.Index(fields=['conversation', 'seq'], name='ai_creation_convers_e9c745_idx'),
|
||||
),
|
||||
migrations.AddConstraint(
|
||||
model_name='creationmessage',
|
||||
constraint=models.UniqueConstraint(fields=('conversation', 'seq'), name='uniq_creation_message_seq'),
|
||||
),
|
||||
]
|
||||
@@ -279,3 +279,108 @@ class PromptTemplate(TimeStampedModel):
|
||||
def __str__(self) -> str:
|
||||
return f"prompt:{self.key}"
|
||||
|
||||
|
||||
class CreationConversation(TeamOwnedModel):
|
||||
"""全能创作的会话。一条会话 = 一次完整创作(可多轮改稿、多次出图/出片)。
|
||||
|
||||
和 ImageConversation 的区别:那张表只是「生图线程」,没有消息实体,历史靠 AITask 拼;
|
||||
这里是真对话,消息落 CreationMessage。两者并存,互不影响(图片工作台仍走旧表)。
|
||||
|
||||
mode 发起时定死,会话内不可切(设计稿顶栏的模型/分辨率/比例跟着 mode 固定)。
|
||||
pinned_refs 是「实体锁定」:本会话引用过的商品/角色/场景,每轮无条件带进上下文 ——
|
||||
这是多次生成之间锁脸、锁商品的唯一手段,别省。
|
||||
"""
|
||||
|
||||
class Mode(models.TextChoices):
|
||||
VIDEO = "video", "视频创作"
|
||||
IMAGE = "image", "图片创作"
|
||||
|
||||
class Status(models.TextChoices):
|
||||
RUNNING = "running", "进行中"
|
||||
COMPLETED = "completed", "已完成"
|
||||
FAILED = "failed", "失败"
|
||||
|
||||
title = models.CharField(max_length=120, default="未命名创作")
|
||||
mode = models.CharField(max_length=16, choices=Mode.choices, default=Mode.VIDEO)
|
||||
preset = models.CharField(max_length=64, blank=True, default="") # "" = 自由创作
|
||||
# 会话级参数:{model, resolution, ratio, duration} —— 设计稿顶栏 meta 就渲染它
|
||||
params = models.JSONField(default=dict, blank=True)
|
||||
# 实体锁定:[Ref],见契约 §1。每轮无条件带上
|
||||
pinned_refs = models.JSONField(default=list, blank=True)
|
||||
# 记忆:{summary, artifacts:[{msg_id,asset_id,prompt,kind}], turn_count}
|
||||
memory = models.JSONField(default=dict, blank=True)
|
||||
status = models.CharField(max_length=16, choices=Status.choices, default=Status.RUNNING)
|
||||
last_active_at = models.DateTimeField(auto_now_add=True)
|
||||
is_deleted = models.BooleanField(default=False)
|
||||
purged_at = models.DateTimeField(null=True, blank=True)
|
||||
|
||||
class Meta:
|
||||
indexes = [
|
||||
# 创作历史页:按团队 + 最近活跃倒序
|
||||
models.Index(fields=["team", "-last_active_at"]),
|
||||
models.Index(fields=["team", "status", "-last_active_at"]),
|
||||
models.Index(fields=["team", "is_deleted", "purged_at"]),
|
||||
]
|
||||
|
||||
def __str__(self) -> str:
|
||||
return f"creation:{self.mode}:{self.title}"
|
||||
|
||||
|
||||
class CreationMessage(TimeStampedModel):
|
||||
"""全能创作对话流里的一条消息。kind 决定前端渲染成哪种卡片(见契约 §2)。
|
||||
|
||||
text 只给 TEXT 用;其余 kind 的内容全在 payload 里 —— 前端按 kind 走不同组件,
|
||||
别把结构化内容塞进 text 再让前端解析。
|
||||
|
||||
生成类消息(GENERATING / RESULT)挂 task:提交时先落一条 GENERATING,
|
||||
轮询到终态后**原地改成 RESULT**(不新增消息),这样对话流不会被中间态刷屏。
|
||||
"""
|
||||
|
||||
class Role(models.TextChoices):
|
||||
USER = "user", "用户"
|
||||
ASSISTANT = "assistant", "AI"
|
||||
SYSTEM = "system", "系统"
|
||||
|
||||
class Kind(models.TextChoices):
|
||||
TEXT = "text", "文字气泡"
|
||||
ELICIT = "elicit", "追问卡" # AI 反问用户,带 单选/多选/填空/选素材 控件
|
||||
STRATEGY = "strategy", "创作策略理解卡"
|
||||
PLAN = "plan", "视频最终方案卡"
|
||||
PROMPT_FILE = "prompt_file", "生成 Prompt 文件卡"
|
||||
CONFIRM = "confirm", "确认闸门(带预计积分)"
|
||||
GENERATING = "generating", "生成中"
|
||||
RESULT = "result", "生成结果"
|
||||
ERROR = "error", "错误"
|
||||
|
||||
conversation = models.ForeignKey(
|
||||
CreationConversation, on_delete=models.CASCADE, related_name="messages"
|
||||
)
|
||||
role = models.CharField(max_length=16, choices=Role.choices)
|
||||
kind = models.CharField(max_length=24, choices=Kind.choices, default=Kind.TEXT)
|
||||
text = models.TextField(blank=True, default="")
|
||||
payload = models.JSONField(default=dict, blank=True)
|
||||
# 本条消息引用的实体:[Ref]。用 id 取事实与参考图,不许只存 "@商品名" 字符串
|
||||
refs = models.JSONField(default=list, blank=True)
|
||||
task = models.ForeignKey(
|
||||
AITask,
|
||||
on_delete=models.SET_NULL,
|
||||
null=True,
|
||||
blank=True,
|
||||
related_name="creation_messages",
|
||||
)
|
||||
# 会话内自增渲染序;并发插入靠 select_for_update 取 max+1
|
||||
seq = models.PositiveIntegerField(default=0)
|
||||
|
||||
class Meta:
|
||||
ordering = ["seq", "created_at"]
|
||||
indexes = [
|
||||
models.Index(fields=["conversation", "seq"]),
|
||||
]
|
||||
constraints = [
|
||||
models.UniqueConstraint(
|
||||
fields=["conversation", "seq"], name="uniq_creation_message_seq"
|
||||
),
|
||||
]
|
||||
|
||||
def __str__(self) -> str:
|
||||
return f"msg:{self.kind}:{self.seq}"
|
||||
|
||||
@@ -1,6 +1,13 @@
|
||||
from rest_framework import serializers
|
||||
|
||||
from .models import AITask, ImageConversation, ModelConfig, ModelProvider
|
||||
from .models import (
|
||||
AITask,
|
||||
CreationConversation,
|
||||
CreationMessage,
|
||||
ImageConversation,
|
||||
ModelConfig,
|
||||
ModelProvider,
|
||||
)
|
||||
|
||||
|
||||
class ModelProviderSerializer(serializers.ModelSerializer):
|
||||
@@ -89,3 +96,71 @@ class AITaskSerializer(serializers.ModelSerializer):
|
||||
# batch_id / mode 是显式声明的 SerializerMethodField(本就只读),不能再列进 read_only_fields(DRF 会报错)
|
||||
read_only_fields = [f for f in fields if f not in ("batch_id", "mode")]
|
||||
|
||||
|
||||
|
||||
class CreationMessageSerializer(serializers.ModelSerializer):
|
||||
"""全能创作对话流里的一条消息。前端**按 kind 分发到不同卡片组件**,
|
||||
结构化内容一律在 payload 里(契约 §2),不要从 text 里解析。"""
|
||||
|
||||
class Meta:
|
||||
model = CreationMessage
|
||||
fields = ["id", "role", "kind", "text", "payload", "refs", "task", "seq", "created_at"]
|
||||
read_only_fields = fields
|
||||
|
||||
|
||||
class CreationConversationSerializer(serializers.ModelSerializer):
|
||||
"""会话列表 / 详情。title 可写(重命名);params 创建后也可改(对话页改模型/比例后立刻生效);
|
||||
mode 创建后不可改。"""
|
||||
|
||||
message_count = serializers.SerializerMethodField()
|
||||
cover_url = serializers.SerializerMethodField()
|
||||
|
||||
class Meta:
|
||||
model = CreationConversation
|
||||
fields = [
|
||||
"id", "title", "mode", "preset", "params", "status",
|
||||
"message_count", "cover_url",
|
||||
"last_active_at", "created_at", "updated_at",
|
||||
]
|
||||
read_only_fields = [
|
||||
"id", "status", "message_count", "cover_url",
|
||||
"last_active_at", "created_at", "updated_at",
|
||||
]
|
||||
|
||||
def get_message_count(self, obj) -> int:
|
||||
cached = getattr(obj, "_message_count", None)
|
||||
return cached if cached is not None else obj.messages.count()
|
||||
|
||||
def get_cover_url(self, obj) -> str:
|
||||
"""历史页封面 = **最新一版**结果(重生成是往下叠加,所以取最后一条 RESULT)。"""
|
||||
last = (
|
||||
obj.messages.filter(kind=CreationMessage.Kind.RESULT)
|
||||
.order_by("-seq")
|
||||
.values_list("payload", flat=True)
|
||||
.first()
|
||||
)
|
||||
if not last:
|
||||
return ""
|
||||
assets = (last or {}).get("assets") or []
|
||||
if not assets:
|
||||
return ""
|
||||
first = assets[0] or {}
|
||||
return first.get("cover") or first.get("url") or ""
|
||||
|
||||
def update(self, instance, validated_data):
|
||||
# mode 定死:允许传但忽略,避免前端误改后顶栏参数与已生成内容对不上
|
||||
validated_data.pop("mode", None)
|
||||
return super().update(instance, validated_data)
|
||||
|
||||
|
||||
class CreationConversationDetailSerializer(CreationConversationSerializer):
|
||||
"""详情:带全量消息,进对话页一次性回填。"""
|
||||
|
||||
messages = CreationMessageSerializer(many=True, read_only=True)
|
||||
pinned_refs = serializers.JSONField(read_only=True)
|
||||
|
||||
class Meta(CreationConversationSerializer.Meta):
|
||||
fields = [*CreationConversationSerializer.Meta.fields, "messages", "pinned_refs"]
|
||||
read_only_fields = [
|
||||
*CreationConversationSerializer.Meta.read_only_fields, "messages", "pinned_refs",
|
||||
]
|
||||
|
||||
@@ -3883,6 +3883,11 @@ def run_standalone_image_task(*, task_id: str) -> None:
|
||||
task=task, project=task.project, recipient=user,
|
||||
stage_label="图片创作", raw=str(exc), hint=friendly_generation_error(str(exc)),
|
||||
)
|
||||
from apps.ai.creation import sync_generating_for_task
|
||||
|
||||
# 全能创作挂在这条任务上的 GENERATING 要立刻改成 RESULT/ERROR,
|
||||
# 不能干等前端下一次轮询 —— 否则页面会一直停在「正在生成」。
|
||||
sync_generating_for_task(task)
|
||||
|
||||
|
||||
# ── 旁白配音(TTS):每镜旁白合成一段语音,导出时作为人声轨混在 BGM 之上 ──
|
||||
|
||||
@@ -0,0 +1,733 @@
|
||||
"""全能创作 · Agent 循环与 SSE(契约 §3/§4)。
|
||||
|
||||
用假 provider 逐帧回放模型输出,验证的是**编排**而不是模型质量:
|
||||
工具调用拼装、追问中断、计费闸门、参考图带入、错误不打崩流。
|
||||
"""
|
||||
import json
|
||||
from unittest.mock import patch
|
||||
|
||||
from django.test import TestCase
|
||||
from rest_framework.test import APIClient
|
||||
|
||||
from apps.accounts.models import Team, TeamMember, User
|
||||
from apps.assets.models import Asset, AssetFile
|
||||
from apps.products.models import Product, ProductImage
|
||||
|
||||
from .creation import append_message
|
||||
from .creation_agent import (
|
||||
COMPRESS_MIN_BATCH,
|
||||
DEFAULT_VIDEO_MODEL,
|
||||
KEEP_RECENT_MESSAGES,
|
||||
SMART_DURATION,
|
||||
_coerce_fields,
|
||||
_image_count,
|
||||
_merge_tool_call_deltas,
|
||||
stream_creation_agent,
|
||||
submit_confirmed_video,
|
||||
video_duration,
|
||||
video_model_name,
|
||||
wanted_asset_pick,
|
||||
wanted_param_keys,
|
||||
)
|
||||
from .models import AITask, CreationConversation, CreationMessage, ModelConfig, ModelProvider
|
||||
|
||||
|
||||
def _text_chunks(text):
|
||||
for piece in text:
|
||||
yield {"type": "delta", "text": piece}
|
||||
yield {"type": "done"}
|
||||
|
||||
|
||||
def _tool_chunks(name, arguments, *, said=""):
|
||||
"""模拟 OpenAI 流式:arguments 被拆成多片下发。"""
|
||||
for piece in said:
|
||||
yield {"type": "delta", "text": piece}
|
||||
yield {"type": "tool_call", "tool_calls": [{"index": 0, "function": {"name": name, "arguments": ""}}]}
|
||||
blob = json.dumps(arguments, ensure_ascii=False)
|
||||
for i in range(0, len(blob), 7):
|
||||
yield {"type": "tool_call", "tool_calls": [{"index": 0, "function": {"arguments": blob[i:i + 7]}}]}
|
||||
yield {"type": "done"}
|
||||
|
||||
|
||||
class FakeProvider:
|
||||
"""按脚本逐轮回放。每调用一次 chat_completion_stream 消费一个剧本。"""
|
||||
|
||||
def __init__(self, scripts):
|
||||
self.scripts = list(scripts)
|
||||
self.calls = []
|
||||
|
||||
def chat_completion_stream(self, **kwargs):
|
||||
self.calls.append(kwargs)
|
||||
if not self.scripts:
|
||||
return iter([{"type": "done"}])
|
||||
return self.scripts.pop(0)
|
||||
|
||||
|
||||
def _events(stream):
|
||||
out = []
|
||||
for frame in stream:
|
||||
for line in frame.strip().splitlines():
|
||||
if line.startswith("data: "):
|
||||
out.append(json.loads(line[6:]))
|
||||
return out
|
||||
|
||||
|
||||
class CreationAgentBaseTests(TestCase):
|
||||
def setUp(self):
|
||||
self.user = User.objects.create_user(username="agent-owner", password="p")
|
||||
self.team = Team.objects.create(name="Agent", owner=self.user)
|
||||
provider = ModelProvider.objects.create(name="fake", display_name="Fake", base_url="https://x")
|
||||
self.model = ModelConfig.objects.create(
|
||||
provider=provider, name="fake-text", display_name="Fake Text",
|
||||
capability=ModelConfig.Capability.TEXT, endpoint="chat/completions",
|
||||
)
|
||||
self.conversation = CreationConversation.objects.create(
|
||||
team=self.team, created_by=self.user, mode="image", title="出图",
|
||||
params={"ratio": "1:1", "model": "Seedream5.0"},
|
||||
)
|
||||
|
||||
def _run(self, scripts, text="来一张商品图", refs=None):
|
||||
fake = FakeProvider(scripts)
|
||||
with patch("apps.ai.creation_agent.build_provider", return_value=fake):
|
||||
events = _events(stream_creation_agent(
|
||||
conversation=self.conversation, user=self.user,
|
||||
text=text, refs=refs or [], model_config=self.model,
|
||||
))
|
||||
return events, fake
|
||||
|
||||
|
||||
class ToolCallAssemblyTests(TestCase):
|
||||
def test_streamed_arguments_are_concatenated_not_overwritten(self):
|
||||
buffer = {}
|
||||
_merge_tool_call_deltas(buffer, [{"index": 0, "function": {"name": "ask_user", "arguments": '{"a"'}}])
|
||||
_merge_tool_call_deltas(buffer, [{"index": 0, "function": {"arguments": ':1}'}}])
|
||||
# 覆盖式赋值会只剩最后一片,tool call 直接废掉
|
||||
self.assertEqual(buffer[0], {"name": "ask_user", "arguments": '{"a":1}'})
|
||||
|
||||
def test_multiple_parallel_tool_calls_keep_separate_slots(self):
|
||||
buffer = {}
|
||||
_merge_tool_call_deltas(buffer, [
|
||||
{"index": 0, "function": {"name": "search_library", "arguments": "{}"}},
|
||||
{"index": 1, "function": {"name": "ask_user", "arguments": "{}"}},
|
||||
])
|
||||
self.assertEqual({buffer[0]["name"], buffer[1]["name"]}, {"search_library", "ask_user"})
|
||||
|
||||
|
||||
class ImageCountTests(TestCase):
|
||||
def test_session_sheet_count_wins(self):
|
||||
self.assertEqual(_image_count({"count": "2 张"}, 4), 2)
|
||||
self.assertEqual(_image_count({"duration": "4 张"}, None), 4)
|
||||
|
||||
def test_tool_count_used_when_user_did_not_pick(self):
|
||||
self.assertEqual(_image_count({"duration": "智能时长"}, 3), 3)
|
||||
self.assertEqual(_image_count({}, None), 1)
|
||||
self.assertEqual(_image_count({}, 9), 8)
|
||||
|
||||
|
||||
class FieldCoercionTests(TestCase):
|
||||
def test_single_choice_without_options_is_dropped(self):
|
||||
fields = _coerce_fields([{"key": "tone", "label": "什么调性?", "type": "single"}])
|
||||
self.assertEqual(fields, []) # 没选项的单选是废卡
|
||||
|
||||
def test_text_and_asset_fields_get_their_defaults(self):
|
||||
fields = _coerce_fields([
|
||||
{"key": "slogan", "label": "想突出哪句话?", "type": "text"},
|
||||
{"key": "who", "label": "用哪个模特?", "type": "asset"},
|
||||
])
|
||||
self.assertEqual(fields[0]["placeholder"], "")
|
||||
self.assertIn("product", fields[1]["asset_types"])
|
||||
|
||||
def test_unknown_type_and_over_limit_are_trimmed(self):
|
||||
raw = [{"key": f"k{i}", "label": "x", "type": "text"} for i in range(5)]
|
||||
raw.append({"key": "bad", "label": "x", "type": "dropdown"})
|
||||
self.assertEqual(len(_coerce_fields(raw)), 3) # 一次最多问 3 项
|
||||
|
||||
def test_product_text_or_single_is_coerced_to_asset_card(self):
|
||||
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"])
|
||||
|
||||
def test_duration_single_stays_single(self):
|
||||
fields = _coerce_fields([
|
||||
{"key": "duration", "label": "改成多长?", "type": "single",
|
||||
"options": [{"value": "10 秒", "label": "10 秒"}]},
|
||||
])
|
||||
self.assertEqual(fields[0]["type"], "single")
|
||||
self.assertEqual(fields[0]["options"][0]["value"], "10 秒")
|
||||
|
||||
|
||||
class AssetPickIntentTests(TestCase):
|
||||
def test_change_product_without_at_should_open_picker(self):
|
||||
self.assertEqual(wanted_asset_pick("我想修改商品", []), "product")
|
||||
self.assertEqual(wanted_asset_pick("换个角色", []), "character")
|
||||
|
||||
def test_named_or_already_referenced_does_not_force_picker(self):
|
||||
self.assertIsNone(wanted_asset_pick("换成净颜精华", []))
|
||||
self.assertIsNone(wanted_asset_pick("我想修改商品", [{"type": "product", "id": "1"}]))
|
||||
|
||||
|
||||
class ParamPickIntentTests(TestCase):
|
||||
def test_change_duration_opens_duration_card(self):
|
||||
self.assertEqual(wanted_param_keys("我想改时长", is_video=True), ["duration"])
|
||||
self.assertEqual(wanted_param_keys("改成 10 秒", is_video=True), ["duration"])
|
||||
|
||||
def test_change_model_does_not_mean_character_model(self):
|
||||
self.assertEqual(wanted_param_keys("换个模型", is_video=True), ["video_model"])
|
||||
self.assertEqual(wanted_param_keys("改模特", is_video=True), [])
|
||||
|
||||
|
||||
class AskUserTests(CreationAgentBaseTests):
|
||||
def test_saying_change_duration_injects_param_card(self):
|
||||
events, _ = self._run([_text_chunks("你想改成多长?")], text="我想改时长")
|
||||
elicit = [e for e in events if e.get("type") == "message" and e["message"]["kind"] == "elicit"]
|
||||
self.assertEqual(len(elicit), 1)
|
||||
self.assertEqual(elicit[0]["message"]["payload"]["fields"][0]["key"], "duration")
|
||||
self.assertEqual(elicit[0]["message"]["payload"]["fields"][0]["type"], "single")
|
||||
|
||||
def test_saying_change_product_injects_asset_card(self):
|
||||
events, _ = self._run([_text_chunks("换成哪个商品?直接选或填名字都行:")], text="我想修改商品")
|
||||
elicit = [e for e in events if e.get("type") == "message" and e["message"]["kind"] == "elicit"]
|
||||
self.assertEqual(len(elicit), 1)
|
||||
self.assertEqual(elicit[0]["message"]["payload"]["fields"][0]["type"], "asset")
|
||||
self.assertEqual(elicit[0]["message"]["payload"]["fields"][0]["asset_types"], ["product"])
|
||||
|
||||
def test_ask_user_emits_elicit_card_and_stops_the_loop(self):
|
||||
events, fake = self._run([
|
||||
_tool_chunks("ask_user", {"fields": [
|
||||
{"key": "tone", "label": "想要什么调性?", "type": "single",
|
||||
"options": [{"value": "warm", "label": "温暖生活感"}, {"value": "cool", "label": "冷淡高级感"}]},
|
||||
]}),
|
||||
_text_chunks("不该跑到这一轮"),
|
||||
])
|
||||
|
||||
elicit = [e for e in events if e.get("type") == "message" and e["message"]["kind"] == "elicit"]
|
||||
self.assertEqual(len(elicit), 1)
|
||||
self.assertEqual(len(elicit[0]["message"]["payload"]["fields"][0]["options"]), 2)
|
||||
# 反问必须中断循环,否则模型会自问自答
|
||||
self.assertEqual(len(fake.calls), 1)
|
||||
self.assertEqual(events[-1]["type"], "done")
|
||||
|
||||
def test_malformed_fields_do_not_stop_the_conversation(self):
|
||||
events, fake = self._run([
|
||||
_tool_chunks("ask_user", {"fields": [{"key": "x", "label": "y", "type": "dropdown"}]}),
|
||||
_text_chunks("那我直接来了"),
|
||||
])
|
||||
kinds = [e["message"]["kind"] for e in events if e.get("type") == "message"]
|
||||
self.assertNotIn("elicit", kinds)
|
||||
self.assertEqual(len(fake.calls), 2) # 废卡不算反问,循环继续
|
||||
|
||||
|
||||
class GenerateImageTests(CreationAgentBaseTests):
|
||||
def _fake_task(self, key):
|
||||
"""真 AITask —— GENERATING 消息要把它挂上 FK,假对象赋不进去。"""
|
||||
return AITask.objects.create(
|
||||
team=self.team, created_by=self.user, task_type=AITask.Type.PRODUCT_IMAGE,
|
||||
model_config=self.model, idempotency_key=key,
|
||||
)
|
||||
|
||||
def setUp(self):
|
||||
super().setUp()
|
||||
self.product = Product.objects.create(team=self.team, created_by=self.user, title="净颜精华")
|
||||
asset = Asset.objects.create(
|
||||
team=self.team, created_by=self.user, name="商品实拍",
|
||||
asset_type=Asset.Type.IMAGE, source=Asset.Source.UPLOAD,
|
||||
category=Asset.Category.PRODUCT_IMAGE,
|
||||
)
|
||||
AssetFile.objects.create(asset=asset, object_key="k", bucket="b", is_primary=True,
|
||||
preview_url="https://cdn/prod.jpg")
|
||||
ProductImage.objects.create(product=self.product, asset=asset, is_primary=True)
|
||||
self.product_asset = asset
|
||||
|
||||
def test_generate_image_submits_with_pinned_reference_and_emits_task(self):
|
||||
with patch("apps.ai.services.enqueue_standalone_images") as enqueue:
|
||||
enqueue.return_value = [self._fake_task("k-ref")]
|
||||
events, _ = self._run(
|
||||
[_tool_chunks("generate_image", {"prompt": "干净棚拍,柔光,居中构图"})],
|
||||
refs=[{"type": "product", "id": str(self.product.id), "name": "净颜精华"}],
|
||||
)
|
||||
kwargs = enqueue.call_args.kwargs
|
||||
|
||||
# @ 引用的商品图必须作为参考图带上,否则出的图跟商品长得不一样
|
||||
self.assertEqual(kwargs["reference_image_ids"], [str(self.product_asset.id)])
|
||||
self.assertEqual(kwargs["ratio"], "1:1") # 会话级参数直接用,不再问用户
|
||||
self.assertEqual(kwargs["count"], 1)
|
||||
|
||||
task_events = [e for e in events if e.get("type") == "task"]
|
||||
self.assertEqual(len(task_events), 1)
|
||||
generating = [e for e in events if e.get("type") == "message" and e["message"]["kind"] == "generating"]
|
||||
self.assertEqual(len(generating), 1)
|
||||
|
||||
def test_session_image_count_overrides_model_count(self):
|
||||
self.conversation.params = {"ratio": "1:1", "count": "2 张"}
|
||||
self.conversation.save(update_fields=["params"])
|
||||
with patch("apps.ai.services.enqueue_standalone_images") as enqueue:
|
||||
enqueue.return_value = [self._fake_task("k-count-1"), self._fake_task("k-count-2")]
|
||||
events, _ = self._run(
|
||||
[_tool_chunks("generate_image", {"prompt": "白底棚拍", "count": 1})],
|
||||
)
|
||||
self.assertEqual(enqueue.call_args.kwargs["count"], 2)
|
||||
generating = [e for e in events if e.get("type") == "message" and e["message"]["kind"] == "generating"]
|
||||
self.assertEqual(len(generating), 2)
|
||||
|
||||
def test_second_generation_in_one_turn_is_blocked(self):
|
||||
with patch("apps.ai.services.enqueue_standalone_images") as enqueue:
|
||||
enqueue.return_value = [self._fake_task("k-twice")]
|
||||
self._run([
|
||||
_tool_chunks("generate_image", {"prompt": "第一版"}),
|
||||
_tool_chunks("generate_image", {"prompt": "第二版"}),
|
||||
])
|
||||
# 一条用户消息只计费一次:一句「多做几版」不能烧掉一堆积分
|
||||
self.assertEqual(enqueue.call_count, 1)
|
||||
|
||||
def test_prompt_is_remembered_for_the_next_revision(self):
|
||||
with patch("apps.ai.services.enqueue_standalone_images") as enqueue:
|
||||
enqueue.return_value = [self._fake_task("k-memory")]
|
||||
self._run([_tool_chunks("generate_image", {"prompt": "白底棚拍"})])
|
||||
self.conversation.refresh_from_db()
|
||||
# 产物索引:下一轮「背景换夜景」要靠它知道在改哪一版
|
||||
self.assertEqual(self.conversation.memory["artifacts"][-1]["prompt"], "白底棚拍")
|
||||
|
||||
def test_video_conversation_is_not_offered_the_image_tool(self):
|
||||
video = CreationConversation.objects.create(
|
||||
team=self.team, created_by=self.user, mode="video", params={}
|
||||
)
|
||||
fake = FakeProvider([_text_chunks("先聊聊")])
|
||||
with patch("apps.ai.creation_agent.build_provider", return_value=fake):
|
||||
list(stream_creation_agent(conversation=video, user=self.user, text="做条视频",
|
||||
model_config=self.model))
|
||||
names = {t["function"]["name"] for t in fake.calls[0]["extra_body"]["tools"]}
|
||||
# 会话 mode 定死,摆出不该用的工具只会诱导模型走错路
|
||||
self.assertNotIn("generate_image", names)
|
||||
|
||||
|
||||
class MissingRefAndFailureTests(CreationAgentBaseTests):
|
||||
def test_deleted_ref_is_reported_but_conversation_continues(self):
|
||||
events, fake = self._run(
|
||||
[_text_chunks("好")],
|
||||
refs=[{"type": "product", "id": "00000000-0000-0000-0000-000000000000", "name": "已删商品"}],
|
||||
)
|
||||
texts = [e["message"]["text"] for e in events if e.get("type") == "message"]
|
||||
self.assertTrue(any("已删商品" in t for t in texts))
|
||||
self.assertEqual(len(fake.calls), 1) # 照常继续,不打断
|
||||
|
||||
def test_provider_blowup_yields_error_event_not_a_hang(self):
|
||||
def explode(**kwargs):
|
||||
raise RuntimeError("provider down")
|
||||
|
||||
fake = FakeProvider([])
|
||||
fake.chat_completion_stream = explode
|
||||
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,
|
||||
))
|
||||
# 未捕获异常会让前端白屏卡死,必须收成 error 事件
|
||||
self.assertEqual(events[-1]["type"], "error")
|
||||
|
||||
def test_no_text_model_configured_fails_fast(self):
|
||||
ModelConfig.objects.update(status=ModelConfig.Status.DISABLED)
|
||||
events = _events(stream_creation_agent(
|
||||
conversation=self.conversation, user=self.user, text="来一张", model_config=None,
|
||||
))
|
||||
self.assertEqual(events[0]["type"], "error")
|
||||
|
||||
|
||||
class SseFramingTests(CreationAgentBaseTests):
|
||||
def test_messages_carrying_uuid_and_datetime_are_serializable(self):
|
||||
"""GENERATING 消息带 task 外键(UUID)和 created_at(datetime)。
|
||||
用标准 json.dumps 会当场 TypeError 把整条流打断 —— 必须走 DjangoJSONEncoder。"""
|
||||
task = AITask.objects.create(
|
||||
team=self.team, created_by=self.user, task_type=AITask.Type.PRODUCT_IMAGE,
|
||||
model_config=self.model, idempotency_key="k-sse",
|
||||
)
|
||||
with patch("apps.ai.services.enqueue_standalone_images", return_value=[task]):
|
||||
events, _ = self._run([_tool_chunks("generate_image", {"prompt": "白底棚拍"})])
|
||||
|
||||
self.assertNotIn("error", [e.get("type") for e in events])
|
||||
generating = next(e for e in events if e.get("type") == "message"
|
||||
and e["message"]["kind"] == "generating")
|
||||
self.assertEqual(generating["message"]["task"], str(task.id))
|
||||
|
||||
|
||||
class SendEndpointTests(TestCase):
|
||||
def setUp(self):
|
||||
self.user = User.objects.create_user(username="send-owner", password="p")
|
||||
self.team = Team.objects.create(name="Send", owner=self.user)
|
||||
TeamMember.objects.create(team=self.team, user=self.user, role="owner")
|
||||
self.conversation = CreationConversation.objects.create(
|
||||
team=self.team, created_by=self.user, mode="image"
|
||||
)
|
||||
self.client = APIClient()
|
||||
self.client.force_authenticate(self.user)
|
||||
|
||||
def test_empty_message_is_rejected(self):
|
||||
response = self.client.post(f"/api/ai/creations/{self.conversation.id}/send/", {}, format="json")
|
||||
self.assertEqual(response.status_code, 400)
|
||||
|
||||
def test_answering_an_elicit_card_marks_it_submitted(self):
|
||||
card = append_message(
|
||||
self.conversation, role="assistant", kind=CreationMessage.Kind.ELICIT,
|
||||
payload={"fields": [{"key": "tone", "label": "什么调性?", "type": "single",
|
||||
"options": [{"value": "warm", "label": "温暖"}]}],
|
||||
"submitted": False, "answers": {}},
|
||||
)
|
||||
fake = FakeProvider([_text_chunks("收到")])
|
||||
with patch("apps.ai.creation_agent.build_provider", return_value=fake):
|
||||
response = self.client.post(
|
||||
f"/api/ai/creations/{self.conversation.id}/send/",
|
||||
{"kind": "elicit_answer", "reply_to": str(card.id), "answers": {"tone": "warm"}},
|
||||
format="json",
|
||||
)
|
||||
list(response.streaming_content)
|
||||
|
||||
card.refresh_from_db()
|
||||
self.assertTrue(card.payload["submitted"])
|
||||
self.assertEqual(card.payload["answers"], {"tone": "warm"})
|
||||
|
||||
def test_elicit_product_choice_pins_the_product(self):
|
||||
"""对话里点选商品只回 answers 时,也必须钉进 pinned_refs,否则出片带不上商品图。"""
|
||||
product = Product.objects.create(team=self.team, created_by=self.user, title="净颜精华")
|
||||
card = append_message(
|
||||
self.conversation, role="assistant", kind=CreationMessage.Kind.ELICIT,
|
||||
payload={"fields": [{"key": "product", "label": "选择商品", "type": "single",
|
||||
"options": [{"value": str(product.id), "label": product.title}]}],
|
||||
"submitted": False, "answers": {}},
|
||||
)
|
||||
fake = FakeProvider([_text_chunks("收到")])
|
||||
with patch("apps.ai.creation_agent.build_provider", return_value=fake):
|
||||
response = self.client.post(
|
||||
f"/api/ai/creations/{self.conversation.id}/send/",
|
||||
{"kind": "elicit_answer", "reply_to": str(card.id),
|
||||
"answers": {"product": str(product.id)}},
|
||||
format="json",
|
||||
)
|
||||
list(response.streaming_content)
|
||||
|
||||
self.conversation.refresh_from_db()
|
||||
pinned = self.conversation.pinned_refs or []
|
||||
self.assertTrue(
|
||||
any(r.get("type") == "product" and str(r.get("id")) == str(product.id) for r in pinned)
|
||||
)
|
||||
|
||||
def test_elicit_product_choice_by_name_pins_the_product(self):
|
||||
product = Product.objects.create(team=self.team, created_by=self.user, title="控油洁面")
|
||||
card = append_message(
|
||||
self.conversation, role="assistant", kind=CreationMessage.Kind.ELICIT,
|
||||
payload={"fields": [{"key": "product", "label": "选择商品", "type": "single",
|
||||
"options": [{"value": product.title, "label": product.title}]}],
|
||||
"submitted": False, "answers": {}},
|
||||
)
|
||||
fake = FakeProvider([_text_chunks("收到")])
|
||||
with patch("apps.ai.creation_agent.build_provider", return_value=fake):
|
||||
response = self.client.post(
|
||||
f"/api/ai/creations/{self.conversation.id}/send/",
|
||||
{"kind": "elicit_answer", "reply_to": str(card.id),
|
||||
"answers": {"product": product.title}},
|
||||
format="json",
|
||||
)
|
||||
list(response.streaming_content)
|
||||
|
||||
self.conversation.refresh_from_db()
|
||||
pinned = self.conversation.pinned_refs or []
|
||||
self.assertTrue(
|
||||
any(r.get("type") == "product" and str(r.get("id")) == str(product.id) for r in pinned)
|
||||
)
|
||||
|
||||
def test_answering_the_same_card_twice_is_refused(self):
|
||||
card = append_message(
|
||||
self.conversation, role="assistant", kind=CreationMessage.Kind.ELICIT,
|
||||
payload={"fields": [], "submitted": True, "answers": {"tone": "warm"}},
|
||||
)
|
||||
response = self.client.post(
|
||||
f"/api/ai/creations/{self.conversation.id}/send/",
|
||||
{"kind": "elicit_answer", "reply_to": str(card.id), "answers": {"tone": "cool"}},
|
||||
format="json",
|
||||
)
|
||||
# 重复提交会让同一个问题在上下文里出现两次答案
|
||||
self.assertEqual(response.status_code, 409)
|
||||
|
||||
|
||||
class VideoPlanAndConfirmTests(CreationAgentBaseTests):
|
||||
"""视频链路:策略卡 → 方案卡 → 确认闸门 → 出片(契约 §0)。"""
|
||||
|
||||
def setUp(self):
|
||||
super().setUp()
|
||||
self.conversation = CreationConversation.objects.create(
|
||||
team=self.team, created_by=self.user, mode="video", title="口播",
|
||||
params={"model": "Seedance 2.5", "ratio": "9:16", "resolution": "720p", "duration": "15 秒"},
|
||||
)
|
||||
|
||||
def _plan_args(self, **overrides):
|
||||
args = {
|
||||
"usp": "核心效果:一整天不泛油光",
|
||||
"points": ["质地轻薄"],
|
||||
"timeline": [{"start": 0, "end": 2.7, "stage": "Hook"}],
|
||||
"matrix": {"shots": 4, "rows": [{"point": "USP", "hits": [1, 3]}]},
|
||||
"voice_chars": [51, 60],
|
||||
"video_prompt": "0-3秒 近景手持商品…",
|
||||
}
|
||||
args.update(overrides)
|
||||
return args
|
||||
|
||||
def test_image_tool_is_hidden_and_video_tools_offered(self):
|
||||
fake = FakeProvider([_text_chunks("先聊聊")])
|
||||
with patch("apps.ai.creation_agent.build_provider", return_value=fake):
|
||||
list(stream_creation_agent(conversation=self.conversation, user=self.user,
|
||||
text="做条视频", model_config=self.model))
|
||||
names = {t["function"]["name"] for t in fake.calls[0]["extra_body"]["tools"]}
|
||||
self.assertIn("write_strategy", names)
|
||||
self.assertIn("write_plan", names)
|
||||
self.assertNotIn("generate_image", names)
|
||||
|
||||
def test_strategy_card_does_not_stop_the_loop(self):
|
||||
fake = FakeProvider([
|
||||
_tool_chunks("write_strategy", {"target": "油皮通勤人群", "trust": "真实使用反馈",
|
||||
"belief": "值得一试", "direction": "达人 UGC 口播"}),
|
||||
_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))
|
||||
strategy = [e for e in events if e.get("type") == "message" and e["message"]["kind"] == "strategy"]
|
||||
self.assertEqual(strategy[0]["message"]["payload"]["target"], "油皮通勤人群")
|
||||
# 策略卡只是「我理解对了吗」,不该停下来
|
||||
self.assertEqual(len(fake.calls), 2)
|
||||
|
||||
def test_plan_emits_three_cards_and_stops_for_confirmation(self):
|
||||
fake = FakeProvider([
|
||||
_tool_chunks("write_plan", self._plan_args()),
|
||||
_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 = [e["message"]["kind"] for e in events if e.get("type") == "message"]
|
||||
self.assertEqual(kinds[-3:], ["plan", "prompt_file", "confirm"])
|
||||
self.assertTrue(any(e.get("type") == "credits" for e in events))
|
||||
# 「仅需确认一次」—— 必须停下等人点,不能自己往下烧钱出片
|
||||
self.assertEqual(len(fake.calls), 1)
|
||||
|
||||
def test_plan_without_video_prompt_is_rejected_without_emitting_cards(self):
|
||||
fake = FakeProvider([
|
||||
_tool_chunks("write_plan", self._plan_args(video_prompt="")),
|
||||
_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 = [e["message"]["kind"] for e in events if e.get("type") == "message"]
|
||||
self.assertNotIn("confirm", kinds) # 没有出片指令的方案不能放行
|
||||
self.assertEqual(len(fake.calls), 2)
|
||||
|
||||
def test_confirm_submits_video_with_prompt_and_session_params(self):
|
||||
card = append_message(
|
||||
self.conversation, role="assistant", kind=CreationMessage.Kind.CONFIRM,
|
||||
payload={"video_prompt": "0-3秒 近景手持商品…", "submitted": False, "estimated_credits": 120},
|
||||
)
|
||||
task = AITask.objects.create(
|
||||
team=self.team, created_by=self.user, task_type=AITask.Type.FREE_VIDEO,
|
||||
model_config=self.model, idempotency_key="k-video",
|
||||
)
|
||||
with patch("apps.ai.free_video.submit_free_video", return_value=task) as submit:
|
||||
message, error = submit_confirmed_video(
|
||||
conversation=self.conversation, user=self.user, confirm_message=card
|
||||
)
|
||||
params = submit.call_args.kwargs["params"]
|
||||
|
||||
self.assertEqual(error, "")
|
||||
self.assertEqual(message.kind, CreationMessage.Kind.GENERATING)
|
||||
self.assertEqual(params["prompt"], "0-3秒 近景手持商品…")
|
||||
# 顶栏参数直接用,label 要翻成火山真名
|
||||
self.assertEqual(params["model"], "doubao-seedance-2-5-260628")
|
||||
self.assertEqual(params["duration"], 15)
|
||||
self.assertEqual(params["aspect_ratio"], "9:16")
|
||||
self.assertTrue(params["generate_audio"])
|
||||
|
||||
def test_confirm_without_stored_prompt_reports_instead_of_submitting(self):
|
||||
card = append_message(
|
||||
self.conversation, role="assistant", kind=CreationMessage.Kind.CONFIRM,
|
||||
payload={"submitted": False},
|
||||
)
|
||||
with patch("apps.ai.free_video.submit_free_video") as submit:
|
||||
message, error = submit_confirmed_video(
|
||||
conversation=self.conversation, user=self.user, confirm_message=card
|
||||
)
|
||||
self.assertIsNone(message)
|
||||
self.assertIn("出片指令", error)
|
||||
submit.assert_not_called()
|
||||
|
||||
|
||||
class VideoParamParsingTests(TestCase):
|
||||
def test_duration_label_and_smart_fallback(self):
|
||||
self.assertEqual(video_duration({"duration": "15 秒"}), 15)
|
||||
self.assertEqual(video_duration({"duration": "智能时长"}), SMART_DURATION)
|
||||
self.assertEqual(video_duration({}), SMART_DURATION)
|
||||
# 火山单次最长 30 秒,超了要夹住而不是让 submit 报错
|
||||
self.assertEqual(video_duration({"duration": "99 秒"}), 30)
|
||||
|
||||
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)
|
||||
|
||||
|
||||
class ConfirmEndpointTests(TestCase):
|
||||
def setUp(self):
|
||||
self.user = User.objects.create_user(username="confirm-owner", password="p")
|
||||
self.team = Team.objects.create(name="Confirm", owner=self.user)
|
||||
TeamMember.objects.create(team=self.team, user=self.user, role="owner")
|
||||
self.conversation = CreationConversation.objects.create(
|
||||
team=self.team, created_by=self.user, mode="video", params={}
|
||||
)
|
||||
self.card = append_message(
|
||||
self.conversation, role="assistant", kind=CreationMessage.Kind.CONFIRM,
|
||||
payload={"video_prompt": "出片指令", "submitted": False},
|
||||
)
|
||||
self.client = APIClient()
|
||||
self.client.force_authenticate(self.user)
|
||||
|
||||
def _post(self):
|
||||
return self.client.post(
|
||||
f"/api/ai/creations/{self.conversation.id}/send/",
|
||||
{"kind": "confirm", "reply_to": str(self.card.id)}, format="json",
|
||||
)
|
||||
|
||||
def test_confirm_twice_is_refused(self):
|
||||
provider = ModelProvider.objects.create(name="fk2", display_name="F", base_url="https://x")
|
||||
model = ModelConfig.objects.create(
|
||||
provider=provider, name="fk-video", display_name="V",
|
||||
capability=ModelConfig.Capability.VIDEO,
|
||||
)
|
||||
task = AITask.objects.create(
|
||||
team=self.team, created_by=self.user, task_type=AITask.Type.FREE_VIDEO,
|
||||
model_config=model, idempotency_key="k-confirm",
|
||||
)
|
||||
with patch("apps.ai.free_video.submit_free_video", return_value=task):
|
||||
self.assertEqual(self._post().status_code, 201)
|
||||
# 连点两下会出两条片、扣两次积分
|
||||
self.assertEqual(self._post().status_code, 409)
|
||||
|
||||
def test_failed_submit_reopens_the_gate(self):
|
||||
with patch("apps.ai.free_video.submit_free_video", side_effect=ValueError("积分不足")):
|
||||
response = self._post()
|
||||
self.card.refresh_from_db()
|
||||
self.assertEqual(response.status_code, 400)
|
||||
# 出片没提交成功,闸门要放回去让用户改完再确认
|
||||
self.assertFalse(self.card.payload["submitted"])
|
||||
|
||||
|
||||
class MemoryCompressionTests(CreationAgentBaseTests):
|
||||
"""长会话记忆压缩(契约 §5)。"""
|
||||
|
||||
def _fill(self, count):
|
||||
for i in range(count):
|
||||
append_message(self.conversation, role="user" if i % 2 == 0 else "assistant", text=f"第{i}句")
|
||||
|
||||
def test_short_conversation_is_not_compressed(self):
|
||||
self._fill(6)
|
||||
fake = FakeProvider([_text_chunks("好")])
|
||||
with patch("apps.ai.creation_agent.build_provider", return_value=fake):
|
||||
list(stream_creation_agent(conversation=self.conversation, user=self.user,
|
||||
text="继续", model_config=self.model))
|
||||
self.conversation.refresh_from_db()
|
||||
self.assertNotIn("summary", self.conversation.memory)
|
||||
self.assertEqual(len(fake.calls), 1) # 没有多花一次压缩调用
|
||||
|
||||
def test_long_conversation_compresses_and_feeds_summary_into_system_prompt(self):
|
||||
self._fill(30)
|
||||
fake = FakeProvider([_text_chunks("这是摘要正文"), _text_chunks("好")])
|
||||
with patch("apps.ai.creation_agent.build_provider", return_value=fake):
|
||||
list(stream_creation_agent(conversation=self.conversation, user=self.user,
|
||||
text="继续", model_config=self.model))
|
||||
self.conversation.refresh_from_db()
|
||||
self.assertEqual(self.conversation.memory["summary"], "这是摘要正文")
|
||||
# 第二次调用才是真正的对话,system 里要带上刚压出来的摘要
|
||||
system = fake.calls[1]["messages"][0]["content"]
|
||||
self.assertIn("前情提要", system)
|
||||
self.assertIn("这是摘要正文", system)
|
||||
# 且只喂最近 KEEP_RECENT_MESSAGES 条原文,不是全量
|
||||
self.assertLessEqual(len(fake.calls[1]["messages"]), KEEP_RECENT_MESSAGES + 2)
|
||||
|
||||
def test_next_turn_does_not_pay_for_another_compression(self):
|
||||
self._fill(30)
|
||||
fake = FakeProvider([_text_chunks("摘要"), _text_chunks("好"), _text_chunks("好")])
|
||||
with patch("apps.ai.creation_agent.build_provider", return_value=fake):
|
||||
list(stream_creation_agent(conversation=self.conversation, user=self.user,
|
||||
text="一", model_config=self.model))
|
||||
calls_after_first = len(fake.calls)
|
||||
list(stream_creation_agent(conversation=self.conversation, user=self.user,
|
||||
text="二", model_config=self.model))
|
||||
# 第二轮只多了 1 次(对话本身)。不设最小批量的话每轮都要重压,长会话成本翻倍。
|
||||
self.assertEqual(len(fake.calls) - calls_after_first, 1)
|
||||
|
||||
def test_compression_resumes_once_enough_new_messages_pile_up(self):
|
||||
self._fill(30)
|
||||
fake = FakeProvider([_text_chunks("摘要一")] + [_text_chunks("好") for _ in range(20)])
|
||||
with patch("apps.ai.creation_agent.build_provider", return_value=fake):
|
||||
list(stream_creation_agent(conversation=self.conversation, user=self.user,
|
||||
text="一", model_config=self.model))
|
||||
self.conversation.refresh_from_db()
|
||||
first_upto = self.conversation.memory["summarized_upto"]
|
||||
self._fill(COMPRESS_MIN_BATCH * 2)
|
||||
list(stream_creation_agent(conversation=self.conversation, user=self.user,
|
||||
text="二", model_config=self.model))
|
||||
self.conversation.refresh_from_db()
|
||||
# 攒够一批之后要接着压,否则早期内容永远进不了摘要
|
||||
self.assertGreater(self.conversation.memory["summarized_upto"], first_upto)
|
||||
|
||||
def test_compression_failure_does_not_break_the_conversation(self):
|
||||
self._fill(30)
|
||||
|
||||
class Flaky(FakeProvider):
|
||||
def chat_completion_stream(self, **kwargs):
|
||||
if len(self.calls) == 0:
|
||||
self.calls.append(kwargs)
|
||||
raise RuntimeError("summary model down")
|
||||
return super().chat_completion_stream(**kwargs)
|
||||
|
||||
fake = Flaky([_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))
|
||||
# 摘要是锦上添花,压缩挂了不该把整条对话打断
|
||||
self.assertEqual(events[-1]["type"], "done")
|
||||
self.conversation.refresh_from_db()
|
||||
self.assertNotIn("summary", self.conversation.memory)
|
||||
|
||||
|
||||
class PresetGuidanceTests(CreationAgentBaseTests):
|
||||
"""预设不只是个名字,要把拍法约束一起给模型(契约 §6)。"""
|
||||
|
||||
def test_preset_guidance_reaches_the_system_prompt(self):
|
||||
conversation = CreationConversation.objects.create(
|
||||
team=self.team, created_by=self.user, mode="video", preset="鱼眼换装", params={},
|
||||
)
|
||||
fake = FakeProvider([_text_chunks("好")])
|
||||
with patch("apps.ai.creation_agent.build_provider", return_value=fake):
|
||||
list(stream_creation_agent(conversation=conversation, user=self.user,
|
||||
text="开始", model_config=self.model))
|
||||
system = fake.calls[0]["messages"][0]["content"]
|
||||
self.assertIn("鱼眼换装", system)
|
||||
# 光有名字模型只能靠猜,拍法约束必须一起给
|
||||
self.assertIn("人物面部和身形必须全程一致", system)
|
||||
|
||||
def test_unknown_preset_degrades_to_name_only(self):
|
||||
conversation = CreationConversation.objects.create(
|
||||
team=self.team, created_by=self.user, mode="video", preset="前端新加的卡", params={},
|
||||
)
|
||||
fake = FakeProvider([_text_chunks("好")])
|
||||
with patch("apps.ai.creation_agent.build_provider", return_value=fake):
|
||||
list(stream_creation_agent(conversation=conversation, user=self.user,
|
||||
text="开始", model_config=self.model))
|
||||
system = fake.calls[0]["messages"][0]["content"]
|
||||
# 前端加了新卡但后端还没写拍法时,退回「只有名字」而不是报错
|
||||
self.assertIn("前端新加的卡", system)
|
||||
|
||||
def test_every_frontend_preset_has_guidance(self):
|
||||
"""前端 8 个视频 + 6 个图片预设都要有拍法,漏一个就等于那张卡是摆设。"""
|
||||
from .creation_presets import IMAGE_PRESETS, VIDEO_PRESETS
|
||||
|
||||
self.assertEqual(len(VIDEO_PRESETS), 8)
|
||||
self.assertEqual(len(IMAGE_PRESETS), 6)
|
||||
self.assertTrue(all(text.strip() for text in {**VIDEO_PRESETS, **IMAGE_PRESETS}.values()))
|
||||
@@ -0,0 +1,238 @@
|
||||
"""全能创作 · 会话与消息底座(契约 §1/§3)。"""
|
||||
from django.test import TestCase
|
||||
from rest_framework.test import APIClient
|
||||
|
||||
from apps.accounts.models import Team, TeamMember, User
|
||||
|
||||
from apps.assets.models import Asset, AssetFile
|
||||
|
||||
from .creation import append_message, finish_generating_message, pin_refs, sync_generating_messages
|
||||
from .models import AITask, CreationConversation, CreationMessage, ModelConfig, ModelProvider
|
||||
|
||||
|
||||
class CreationMessageServiceTests(TestCase):
|
||||
def setUp(self):
|
||||
self.user = User.objects.create_user(username="omni-svc", password="p")
|
||||
self.team = Team.objects.create(name="Omni SVC", owner=self.user)
|
||||
self.conversation = CreationConversation.objects.create(
|
||||
team=self.team, created_by=self.user, title="净颜精华口播", mode="video"
|
||||
)
|
||||
|
||||
def test_seq_is_monotonic_per_conversation(self):
|
||||
first = append_message(self.conversation, role="user", text="做一条口播")
|
||||
second = append_message(self.conversation, role="assistant", text="好的")
|
||||
other = CreationConversation.objects.create(team=self.team, created_by=self.user, mode="image")
|
||||
other_first = append_message(other, role="user", text="来张主图")
|
||||
|
||||
self.assertEqual([first.seq, second.seq], [1, 2])
|
||||
# seq 是会话内自增,不是全局 —— 换一条会话要从 1 重新开始
|
||||
self.assertEqual(other_first.seq, 1)
|
||||
|
||||
def test_append_refreshes_last_active_at(self):
|
||||
before = self.conversation.last_active_at
|
||||
append_message(self.conversation, role="user", text="改一下背景")
|
||||
self.conversation.refresh_from_db()
|
||||
self.assertGreater(self.conversation.last_active_at, before)
|
||||
|
||||
def test_generating_message_is_replaced_in_place_not_appended(self):
|
||||
placeholder = append_message(
|
||||
self.conversation,
|
||||
role="assistant",
|
||||
kind=CreationMessage.Kind.GENERATING,
|
||||
payload={"task_id": "t-1", "kind": "video"},
|
||||
)
|
||||
finish_generating_message(
|
||||
placeholder,
|
||||
assets=[{"id": "a-1", "url": "https://x/v.mp4", "cover": "https://x/c.jpg", "type": "video"}],
|
||||
meta={"model": "Seedance 2.5", "resolution": "1080p", "ratio": "9:16"},
|
||||
)
|
||||
placeholder.refresh_from_db()
|
||||
self.conversation.refresh_from_db()
|
||||
|
||||
self.assertEqual(placeholder.kind, CreationMessage.Kind.RESULT)
|
||||
self.assertEqual(placeholder.payload["assets"][0]["id"], "a-1")
|
||||
self.assertEqual(placeholder.payload["task_id"], "t-1") # 原 payload 不能被覆盖掉
|
||||
self.assertEqual(self.conversation.messages.count(), 1) # 中间态不刷屏
|
||||
self.assertEqual(self.conversation.status, CreationConversation.Status.COMPLETED)
|
||||
|
||||
def test_pin_refs_dedupes_and_keeps_first_seen_order(self):
|
||||
pin_refs(self.conversation, [
|
||||
{"type": "product", "id": "p1", "name": "净颜精华"},
|
||||
{"type": "character", "id": "c1", "name": "白领女性"},
|
||||
])
|
||||
pin_refs(self.conversation, [
|
||||
{"type": "product", "id": "p1", "name": "净颜精华"}, # 重复,不入
|
||||
{"type": "scene", "id": "s1", "name": "居家早餐台"},
|
||||
{"type": "scene"}, # 缺 id,丢弃
|
||||
])
|
||||
self.conversation.refresh_from_db()
|
||||
|
||||
self.assertEqual(
|
||||
[(r["type"], r["id"]) for r in self.conversation.pinned_refs],
|
||||
[("product", "p1"), ("character", "c1"), ("scene", "s1")],
|
||||
)
|
||||
|
||||
|
||||
class GenerationBackfillTests(TestCase):
|
||||
"""出图在 worker 里跑完后,GENERATING 必须被回填,否则对话会一直转圈。"""
|
||||
|
||||
def setUp(self):
|
||||
self.user = User.objects.create_user(username="omni-backfill", password="p")
|
||||
self.team = Team.objects.create(name="Omni Backfill", owner=self.user)
|
||||
TeamMember.objects.create(team=self.team, user=self.user, role="owner")
|
||||
provider = ModelProvider.objects.create(name="img", display_name="Img")
|
||||
self.model = ModelConfig.objects.create(
|
||||
provider=provider, name="gpt-image-2", display_name="YQ image2",
|
||||
capability=ModelConfig.Capability.IMAGE,
|
||||
)
|
||||
self.conversation = CreationConversation.objects.create(
|
||||
team=self.team, created_by=self.user, title="回填", mode="image"
|
||||
)
|
||||
self.client = APIClient()
|
||||
self.client.force_authenticate(self.user)
|
||||
|
||||
def _task(self, status, key="k-backfill"):
|
||||
return AITask.objects.create(
|
||||
team=self.team, created_by=self.user, task_type=AITask.Type.PRODUCT_IMAGE,
|
||||
model_config=self.model, status=status, idempotency_key=key,
|
||||
error_message="额度不足" if status == AITask.Status.FAILED else "",
|
||||
)
|
||||
|
||||
def _asset(self, task, url="https://cdn.example/done.png"):
|
||||
asset = Asset.objects.create(
|
||||
team=self.team, created_by=self.user, name="成图",
|
||||
asset_type=Asset.Type.IMAGE, source=Asset.Source.AI_GENERATED,
|
||||
category=Asset.Category.PRODUCT_IMAGE, origin_task=task,
|
||||
)
|
||||
AssetFile.objects.create(
|
||||
asset=asset, object_key="k", bucket="b", is_primary=True, preview_url=url,
|
||||
)
|
||||
return asset
|
||||
|
||||
def test_succeeded_task_turns_generating_into_result(self):
|
||||
task = self._task(AITask.Status.SUCCEEDED)
|
||||
self._asset(task)
|
||||
message = append_message(
|
||||
self.conversation, role="assistant", kind=CreationMessage.Kind.GENERATING,
|
||||
payload={"task_id": str(task.id), "kind": "image", "prompt": "白底"},
|
||||
task=task,
|
||||
)
|
||||
|
||||
self.assertEqual(sync_generating_messages(self.conversation), 1)
|
||||
message.refresh_from_db()
|
||||
self.conversation.refresh_from_db()
|
||||
self.assertEqual(message.kind, CreationMessage.Kind.RESULT)
|
||||
self.assertEqual(message.payload["assets"][0]["url"], "https://cdn.example/done.png")
|
||||
self.assertEqual(message.payload["task_id"], str(task.id))
|
||||
self.assertEqual(self.conversation.status, CreationConversation.Status.COMPLETED)
|
||||
|
||||
def test_failed_task_turns_generating_into_error(self):
|
||||
task = self._task(AITask.Status.FAILED, key="k-fail")
|
||||
message = append_message(
|
||||
self.conversation, role="assistant", kind=CreationMessage.Kind.GENERATING,
|
||||
payload={"task_id": str(task.id), "kind": "image"},
|
||||
task=task,
|
||||
)
|
||||
|
||||
self.assertEqual(sync_generating_messages(self.conversation), 1)
|
||||
message.refresh_from_db()
|
||||
self.assertEqual(message.kind, CreationMessage.Kind.ERROR)
|
||||
self.assertIn("额度不足", message.text)
|
||||
|
||||
def test_running_task_stays_generating(self):
|
||||
task = self._task(AITask.Status.RESERVED, key="k-run")
|
||||
message = append_message(
|
||||
self.conversation, role="assistant", kind=CreationMessage.Kind.GENERATING,
|
||||
payload={"task_id": str(task.id), "kind": "image"},
|
||||
task=task,
|
||||
)
|
||||
|
||||
self.assertEqual(sync_generating_messages(self.conversation), 0)
|
||||
message.refresh_from_db()
|
||||
self.assertEqual(message.kind, CreationMessage.Kind.GENERATING)
|
||||
|
||||
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")
|
||||
append_message(
|
||||
self.conversation, role="assistant", kind=CreationMessage.Kind.GENERATING,
|
||||
payload={"task_id": str(task.id), "kind": "image"},
|
||||
task=task,
|
||||
)
|
||||
|
||||
detail = self.client.get(f"/api/ai/creations/{self.conversation.id}/")
|
||||
self.assertEqual(detail.status_code, 200)
|
||||
self.assertEqual(detail.data["messages"][0]["kind"], "result")
|
||||
self.assertEqual(detail.data["messages"][0]["payload"]["assets"][0]["url"], "https://cdn.example/api.png")
|
||||
|
||||
|
||||
class CreationConversationAPITests(TestCase):
|
||||
def setUp(self):
|
||||
self.user = User.objects.create_user(username="omni-api", password="p")
|
||||
self.team = Team.objects.create(name="Omni API", owner=self.user)
|
||||
TeamMember.objects.create(team=self.team, user=self.user, role="owner")
|
||||
self.client = APIClient()
|
||||
self.client.force_authenticate(self.user)
|
||||
|
||||
def test_create_list_and_rename(self):
|
||||
created = self.client.post(
|
||||
"/api/ai/creations/",
|
||||
{"title": "净颜精华口播", "mode": "video", "preset": "达人口播种草",
|
||||
"params": {"model": "Seedance 2.5", "ratio": "9:16", "resolution": "1080p", "duration": "智能时长"}},
|
||||
format="json",
|
||||
)
|
||||
self.assertEqual(created.status_code, 201, created.data)
|
||||
conv_id = created.data["id"]
|
||||
|
||||
listed = self.client.get("/api/ai/creations/?mode=video")
|
||||
self.assertEqual(listed.status_code, 200)
|
||||
self.assertEqual(len(listed.data["results"] if "results" in listed.data else listed.data), 1)
|
||||
|
||||
renamed = self.client.patch(f"/api/ai/creations/{conv_id}/", {"title": "改个名", "mode": "image"}, format="json")
|
||||
self.assertEqual(renamed.status_code, 200)
|
||||
self.assertEqual(renamed.data["title"], "改个名")
|
||||
# mode 定死:传了也不生效,否则顶栏参数会和已生成内容对不上
|
||||
self.assertEqual(renamed.data["mode"], "video")
|
||||
|
||||
def test_retrieve_returns_full_thread_and_cover_is_latest_result(self):
|
||||
conversation = CreationConversation.objects.create(
|
||||
team=self.team, created_by=self.user, title="叠加测试", mode="image"
|
||||
)
|
||||
append_message(conversation, role="user", text="来一张")
|
||||
append_message(conversation, role="assistant", kind=CreationMessage.Kind.RESULT,
|
||||
payload={"assets": [{"id": "a1", "url": "u1", "cover": "c1"}]})
|
||||
append_message(conversation, role="user", text="背景换夜景")
|
||||
append_message(conversation, role="assistant", kind=CreationMessage.Kind.RESULT,
|
||||
payload={"assets": [{"id": "a2", "url": "u2", "cover": "c2"}]})
|
||||
|
||||
detail = self.client.get(f"/api/ai/creations/{conversation.id}/")
|
||||
self.assertEqual(detail.status_code, 200)
|
||||
self.assertEqual(len(detail.data["messages"]), 4)
|
||||
# 重生成往下叠加,旧的留着;封面取最新一版
|
||||
self.assertEqual(detail.data["cover_url"], "c2")
|
||||
|
||||
def test_messages_endpoint_supports_incremental_pull(self):
|
||||
conversation = CreationConversation.objects.create(team=self.team, created_by=self.user, mode="video")
|
||||
append_message(conversation, role="user", text="一")
|
||||
append_message(conversation, role="assistant", text="二")
|
||||
|
||||
incremental = self.client.get(f"/api/ai/creations/{conversation.id}/messages/?after_seq=1")
|
||||
self.assertEqual(incremental.status_code, 200)
|
||||
self.assertEqual([m["text"] for m in incremental.data], ["二"])
|
||||
|
||||
def test_other_team_cannot_read_conversation(self):
|
||||
conversation = CreationConversation.objects.create(team=self.team, created_by=self.user, mode="video")
|
||||
stranger = User.objects.create_user(username="omni-stranger", password="p")
|
||||
other_team = Team.objects.create(name="Other", owner=stranger)
|
||||
TeamMember.objects.create(team=other_team, user=stranger, role="owner")
|
||||
other_client = APIClient()
|
||||
other_client.force_authenticate(stranger)
|
||||
|
||||
self.assertEqual(other_client.get(f"/api/ai/creations/{conversation.id}/").status_code, 404)
|
||||
|
||||
def test_destroy_is_soft_delete(self):
|
||||
conversation = CreationConversation.objects.create(team=self.team, created_by=self.user, mode="video")
|
||||
self.assertEqual(self.client.delete(f"/api/ai/creations/{conversation.id}/").status_code, 204)
|
||||
conversation.refresh_from_db()
|
||||
self.assertTrue(conversation.is_deleted)
|
||||
self.assertEqual(self.client.get(f"/api/ai/creations/{conversation.id}/").status_code, 404)
|
||||
@@ -0,0 +1,179 @@
|
||||
"""全能创作 · @引用检索与 Ref 解析(契约 §1/§3)。"""
|
||||
from unittest.mock import patch
|
||||
|
||||
from django.test import TestCase
|
||||
from rest_framework.test import APIClient
|
||||
|
||||
from apps.accounts.models import Team, TeamMember, User
|
||||
from apps.assets.models import Asset, AssetFile, Model
|
||||
from apps.products.models import Product, ProductImage, ProductSellingPoint
|
||||
|
||||
from .mentions import product_facts_text, resolve_refs, search_mentions
|
||||
|
||||
|
||||
def _image_asset(team, user, name, category, *, source=Asset.Source.UPLOAD, url="", **kwargs):
|
||||
asset = Asset.objects.create(
|
||||
team=team, created_by=user, name=name, asset_type=Asset.Type.IMAGE,
|
||||
source=source, category=category, **kwargs,
|
||||
)
|
||||
AssetFile.objects.create(
|
||||
asset=asset, object_key=f"k/{name}", bucket="b", is_primary=True,
|
||||
preview_url=url or f"https://cdn/{name}.jpg",
|
||||
)
|
||||
return asset
|
||||
|
||||
|
||||
class MentionSearchTests(TestCase):
|
||||
def setUp(self):
|
||||
self.user = User.objects.create_user(username="mention-owner", password="p")
|
||||
self.team = Team.objects.create(name="Mention", owner=self.user)
|
||||
self.product = Product.objects.create(team=self.team, created_by=self.user, title="净颜精华")
|
||||
self.person = _image_asset(self.team, self.user, "白领女性", Asset.Category.PERSON)
|
||||
self.scene = _image_asset(self.team, self.user, "居家早餐台", Asset.Category.SCENE)
|
||||
|
||||
def test_search_filters_by_keyword_and_type(self):
|
||||
Product.objects.create(team=self.team, created_by=self.user, title="控油洁面")
|
||||
|
||||
hits = search_mentions(self.team, q="净颜", types=["product"])
|
||||
self.assertEqual([h["name"] for h in hits], ["净颜精华"])
|
||||
self.assertEqual(hits[0]["type"], "product")
|
||||
|
||||
def test_search_returns_all_types_when_unspecified(self):
|
||||
found = {(h["type"], h["name"]) for h in search_mentions(self.team, q="")}
|
||||
self.assertIn(("product", "净颜精华"), found)
|
||||
self.assertIn(("character", "白领女性"), found)
|
||||
self.assertIn(("scene", "居家早餐台"), found)
|
||||
|
||||
def test_asset_type_only_lists_items_added_to_library(self):
|
||||
_image_asset(self.team, self.user, "工作台试验图", Asset.Category.FREE_CREATE, in_library=False)
|
||||
_image_asset(self.team, self.user, "入库图", Asset.Category.FREE_CREATE)
|
||||
names = [h["name"] for h in search_mentions(self.team, types=["asset"])]
|
||||
self.assertNotIn("工作台试验图", names) # 工作台的试验图不该冒进 @ 菜单
|
||||
self.assertIn("入库图", names)
|
||||
|
||||
def test_asset_group_does_not_duplicate_entries_with_their_own_group(self):
|
||||
found = [(h["type"], h["name"]) for h in search_mentions(self.team)]
|
||||
# 定妆照只该出现在「角色」里;再在「资产库」列一遍,菜单里看着像两个素材
|
||||
self.assertEqual(found.count(("character", "白领女性")), 1)
|
||||
self.assertNotIn(("asset", "白领女性"), found)
|
||||
self.assertNotIn(("asset", "居家早餐台"), found)
|
||||
|
||||
def test_other_team_entities_are_invisible(self):
|
||||
stranger = User.objects.create_user(username="mention-stranger", password="p")
|
||||
other_team = Team.objects.create(name="Other", owner=stranger)
|
||||
Product.objects.create(team=other_team, created_by=stranger, title="别家的商品")
|
||||
|
||||
names = [h["name"] for h in search_mentions(self.team)]
|
||||
self.assertNotIn("别家的商品", names)
|
||||
|
||||
|
||||
class ResolveRefsTests(TestCase):
|
||||
def setUp(self):
|
||||
self.user = User.objects.create_user(username="resolve-owner", password="p")
|
||||
self.team = Team.objects.create(name="Resolve", owner=self.user)
|
||||
self.product = Product.objects.create(
|
||||
team=self.team, created_by=self.user, title="净颜精华", brand="影擎",
|
||||
category="护肤", description="早晚各一次", specs={"容量": "30ml"},
|
||||
)
|
||||
ProductSellingPoint.objects.create(product=self.product, title="控油", detail="12 小时不脱妆")
|
||||
product_asset = _image_asset(
|
||||
self.team, self.user, "商品实拍", Asset.Category.PRODUCT_IMAGE, url="https://cdn/prod.jpg"
|
||||
)
|
||||
ProductImage.objects.create(product=self.product, asset=product_asset, is_primary=True)
|
||||
self.person = _image_asset(
|
||||
self.team, self.user, "白领女性", Asset.Category.PERSON,
|
||||
url="https://cdn/person.jpg", review_status="active", review_remote_id="R-1",
|
||||
)
|
||||
self.scene = _image_asset(
|
||||
self.team, self.user, "居家早餐台", Asset.Category.SCENE, url="https://cdn/scene.jpg"
|
||||
)
|
||||
|
||||
def test_product_ref_yields_selling_points_and_reference_image(self):
|
||||
resolved = resolve_refs(self.team, [{"type": "product", "id": str(self.product.id)}])
|
||||
|
||||
self.assertIn("控油", resolved.facts_text)
|
||||
self.assertIn("12 小时不脱妆", resolved.facts_text)
|
||||
self.assertIn("30ml", resolved.facts_text)
|
||||
self.assertEqual([r["url"] for r in resolved.references], ["https://cdn/prod.jpg"])
|
||||
|
||||
def test_reference_order_is_character_then_scene_then_product(self):
|
||||
resolved = resolve_refs(self.team, [
|
||||
{"type": "product", "id": str(self.product.id)},
|
||||
{"type": "scene", "id": str(self.scene.id)},
|
||||
{"type": "character", "id": str(self.person.id)},
|
||||
])
|
||||
# 顺序是 @图N 的语义依据:角色 → 场景 → 商品,不能跟着用户 @ 的先后走
|
||||
self.assertEqual([r["type"] for r in resolved.references], ["character", "scene", "product"])
|
||||
|
||||
def test_character_ref_carries_review_status_for_asset_scheme_swap(self):
|
||||
resolved = resolve_refs(self.team, [{"type": "character", "id": str(self.person.id)}])
|
||||
entry = resolved.references[0]
|
||||
# 视频路要靠这两个字段把真人图换成火山 asset:// 引用,否则会被判「疑似真人」拒
|
||||
self.assertEqual(entry["review_status"], "active")
|
||||
self.assertEqual(entry["review_remote_id"], "R-1")
|
||||
|
||||
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(
|
||||
team=self.team, created_by=self.user, name="小夏",
|
||||
portrait_asset=portrait, triview_asset=triview,
|
||||
)
|
||||
|
||||
resolved = resolve_refs(self.team, [{"type": "model", "id": str(model.id)}])
|
||||
# 三视图信息量最大,锁脸优先用它
|
||||
self.assertEqual(resolved.references[0]["url"], "https://cdn/t.jpg")
|
||||
|
||||
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")
|
||||
model = Model.objects.create(team=self.team, created_by=self.user, name="阿岚", portrait_asset=portrait)
|
||||
|
||||
resolved = resolve_refs(self.team, [{"type": "model", "id": str(model.id)}])
|
||||
self.assertEqual(resolved.references[0]["url"], "https://cdn/p2.jpg")
|
||||
|
||||
def test_deleted_or_foreign_refs_go_to_missing_and_never_raise(self):
|
||||
stranger = User.objects.create_user(username="resolve-stranger", password="p")
|
||||
other_team = Team.objects.create(name="Other", owner=stranger)
|
||||
foreign = Product.objects.create(team=other_team, created_by=stranger, title="别家的")
|
||||
|
||||
resolved = resolve_refs(self.team, [
|
||||
{"type": "product", "id": str(foreign.id)},
|
||||
{"type": "character", "id": "00000000-0000-0000-0000-000000000000"},
|
||||
{"type": "product", "id": str(self.product.id)},
|
||||
])
|
||||
# 素材被删/跨团队不该炸掉整条对话,交给 agent 在对话里说明
|
||||
self.assertEqual(len(resolved.missing), 2)
|
||||
self.assertEqual(len(resolved.references), 1)
|
||||
|
||||
def test_same_image_referenced_twice_is_deduped(self):
|
||||
resolved = resolve_refs(self.team, [
|
||||
{"type": "character", "id": str(self.person.id)},
|
||||
{"type": "asset", "id": str(self.person.id)},
|
||||
])
|
||||
# 编号错位会让 @图N 指错,必须去重
|
||||
self.assertEqual(len(resolved.references), 1)
|
||||
|
||||
def test_product_facts_text_without_selling_points_still_has_title(self):
|
||||
bare = Product.objects.create(team=self.team, created_by=self.user, title="裸商品")
|
||||
self.assertIn("裸商品", product_facts_text(bare))
|
||||
|
||||
|
||||
class MentionAPITests(TestCase):
|
||||
def setUp(self):
|
||||
self.user = User.objects.create_user(username="mention-api", password="p")
|
||||
self.team = Team.objects.create(name="Mention API", owner=self.user)
|
||||
TeamMember.objects.create(team=self.team, user=self.user, role="owner")
|
||||
Product.objects.create(team=self.team, created_by=self.user, title="净颜精华")
|
||||
self.client = APIClient()
|
||||
self.client.force_authenticate(self.user)
|
||||
|
||||
def test_endpoint_returns_refs_with_group_labels(self):
|
||||
response = self.client.get("/api/ai/mentions/?q=净颜&types=product")
|
||||
self.assertEqual(response.status_code, 200)
|
||||
self.assertEqual(response.data["results"][0]["name"], "净颜精华")
|
||||
self.assertEqual(response.data["type_labels"]["product"], "商品库")
|
||||
|
||||
def test_unknown_type_is_rejected(self):
|
||||
response = self.client.get("/api/ai/mentions/?types=product,ghost")
|
||||
self.assertEqual(response.status_code, 400)
|
||||
self.assertIn("ghost", response.data["detail"])
|
||||
@@ -3,6 +3,7 @@ from rest_framework.routers import DefaultRouter
|
||||
|
||||
from .views import (
|
||||
AITaskViewSet,
|
||||
CreationConversationViewSet,
|
||||
VideoDigestDetailView,
|
||||
VideoDigestView,
|
||||
FreeVideoDetailView,
|
||||
@@ -15,6 +16,7 @@ from .views import (
|
||||
FreeVideoView,
|
||||
GenerateImageView,
|
||||
ImageConversationViewSet,
|
||||
MentionSearchView,
|
||||
ModelConfigViewSet,
|
||||
VideoReplacePollView,
|
||||
VideoReplaceView,
|
||||
@@ -24,8 +26,10 @@ router = DefaultRouter()
|
||||
router.register("tasks", AITaskViewSet, basename="ai-task")
|
||||
router.register("models", ModelConfigViewSet, basename="model-config")
|
||||
router.register("image-conversations", ImageConversationViewSet, basename="image-conversation")
|
||||
router.register("creations", CreationConversationViewSet, basename="creation-conversation")
|
||||
|
||||
urlpatterns = [
|
||||
path("mentions/", MentionSearchView.as_view(), name="ai-mentions"),
|
||||
path("generate-image/", GenerateImageView.as_view(), name="ai-generate-image"),
|
||||
path("video-digest/", VideoDigestView.as_view(), name="ai-video-digest"),
|
||||
path("video-digest/<uuid:task_id>/", VideoDigestDetailView.as_view(), name="ai-video-digest-detail"),
|
||||
|
||||
@@ -2,6 +2,7 @@ import logging
|
||||
import uuid
|
||||
|
||||
from django.db import transaction
|
||||
from django.http import JsonResponse, StreamingHttpResponse
|
||||
from django.db.models import Count, Exists, OuterRef, Q
|
||||
from django.utils import timezone
|
||||
from rest_framework import status
|
||||
@@ -14,14 +15,20 @@ from rest_framework.viewsets import ModelViewSet, ReadOnlyModelViewSet
|
||||
|
||||
from apps.assets.models import Asset
|
||||
from apps.assets.serializers import AssetFileSerializer, AssetSerializer
|
||||
from apps.common.api import TeamScopedViewSetMixin, get_current_team
|
||||
from apps.common.api import ServerSentEventRenderer, TeamScopedViewSetMixin, get_current_team
|
||||
from apps.common.celery_health import require_worker, require_worker_task
|
||||
from apps.products.models import Product
|
||||
|
||||
from .generation_errors import classify_generation_error, public_error_for_task
|
||||
from .models import AITask, ImageConversation, ModelConfig
|
||||
from .creation import append_message, sync_generating_messages
|
||||
from .creation_agent import apply_session_params, stream_creation_agent, submit_confirmed_video
|
||||
from .mentions import TYPE_LABELS, VALID_TYPES, refs_from_elicit_answers, search_mentions
|
||||
from .models import AITask, CreationConversation, CreationMessage, ImageConversation, ModelConfig
|
||||
from .serializers import (
|
||||
AITaskSerializer,
|
||||
CreationConversationDetailSerializer,
|
||||
CreationConversationSerializer,
|
||||
CreationMessageSerializer,
|
||||
ImageConversationSerializer,
|
||||
ImageConversationTrashSerializer,
|
||||
ModelConfigSerializer,
|
||||
@@ -1244,3 +1251,218 @@ class ModelConfigViewSet(ReadOnlyModelViewSet):
|
||||
search_fields = ["name", "display_name", "capability"]
|
||||
ordering_fields = ["created_at", "display_name"]
|
||||
|
||||
|
||||
|
||||
class CreationConversationViewSet(TeamScopedViewSetMixin, ModelViewSet):
|
||||
"""全能创作会话 CRUD(契约 §3)。
|
||||
|
||||
list 创作历史页,按 ?status=running|completed 过滤,-last_active_at 倒序
|
||||
retrieve 进对话页,一次性带回全部消息
|
||||
create 从首页「开始创作」发起,带 mode/preset/params;首条用户消息由 messages 接口发
|
||||
partial_update 只允许改 title(重命名)
|
||||
destroy 软删(不连带删已生成的资产 —— 图还在资产库里)
|
||||
|
||||
发消息走 POST {id}/messages/(SSE),不在这里。
|
||||
"""
|
||||
|
||||
serializer_class = CreationConversationSerializer
|
||||
queryset = CreationConversation.objects.order_by("-last_active_at")
|
||||
|
||||
def get_serializer_class(self):
|
||||
if self.action == "retrieve":
|
||||
return CreationConversationDetailSerializer
|
||||
return super().get_serializer_class()
|
||||
|
||||
def get_queryset(self):
|
||||
queryset = super().get_queryset().filter(is_deleted=False, purged_at__isnull=True)
|
||||
if self.action == "retrieve":
|
||||
queryset = queryset.prefetch_related("messages")
|
||||
else:
|
||||
queryset = queryset.annotate(_message_count=Count("messages"))
|
||||
mode = self.request.query_params.get("mode", "").strip()
|
||||
if mode:
|
||||
if mode not in CreationConversation.Mode.values:
|
||||
raise ValidationError({"detail": "mode 仅支持 video / image"})
|
||||
queryset = queryset.filter(mode=mode)
|
||||
conv_status = self.request.query_params.get("status", "").strip()
|
||||
if conv_status:
|
||||
if conv_status not in CreationConversation.Status.values:
|
||||
raise ValidationError({"detail": "status 仅支持 running / completed / failed"})
|
||||
queryset = queryset.filter(status=conv_status)
|
||||
return queryset
|
||||
|
||||
def perform_destroy(self, instance):
|
||||
# 只软删会话本身。已生成的图/视频留在资产库 —— 用户删对话不等于要删素材。
|
||||
instance.is_deleted = True
|
||||
instance.save(update_fields=["is_deleted", "updated_at"])
|
||||
|
||||
def _synced_conversation(self):
|
||||
"""拉会话前先把已结束的 GENERATING 回填成 RESULT/ERROR。
|
||||
|
||||
出图在 worker 里跑,agent 只提交。前端轮询 GET 本接口拿结果。
|
||||
prefetch 缓存里是回填前的旧对象,改过必须丢掉再读。
|
||||
"""
|
||||
conversation = self.get_object()
|
||||
if sync_generating_messages(conversation):
|
||||
conversation.refresh_from_db()
|
||||
cache = getattr(conversation, "_prefetched_objects_cache", None)
|
||||
if cache is not None:
|
||||
cache.pop("messages", None)
|
||||
return conversation
|
||||
|
||||
def retrieve(self, request, *args, **kwargs):
|
||||
conversation = self._synced_conversation()
|
||||
serializer = self.get_serializer(conversation)
|
||||
return Response(serializer.data)
|
||||
|
||||
@action(detail=True, methods=["get"], url_path="messages")
|
||||
def messages(self, request, pk=None):
|
||||
"""按 ?after_seq= 增量拉消息。轮询视频结果时前端只补新的,不重拉整条会话。"""
|
||||
conversation = self._synced_conversation()
|
||||
queryset = conversation.messages.all()
|
||||
after_seq = request.query_params.get("after_seq", "").strip()
|
||||
if after_seq:
|
||||
try:
|
||||
queryset = queryset.filter(seq__gt=int(after_seq))
|
||||
except (TypeError, ValueError) as exc:
|
||||
raise ValidationError({"detail": "after_seq 必须是整数"}) from exc
|
||||
return Response(CreationMessageSerializer(queryset, many=True).data)
|
||||
|
||||
@action(
|
||||
detail=True, methods=["post"], url_path="send",
|
||||
renderer_classes=[ServerSentEventRenderer],
|
||||
)
|
||||
def send(self, request, pk=None):
|
||||
"""发一条消息 → SSE 流(契约 §3)。
|
||||
|
||||
kind=text 普通发言,text + refs
|
||||
kind=elicit_answer 回答追问卡,reply_to + answers
|
||||
kind=confirm 点确认闸门 → **不跑模型**,直接按方案卡存的 video_prompt 出片
|
||||
|
||||
响应 text/event-stream。**必须挂 ServerSentEventRenderer,否则 DRF 内容协商直接 406。**
|
||||
"""
|
||||
conversation = self.get_object()
|
||||
# 模型/比例/分辨率/时长在新建会话时锁定,发送和确认出片都按当时那套,
|
||||
# 否则 5 秒方案被改成 10 秒再出片会对不上。
|
||||
kind = str(request.data.get("kind") or "text")
|
||||
text = str(request.data.get("text") or "").strip()
|
||||
refs = request.data.get("refs") or []
|
||||
if not isinstance(refs, list):
|
||||
return JsonResponse({"detail": "refs 必须是数组"}, status=400)
|
||||
|
||||
if kind == "confirm":
|
||||
reply_to = str(request.data.get("reply_to") or "").strip()
|
||||
card = conversation.messages.filter(
|
||||
id=reply_to, kind=CreationMessage.Kind.CONFIRM
|
||||
).first() if reply_to else None
|
||||
if card is None:
|
||||
return JsonResponse({"detail": "确认卡不存在"}, status=404)
|
||||
if (card.payload or {}).get("submitted"):
|
||||
# 确认闸门是一次性的:连点两下会出两条片、扣两次积分
|
||||
return JsonResponse({"detail": "这条方案已经确认过了"}, status=409)
|
||||
card.payload = {**(card.payload or {}), "submitted": True}
|
||||
card.save(update_fields=["payload", "updated_at"])
|
||||
message, error = submit_confirmed_video(
|
||||
conversation=conversation, user=request.user, confirm_message=card
|
||||
)
|
||||
if error:
|
||||
# 出片没提交成功 → 把闸门放回去,用户可以改完再确认
|
||||
card.payload = {**(card.payload or {}), "submitted": False}
|
||||
card.save(update_fields=["payload", "updated_at"])
|
||||
failure = append_message(
|
||||
conversation, role="assistant",
|
||||
kind=CreationMessage.Kind.ERROR, text=error,
|
||||
)
|
||||
return JsonResponse(
|
||||
{"detail": error, "message": CreationMessageSerializer(failure).data}, status=400
|
||||
)
|
||||
# 纯 Django 响应:这个 action 只挂了 SSE renderer,走 DRF Response 会渲染失败
|
||||
return JsonResponse(
|
||||
{"message": CreationMessageSerializer(message).data}, status=201
|
||||
)
|
||||
|
||||
if kind == "elicit_answer":
|
||||
reply_to = str(request.data.get("reply_to") or "").strip()
|
||||
answers = request.data.get("answers")
|
||||
if not reply_to or not isinstance(answers, dict):
|
||||
return JsonResponse({"detail": "回答追问需要 reply_to 与 answers"}, status=400)
|
||||
card = conversation.messages.filter(
|
||||
id=reply_to, kind=CreationMessage.Kind.ELICIT
|
||||
).first()
|
||||
if card is None:
|
||||
return JsonResponse({"detail": "追问卡不存在"}, status=404)
|
||||
if (card.payload or {}).get("submitted"):
|
||||
# 追问卡是一次性的:重复提交会让同一个问题在上下文里出现两次答案
|
||||
return JsonResponse({"detail": "这个问题已经回答过了"}, status=409)
|
||||
payload = dict(card.payload or {})
|
||||
payload["answers"] = answers
|
||||
payload["submitted"] = True
|
||||
card.payload = payload
|
||||
card.save(update_fields=["payload", "updated_at"])
|
||||
labels = {f["key"]: f["label"] for f in payload.get("fields", [])}
|
||||
text = ";".join(
|
||||
f"{labels.get(k, k)}:{'、'.join(v) if isinstance(v, list) else v}"
|
||||
for k, v in answers.items()
|
||||
)
|
||||
if apply_session_params(conversation, payload.get("fields") or [], answers):
|
||||
text = f"{text}。请按新的会话参数重新写方案,旧方案作废"
|
||||
# 点选商品/角色必须钉成 Ref:模型常把选项做成单选文字,前端只回 answers。
|
||||
refs = list(refs)
|
||||
existing = {(item.get("type"), str(item.get("id"))) for item in refs if isinstance(item, dict)}
|
||||
for extra in refs_from_elicit_answers(conversation.team, payload.get("fields") or [], answers):
|
||||
mark = (extra.get("type"), str(extra.get("id")))
|
||||
if mark in existing:
|
||||
continue
|
||||
refs.append(extra)
|
||||
existing.add(mark)
|
||||
elif not text and not refs:
|
||||
return JsonResponse({"detail": "消息不能为空"}, status=400)
|
||||
|
||||
model_config = None
|
||||
requested = request.data.get("model_config_id")
|
||||
if requested:
|
||||
model_config = (
|
||||
ModelConfig.objects.select_related("provider")
|
||||
.filter(id=requested, capability=ModelConfig.Capability.TEXT, status=ModelConfig.Status.ACTIVE)
|
||||
.first()
|
||||
)
|
||||
|
||||
stream = stream_creation_agent(
|
||||
conversation=conversation,
|
||||
user=request.user,
|
||||
text=text,
|
||||
refs=refs,
|
||||
model_config=model_config,
|
||||
)
|
||||
response = StreamingHttpResponse(stream, content_type="text/event-stream")
|
||||
response["Cache-Control"] = "no-cache"
|
||||
response["X-Accel-Buffering"] = "no" # 关 nginx 缓冲,保证逐帧下发
|
||||
return response
|
||||
|
||||
|
||||
class MentionSearchView(APIView):
|
||||
"""@ 引用检索(契约 §3)。
|
||||
|
||||
GET /api/ai/mentions/?q=净颜&types=product,character&limit=8
|
||||
|
||||
返回 [Ref] —— 前端把它按 type 分组渲染成 @ 菜单(设计稿 .omni-mention-group)。
|
||||
**前端拿到后必须整条 Ref 存进消息的 refs 字段**,不能只把 name 拼进文本,
|
||||
否则后端取不到卖点和参考图(契约 §1)。
|
||||
"""
|
||||
|
||||
def get(self, request):
|
||||
team = get_current_team(request.user)
|
||||
q = str(request.query_params.get("q") or "").strip()
|
||||
raw_types = str(request.query_params.get("types") or "").strip()
|
||||
types = [t.strip() for t in raw_types.split(",") if t.strip()] if raw_types else None
|
||||
if types:
|
||||
unknown = [t for t in types if t not in VALID_TYPES]
|
||||
if unknown:
|
||||
raise ValidationError({"detail": f"未知引用类型:{'、'.join(unknown)}"})
|
||||
try:
|
||||
limit = int(request.query_params.get("limit") or 8)
|
||||
except (TypeError, ValueError) as exc:
|
||||
raise ValidationError({"detail": "limit 必须是整数"}) from exc
|
||||
limit = max(1, min(limit, 20))
|
||||
results = search_mentions(team, q=q, types=types, limit=limit)
|
||||
return Response({"results": results, "type_labels": TYPE_LABELS})
|
||||
|
||||
Reference in New Issue
Block a user