添加全能创作功能

This commit is contained in:
Azmat@qq.com
2026-09-03 13:11:46 +08:00
parent 22ed2833ad
commit 6a628b0ca7
53 changed files with 9115 additions and 108 deletions
+203
View File
@@ -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
+64
View File
@@ -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(), "")
+9
View File
@@ -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
+366
View File
@@ -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'),
),
]
+105
View File
@@ -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}"
+76 -1
View File
@@ -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",
]
+5
View File
@@ -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 之上 ──
+733
View File
@@ -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"])
+4
View File
@@ -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"),
+224 -2
View File
@@ -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})
+13
View File
@@ -1,4 +1,5 @@
from rest_framework.exceptions import PermissionDenied
from rest_framework.renderers import BaseRenderer
def get_current_team(user):
@@ -27,3 +28,15 @@ class TeamScopedViewSetMixin:
def perform_create(self, serializer):
serializer.save(team=self.get_team(), created_by=self.request.user)
class ServerSentEventRenderer(BaseRenderer):
"""让 DRF 内容协商接受 Accept: text/event-stream(否则流式端点直接 406)。
实际响应由视图返回 StreamingHttpResponse 直接下发,这个 renderer 只用于通过协商"""
media_type = "text/event-stream"
format = "event-stream"
charset = None
def render(self, data, accepted_media_type=None, renderer_context=None):
return data
@@ -1358,3 +1358,7 @@ class QuickCreateCoordinatorTests(TestCase):
+1 -14
View File
@@ -11,7 +11,6 @@ from rest_framework import status
from rest_framework.decorators import action
from rest_framework.exceptions import APIException, ValidationError
from rest_framework.parsers import FormParser, MultiPartParser
from rest_framework.renderers import BaseRenderer
from rest_framework.response import Response
from rest_framework.viewsets import ModelViewSet
@@ -42,7 +41,7 @@ from apps.ai.services import (
from apps.assets.models import Asset, AssetFile
from apps.assets.serializers import AssetFileSerializer
from apps.assets.storage import TosStorage
from apps.common.api import TeamScopedViewSetMixin
from apps.common.api import ServerSentEventRenderer, TeamScopedViewSetMixin
from apps.common.celery_health import require_worker, require_worker_task
from apps.ai.generation_errors import classify_generation_error, public_error_for_task
from apps.ai.video_digest import VideoDigestError, digest_project_video
@@ -94,18 +93,6 @@ class QuickCreateInProgress(APIException):
default_code = "quick_create_running"
class ServerSentEventRenderer(BaseRenderer):
"""让 DRF 内容协商接受 Accept: text/event-stream(否则流式端点直接 406)。
实际响应由视图返回 StreamingHttpResponse 直接下发,这个 renderer 只用于通过协商"""
media_type = "text/event-stream"
format = "event-stream"
charset = None
def render(self, data, accepted_media_type=None, renderer_context=None):
return data
def _store_uploaded_asset(*, team, user, upload, asset_type: str, category: str, name: str) -> Asset:
"""把上传的文件落到 TOS,建 Asset+AssetFile(主文件)。供上传视频段 / 上传 BGM 复用。"""
fallback_suffix = {