修改全能创作发现问题
This commit is contained in:
@@ -29,12 +29,23 @@ YUNQI_MODELS = [
|
||||
]
|
||||
|
||||
VOLCANO_MODELS = [
|
||||
{
|
||||
"display_name": "Doubao-Seed-2.1-Pro",
|
||||
"name": "doubao-seed-2-1-pro-260628",
|
||||
"capability": "text",
|
||||
"endpoint": "chat/completions",
|
||||
"metadata": {
|
||||
"think": True,
|
||||
"vision": True,
|
||||
"source": "volcengine ark · Seed 2.1 Pro",
|
||||
},
|
||||
},
|
||||
{
|
||||
"display_name": "Doubao-Seed-2.0-Pro",
|
||||
"name": "doubao-seed-2-0-pro-260215",
|
||||
"capability": "text",
|
||||
"endpoint": "chat/completions",
|
||||
"metadata": {"think": True, "source": "video-flow/data/vendor/volcengine.ts"},
|
||||
"metadata": {"think": True, "vision": True, "source": "video-flow/data/vendor/volcengine.ts"},
|
||||
},
|
||||
{
|
||||
"display_name": "Doubao-Seed-2.0-Lite",
|
||||
@@ -184,7 +195,10 @@ def _catalog_capabilities(item: dict) -> dict:
|
||||
capability = item["capability"]
|
||||
metadata = item["metadata"]
|
||||
if capability == "text":
|
||||
return {"operations": ["chat"], "features": ["streaming", "structured_output"]}
|
||||
features = ["streaming", "structured_output"]
|
||||
if metadata.get("vision"):
|
||||
features.extend(["vision", "image_input", "multimodal"])
|
||||
return {"operations": ["chat"], "features": features}
|
||||
if capability == "image":
|
||||
modes = set(metadata.get("modes") or [])
|
||||
supports_reference = bool(metadata.get("supports_reference")) or bool(
|
||||
|
||||
@@ -29,7 +29,7 @@ from .creation import append_message, pin_refs
|
||||
from .creation_presets import preset_guidance
|
||||
from .mentions import TYPE_LABELS, infer_field_types, resolve_refs, search_mentions
|
||||
from .models import CreationConversation, CreationMessage, ModelConfig
|
||||
from .services import build_provider, get_default_model
|
||||
from .services import build_provider, get_default_model, resolve_text_model
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -209,6 +209,52 @@ def apply_session_params(conversation, fields, answers: dict) -> bool:
|
||||
return changed
|
||||
|
||||
|
||||
def snapshot_session_params(conversation) -> dict:
|
||||
params = conversation.params or {}
|
||||
return {
|
||||
"model": str(params.get("model") or ""),
|
||||
"resolution": str(params.get("resolution") or ""),
|
||||
"ratio": str(params.get("ratio") or ""),
|
||||
"duration": str(params.get("duration") or ""),
|
||||
"count": str(params.get("count") or params.get("duration") or ""),
|
||||
}
|
||||
|
||||
|
||||
def confirm_param_options(is_video: bool) -> dict:
|
||||
return {
|
||||
"model": VIDEO_MODELS if is_video else IMAGE_MODELS,
|
||||
"resolution": RESOLUTIONS if is_video else [],
|
||||
"ratio": RATIOS,
|
||||
"duration": VIDEO_DURATIONS if is_video else [],
|
||||
"count": IMAGE_COUNTS if not is_video else [],
|
||||
}
|
||||
|
||||
|
||||
def apply_confirm_params(conversation, incoming: dict | None) -> tuple[dict, bool]:
|
||||
"""确认卡上改的参数写回会话。返回 (最新 params, 视频时长是否变了)。"""
|
||||
current = dict(conversation.params or {})
|
||||
old_duration = str(current.get("duration") or "")
|
||||
changed = False
|
||||
for key, raw in (incoming or {}).items():
|
||||
if key not in {"model", "ratio", "resolution", "duration", "count"}:
|
||||
continue
|
||||
value = str(raw or "").strip()
|
||||
if not value or current.get(key) == value:
|
||||
continue
|
||||
current[key] = value
|
||||
changed = True
|
||||
duration_changed = (
|
||||
conversation.mode == CreationConversation.Mode.VIDEO
|
||||
and str(current.get("duration") or "") != old_duration
|
||||
and bool(str(current.get("duration") or ""))
|
||||
and bool(old_duration)
|
||||
)
|
||||
if changed:
|
||||
conversation.params = current
|
||||
conversation.save(update_fields=["params", "updated_at"])
|
||||
return snapshot_session_params(conversation), duration_changed
|
||||
|
||||
|
||||
# ---------------------------------------------------------------- 工具 schema
|
||||
|
||||
def tool_schemas(context: AgentContext) -> list[dict]:
|
||||
@@ -330,24 +376,6 @@ def tool_schemas(context: AgentContext) -> list[dict]:
|
||||
"required": ["start", "end", "stage"],
|
||||
},
|
||||
},
|
||||
"matrix": {
|
||||
"type": "object",
|
||||
"description": "卖点覆盖矩阵:哪个卖点落在第几个镜头",
|
||||
"properties": {
|
||||
"shots": {"type": "integer"},
|
||||
"rows": {
|
||||
"type": "array",
|
||||
"items": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"point": {"type": "string"},
|
||||
"hits": {"type": "array", "items": {"type": "integer"}},
|
||||
},
|
||||
"required": ["point", "hits"],
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
"voice_chars": {
|
||||
"type": "array", "items": {"type": "integer"},
|
||||
"description": "口播字数区间 [下限, 上限]",
|
||||
@@ -587,9 +615,134 @@ def submit_confirmed_video(*, conversation: CreationConversation, user, confirm_
|
||||
return message, ""
|
||||
|
||||
|
||||
def submit_confirmed_image(*, conversation: CreationConversation, user, confirm_message: CreationMessage):
|
||||
"""用户点了确认 → 按确认卡里存的画面描述出图。同样不跑一轮模型。"""
|
||||
payload = confirm_message.payload or {}
|
||||
prompt = str(payload.get("prompt") or payload.get("image_prompt") or "").strip()
|
||||
if not prompt:
|
||||
return None, "这条方案没有存下出图指令,请让我重新写一次。"
|
||||
|
||||
context = AgentContext(conversation=conversation, user=user, model_config=None)
|
||||
try:
|
||||
_result, tasks = _run_generate_image(context, {"prompt": prompt})
|
||||
except AgentError as exc:
|
||||
return None, str(exc)
|
||||
except ValueError as exc:
|
||||
return None, str(exc)
|
||||
|
||||
message = None
|
||||
for task in tasks:
|
||||
message = append_message(
|
||||
conversation, role="assistant",
|
||||
kind=CreationMessage.Kind.GENERATING,
|
||||
payload={"task_id": str(task.id), "kind": "image", "prompt": prompt},
|
||||
task=task,
|
||||
)
|
||||
if message is None:
|
||||
return None, "出图没有提交成功,请再试一次。"
|
||||
_remember_artifact(conversation, prompt, "image")
|
||||
return message, ""
|
||||
|
||||
|
||||
# ---------------------------------------------------------------- 提示词
|
||||
|
||||
|
||||
|
||||
def get_creation_chat_model(requested: ModelConfig | None = None) -> ModelConfig | None:
|
||||
"""全能创作对话模型:走平台文本解析,DeepSeek 一律换成 Seed 2.1 Pro。"""
|
||||
return resolve_text_model(requested)
|
||||
|
||||
|
||||
def _creation_model_sees_images(model_config: ModelConfig | None) -> bool:
|
||||
"""对话模型能不能收图。参考图只在能看图时才塞进 chat messages,避免纯文本模型整轮失败。"""
|
||||
if model_config is None:
|
||||
return False
|
||||
if getattr(model_config, "capability", "") == ModelConfig.Capability.VISION:
|
||||
return True
|
||||
name = str(getattr(model_config, "name", "") or "").lower()
|
||||
# 豆包 Seed 2.x / 1.6 文本档都支持图文;vl / vision 后缀同理。
|
||||
if name.startswith("doubao-seed-") or "vision" in name or name.endswith("-vl") or "-vl-" in name:
|
||||
return True
|
||||
metadata = model_config.metadata if isinstance(getattr(model_config, "metadata", None), dict) else {}
|
||||
capabilities = metadata.get("capabilities") if isinstance(metadata.get("capabilities"), dict) else {}
|
||||
features = {str(item) for item in capabilities.get("features") or []}
|
||||
return bool({"vision", "image_input", "multimodal"} & features)
|
||||
|
||||
|
||||
def _prefer_vision_text_model(current: ModelConfig | None, team, refs: list | None) -> ModelConfig | None:
|
||||
"""有参考图时,尽量换成能看图的文本模型(豆包 Seed 等),否则聊天侧完全看不见男女。"""
|
||||
if current is not None and _creation_model_sees_images(current):
|
||||
return current
|
||||
if not _ref_image_urls(team, refs):
|
||||
return current
|
||||
qs = (
|
||||
ModelConfig.objects.select_related("provider")
|
||||
.filter(
|
||||
capability=ModelConfig.Capability.TEXT,
|
||||
status=ModelConfig.Status.ACTIVE,
|
||||
provider__status="active",
|
||||
)
|
||||
.order_by("created_at")
|
||||
)
|
||||
for candidate in qs:
|
||||
if _creation_model_sees_images(candidate):
|
||||
return candidate
|
||||
return get_default_model(ModelConfig.Capability.VISION) or current
|
||||
|
||||
|
||||
def _ref_image_urls(team, refs: list | None) -> list[str]:
|
||||
resolved = resolve_refs(team, refs or [])
|
||||
urls: list[str] = []
|
||||
seen: set[str] = set()
|
||||
for item in resolved.references:
|
||||
url = str(item.get("url") or "").strip()
|
||||
if not url or url in seen:
|
||||
continue
|
||||
seen.add(url)
|
||||
urls.append(url)
|
||||
return urls[:6]
|
||||
|
||||
|
||||
def _attach_ref_images(messages: list[dict], image_urls: list[str]) -> list[dict]:
|
||||
"""把锁定素材图挂到最近一条 user 消息上(OpenAI image_url 格式)。"""
|
||||
if not image_urls or not messages:
|
||||
return messages
|
||||
note = (
|
||||
f"【参考图·请亲眼看】下面 {len(image_urls)} 张是用户锁定的素材。"
|
||||
"人物的性别、年龄段、发型、服装必须以图为准;看不清再问用户,禁止凭文件名猜测性别。"
|
||||
)
|
||||
out = [dict(message) for message in messages]
|
||||
index = next((i for i in range(len(out) - 1, -1, -1) if out[i].get("role") == "user"), None)
|
||||
if index is None:
|
||||
content = [{"type": "text", "text": note}]
|
||||
content.extend({"type": "image_url", "image_url": {"url": url}} for url in image_urls)
|
||||
out.append({"role": "user", "content": content})
|
||||
return out
|
||||
last = dict(out[index])
|
||||
raw = last.get("content")
|
||||
if isinstance(raw, list):
|
||||
content = list(raw)
|
||||
text_bits = [str(item.get("text") or "") for item in content if isinstance(item, dict) and item.get("type") == "text"]
|
||||
if not any(note[:8] in bit for bit in text_bits):
|
||||
content.append({"type": "text", "text": note})
|
||||
existing = {
|
||||
(item.get("image_url") or {}).get("url")
|
||||
for item in content
|
||||
if isinstance(item, dict) and item.get("type") == "image_url"
|
||||
}
|
||||
content.extend(
|
||||
{"type": "image_url", "image_url": {"url": url}}
|
||||
for url in image_urls
|
||||
if url not in existing
|
||||
)
|
||||
else:
|
||||
content = [{"type": "text", "text": f"{raw or ''}\n\n{note}".strip()}]
|
||||
content.extend({"type": "image_url", "image_url": {"url": url}} for url in image_urls)
|
||||
last["content"] = content
|
||||
out[index] = last
|
||||
return out
|
||||
|
||||
|
||||
def build_system_prompt(context: AgentContext) -> str:
|
||||
conversation = context.conversation
|
||||
params = conversation.params or {}
|
||||
@@ -635,7 +788,9 @@ def build_system_prompt(context: AgentContext) -> str:
|
||||
lines.append("\n【本次会话已锁定的素材事实】")
|
||||
lines.append(resolved.facts_text)
|
||||
lines.append(
|
||||
"以上素材的参考图会自动附给生成模型锁人锁物,你在 prompt 里不需要重复描述它们的外观。"
|
||||
"以上素材的参考图会附给你看,出片时也会自动锁人锁物。"
|
||||
"人物性别、年龄段、发型、服装、商品颜色外形必须以图为准;"
|
||||
"图上看不清或没附图时,必须问用户,禁止凭文件名猜测男女。"
|
||||
)
|
||||
memory = conversation.memory or {}
|
||||
if memory.get("summary"):
|
||||
@@ -682,6 +837,11 @@ def build_messages(context: AgentContext) -> list[dict]:
|
||||
messages.append({"role": "assistant", "content": f"(我生成了一版,prompt:{prompt[:200]})"})
|
||||
elif message.kind == CreationMessage.Kind.ERROR:
|
||||
messages.append({"role": "assistant", "content": f"(上一次生成失败:{message.text})"})
|
||||
if _creation_model_sees_images(context.model_config):
|
||||
messages = _attach_ref_images(
|
||||
messages,
|
||||
_ref_image_urls(context.team, context.conversation.pinned_refs or []),
|
||||
)
|
||||
return messages
|
||||
|
||||
|
||||
@@ -727,7 +887,7 @@ def stream_creation_agent(
|
||||
) -> Iterator[str]:
|
||||
"""一条用户消息 → SSE 流。生成器,由 StreamingHttpResponse 逐帧下发。"""
|
||||
refs = refs or []
|
||||
model_config = model_config or get_default_model(ModelConfig.Capability.TEXT)
|
||||
model_config = get_creation_chat_model(model_config)
|
||||
if model_config is None:
|
||||
yield _sse({"type": "error", "detail": "没有可用的文本模型,请先在模型库配置"})
|
||||
return
|
||||
@@ -738,6 +898,8 @@ def stream_creation_agent(
|
||||
pin_refs(conversation, refs)
|
||||
user_message = append_message(conversation, role="user", text=text, refs=refs)
|
||||
yield _sse({"type": "message", "message": _message_payload(user_message)})
|
||||
model_config = _prefer_vision_text_model(model_config, conversation.team, conversation.pinned_refs or [])
|
||||
context.model_config = model_config
|
||||
|
||||
resolved = resolve_refs(context.team, refs)
|
||||
if resolved.missing:
|
||||
@@ -898,7 +1060,6 @@ def _dispatch_tool(context: AgentContext, name: str, args: dict) -> tuple[dict,
|
||||
"usp": str(args.get("usp") or ""),
|
||||
"points": [str(p) for p in (args.get("points") or [])][:3],
|
||||
"timeline": args.get("timeline") or [],
|
||||
"matrix": args.get("matrix") or {},
|
||||
"voice_chars": args.get("voice_chars") or [],
|
||||
"ref_count": len(context.conversation.pinned_refs or []),
|
||||
}
|
||||
@@ -922,8 +1083,15 @@ def _dispatch_tool(context: AgentContext, name: str, args: dict) -> tuple[dict,
|
||||
confirm = append_message(
|
||||
context.conversation, role="assistant",
|
||||
kind=CreationMessage.Kind.CONFIRM,
|
||||
payload={"label": "开始生成", "estimated_credits": credits,
|
||||
"video_prompt": video_prompt, "submitted": False},
|
||||
payload={
|
||||
"kind": "video",
|
||||
"label": "开始生成",
|
||||
"estimated_credits": credits,
|
||||
"video_prompt": video_prompt,
|
||||
"submitted": False,
|
||||
"params": snapshot_session_params(context.conversation),
|
||||
"param_options": confirm_param_options(True),
|
||||
},
|
||||
)
|
||||
events.append({"type": "message", "message": _message_payload(confirm)})
|
||||
events.append({"type": "credits", "estimated": credits})
|
||||
@@ -934,19 +1102,25 @@ def _dispatch_tool(context: AgentContext, name: str, args: dict) -> tuple[dict,
|
||||
if context.generations_used >= MAX_BILLED_GENERATIONS:
|
||||
# 一条用户消息只计费一次。模型想连出好几版时在这里挡住。
|
||||
return {"payload": {"error": "本轮已经生成过一次了,请让用户看过再决定要不要改"}}, True
|
||||
payload, tasks = _run_generate_image(context, args)
|
||||
events = []
|
||||
for task in tasks:
|
||||
message = append_message(
|
||||
context.conversation, role="assistant",
|
||||
kind=CreationMessage.Kind.GENERATING,
|
||||
payload={"task_id": str(task.id), "kind": "image", "prompt": args.get("prompt", "")},
|
||||
task=task,
|
||||
)
|
||||
events.append({"type": "message", "message": _message_payload(message)})
|
||||
events.append({"type": "task", "task_id": str(task.id), "kind": "image"})
|
||||
_remember_artifact(context.conversation, args.get("prompt", ""), "image")
|
||||
return {"payload": payload, "_events": events}, True
|
||||
prompt = str(args.get("prompt") or "").strip()
|
||||
if not prompt:
|
||||
return {"payload": {"error": "生成失败:模型没有给出画面描述"}}, False
|
||||
# 出图也走确认卡:用户先看当前模型/比例/张数,点了才提交。
|
||||
confirm = append_message(
|
||||
context.conversation, role="assistant",
|
||||
kind=CreationMessage.Kind.CONFIRM,
|
||||
payload={
|
||||
"kind": "image",
|
||||
"label": "开始生成",
|
||||
"estimated_credits": 0,
|
||||
"prompt": prompt,
|
||||
"submitted": False,
|
||||
"params": snapshot_session_params(context.conversation),
|
||||
"param_options": confirm_param_options(False),
|
||||
},
|
||||
)
|
||||
events = [{"type": "message", "message": _message_payload(confirm)}]
|
||||
return {"payload": {"awaiting_confirmation": True}, "_events": events}, True
|
||||
|
||||
return {"payload": {"error": f"未知工具 {name}"}}, False
|
||||
|
||||
|
||||
@@ -839,6 +839,8 @@ def _store_free_video_media(*, task: AITask, media: str = "", video_bytes: bytes
|
||||
if feature == "video_replace":
|
||||
mode_label = "角色复刻" if payload.get("replace_mode") == "character" else "商品复刻"
|
||||
asset_name = f"{subject}{mode_label}" if subject else mode_label
|
||||
elif feature == "omni_create":
|
||||
asset_name = prompt[:255] or "全能创作视频"
|
||||
else:
|
||||
asset_name = prompt[:255] or "自由创作视频"
|
||||
asset = Asset.objects.create(
|
||||
@@ -852,6 +854,8 @@ def _store_free_video_media(*, task: AITask, media: str = "", video_bytes: bytes
|
||||
source=Asset.Source.AI_GENERATED,
|
||||
category=Asset.Category.FREE_CREATE,
|
||||
origin_task=task,
|
||||
# 全能创作成品只挂在会话结果卡上,不进资产库/专业创作用的视频列表。
|
||||
in_library=feature != "omni_create",
|
||||
metadata={"feature": feature},
|
||||
)
|
||||
AssetFile.objects.create(
|
||||
|
||||
@@ -231,6 +231,19 @@ def refs_from_elicit_answers(team, fields, answers: dict) -> list[dict]:
|
||||
return refs
|
||||
|
||||
|
||||
def _person_meta_facts(metadata) -> list[str]:
|
||||
"""角色/模特库里填过的性别、年龄。聊天模型不一定看得到图,这些字必须进事实块。"""
|
||||
meta = metadata if isinstance(metadata, dict) else {}
|
||||
lines = []
|
||||
gender = str(meta.get("gender") or meta.get("sex") or "").strip()
|
||||
age = str(meta.get("age") or meta.get("age_range") or "").strip()
|
||||
if gender:
|
||||
lines.append(f"性别:{gender}")
|
||||
if age:
|
||||
lines.append(f"年龄:{age}")
|
||||
return lines
|
||||
|
||||
|
||||
def product_facts_text(product) -> str:
|
||||
"""商品事实块。全能创作没有 project,所以不能复用 script_agent._product_context()。
|
||||
这里只给**客观事实**(标题/品牌/品类/规格/卖点),不带人设和口吻 —— 那些由策略卡决定。"""
|
||||
@@ -270,6 +283,56 @@ def _asset_reference(asset, type_: str, label: str) -> dict | None:
|
||||
}
|
||||
|
||||
|
||||
|
||||
def _product_triview_asset(product):
|
||||
"""商品三视图是独立 Asset(metadata.view=three_view),不在商品图列表里。没有就返回 None。"""
|
||||
return (
|
||||
Asset.objects.filter(
|
||||
team_id=product.team_id,
|
||||
is_deleted=False,
|
||||
purged_at__isnull=True,
|
||||
metadata__product_id=str(product.id),
|
||||
metadata__view="three_view",
|
||||
)
|
||||
.order_by("-created_at")
|
||||
.first()
|
||||
)
|
||||
|
||||
|
||||
def _character_triview_asset(team, portrait_id):
|
||||
"""角色立绘配对的三视图:metadata.triview_of = 立绘 id。没有就返回 None。"""
|
||||
if not portrait_id:
|
||||
return None
|
||||
return (
|
||||
Asset.objects.filter(
|
||||
team=team,
|
||||
is_deleted=False,
|
||||
purged_at__isnull=True,
|
||||
category=Asset.Category.TRI_VIEW,
|
||||
metadata__triview_of=str(portrait_id),
|
||||
)
|
||||
.order_by("-created_at")
|
||||
.first()
|
||||
) or (
|
||||
Asset.objects.filter(
|
||||
team=team,
|
||||
is_deleted=False,
|
||||
purged_at__isnull=True,
|
||||
metadata__triview_of=str(portrait_id),
|
||||
)
|
||||
.order_by("-created_at")
|
||||
.first()
|
||||
)
|
||||
|
||||
|
||||
def _append_ref(resolved: ResolvedRefs, asset, type_: str, label: str) -> bool:
|
||||
entry = _asset_reference(asset, type_, label)
|
||||
if not entry:
|
||||
return False
|
||||
resolved.references.append(entry)
|
||||
return True
|
||||
|
||||
|
||||
def _product_reference(product) -> dict | None:
|
||||
"""商品参考图:**真实上传图优先,排除 AI 生成图** —— 拿生成图当真相再喂回模型会误差累积。
|
||||
一张真实图都没有才回落封面(可能是 AI 图,但好过纯文生图)。
|
||||
@@ -314,6 +377,9 @@ def resolve_refs(team, refs: list[dict]) -> ResolvedRefs:
|
||||
entry = _product_reference(product)
|
||||
if entry:
|
||||
resolved.references.append(entry)
|
||||
triview = _product_triview_asset(product)
|
||||
if triview is not None:
|
||||
_append_ref(resolved, triview, "product", f"{product.title}三视图")
|
||||
continue
|
||||
if type_ == "model":
|
||||
model = Model.objects.filter(
|
||||
@@ -322,15 +388,23 @@ def resolve_refs(team, refs: list[dict]) -> ResolvedRefs:
|
||||
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)
|
||||
person_lines = []
|
||||
desc = (model.description or "").strip()
|
||||
meta_lines = _person_meta_facts(model.metadata)
|
||||
if desc:
|
||||
person_lines.append(f"模特「{model.name}」:{desc}")
|
||||
else:
|
||||
person_lines.append(f"模特「{model.name}」")
|
||||
person_lines.extend(meta_lines)
|
||||
if desc or meta_lines:
|
||||
resolved.facts.append("\n".join(person_lines))
|
||||
# 形象图和三视图都带上(有就带,没有不加、不报错)。三视图锁脸更稳。
|
||||
added = False
|
||||
if _append_ref(resolved, model.portrait_asset, "model", model.name):
|
||||
added = True
|
||||
if _append_ref(resolved, model.triview_asset, "model", f"{model.name}三视图"):
|
||||
added = True
|
||||
if not added:
|
||||
resolved.missing.append(ref)
|
||||
continue
|
||||
asset = Asset.objects.filter(
|
||||
@@ -339,12 +413,23 @@ def resolve_refs(team, refs: list[dict]) -> ResolvedRefs:
|
||||
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:
|
||||
desc = (asset.description or "").strip()
|
||||
meta_lines = _person_meta_facts(asset.metadata)
|
||||
label = TYPE_LABELS[type_]
|
||||
person_lines = [f"{label}「{asset.name}」:{desc}" if desc else f"{label}「{asset.name}」"]
|
||||
person_lines.extend(meta_lines)
|
||||
if desc or meta_lines:
|
||||
resolved.facts.append("\n".join(person_lines))
|
||||
added = _append_ref(resolved, asset, type_, asset.name)
|
||||
# 角色/模特立绘若有配对三视图,一并带给出片;没有就跳过,不要当成引用失败。
|
||||
if type_ in {"character", "model", "asset"}:
|
||||
portrait_id = asset.id
|
||||
# 用户 @ 的就是三视图本身时,不再反查。
|
||||
if asset.category != Asset.Category.TRI_VIEW:
|
||||
paired = _character_triview_asset(team, portrait_id)
|
||||
if paired is not None:
|
||||
_append_ref(resolved, paired, type_, f"{asset.name}三视图")
|
||||
if not added:
|
||||
resolved.missing.append(ref)
|
||||
|
||||
# 角色 → 场景 → 商品。同优先级内保持用户 @ 的先后。
|
||||
|
||||
@@ -0,0 +1,85 @@
|
||||
"""把默认文本模型换成火山 Doubao-Seed-2.1-Pro。
|
||||
|
||||
全能创作 / 实体提取 / 脚本思考都走 get_default_model(TEXT)。
|
||||
此前后台默认是 DeepSeek(DP V4 Pro),看不了参考图,聊天里分不出男女。
|
||||
Seed 2.1 Pro 支持图文 + function calling,模型 ID: doubao-seed-2-1-pro-260628。
|
||||
|
||||
幂等:可重复 apply。不改其它能力(生图/视频)的 is_default。
|
||||
"""
|
||||
|
||||
from django.db import migrations
|
||||
|
||||
NAME = "doubao-seed-2-1-pro-260628"
|
||||
DISPLAY = "Doubao-Seed-2.1-Pro"
|
||||
METADATA = {
|
||||
"think": True,
|
||||
"vision": True,
|
||||
"source": "volcengine ark · Seed 2.1 Pro",
|
||||
"routing": {"fallback_on_failure": False, "fallback_candidate": True},
|
||||
"capabilities": {
|
||||
"operations": ["chat"],
|
||||
"features": ["streaming", "structured_output", "vision", "image_input", "multimodal"],
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def seed(apps, schema_editor):
|
||||
ModelProvider = apps.get_model("ai", "ModelProvider")
|
||||
ModelConfig = apps.get_model("ai", "ModelConfig")
|
||||
|
||||
provider, _ = ModelProvider.objects.get_or_create(
|
||||
name="volcengine",
|
||||
defaults={
|
||||
"display_name": "火山引擎(豆包)",
|
||||
"status": "active",
|
||||
"base_url": "https://ark.cn-beijing.volces.com/api/v3",
|
||||
},
|
||||
)
|
||||
config, created = ModelConfig.objects.get_or_create(
|
||||
provider=provider,
|
||||
name=NAME,
|
||||
capability="text",
|
||||
defaults={
|
||||
"display_name": DISPLAY,
|
||||
"endpoint": "chat/completions",
|
||||
"status": "active",
|
||||
"is_default": True,
|
||||
"metadata": METADATA,
|
||||
},
|
||||
)
|
||||
metadata = dict(config.metadata or {})
|
||||
for key, value in METADATA.items():
|
||||
if key == "capabilities":
|
||||
caps = dict(metadata.get("capabilities") or {})
|
||||
caps.setdefault("operations", value["operations"])
|
||||
features = list(caps.get("features") or [])
|
||||
for feat in value["features"]:
|
||||
if feat not in features:
|
||||
features.append(feat)
|
||||
caps["features"] = features
|
||||
metadata["capabilities"] = caps
|
||||
else:
|
||||
metadata.setdefault(key, value)
|
||||
config.metadata = metadata
|
||||
config.display_name = config.display_name or DISPLAY
|
||||
config.endpoint = config.endpoint or "chat/completions"
|
||||
config.status = "active"
|
||||
config.is_default = True
|
||||
config.save()
|
||||
|
||||
# 同一能力只留一个默认:清掉 DeepSeek 等其它文本模型的 is_default。
|
||||
ModelConfig.objects.filter(capability="text", is_default=True).exclude(pk=config.pk).update(is_default=False)
|
||||
|
||||
|
||||
def unseed(apps, schema_editor):
|
||||
ModelConfig = apps.get_model("ai", "ModelConfig")
|
||||
ModelConfig.objects.filter(name=NAME, capability="text").update(is_default=False)
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
dependencies = [
|
||||
("ai", "0034_creationconversation_creationmessage_and_more"),
|
||||
]
|
||||
operations = [
|
||||
migrations.RunPython(seed, unseed),
|
||||
]
|
||||
@@ -0,0 +1,39 @@
|
||||
"""彻底停用 DeepSeek 文本模型,默认只留 Seed 2.1 Pro。
|
||||
|
||||
幂等:可重复 apply。不删行,方便后台还能看见历史配置。
|
||||
"""
|
||||
|
||||
from django.db import migrations
|
||||
from django.db.models import Q
|
||||
|
||||
|
||||
def apply(apps, schema_editor):
|
||||
ModelConfig = apps.get_model("ai", "ModelConfig")
|
||||
retired = ModelConfig.objects.filter(capability="text").filter(
|
||||
Q(name__icontains="deepseek") | Q(display_name__icontains="deepseek") | Q(display_name__icontains="DP V4")
|
||||
)
|
||||
retired.update(status="disabled", is_default=False)
|
||||
seed = (
|
||||
ModelConfig.objects.filter(name="doubao-seed-2-1-pro-260628", capability="text")
|
||||
.order_by("created_at")
|
||||
.first()
|
||||
)
|
||||
if seed is None:
|
||||
return
|
||||
seed.status = "active"
|
||||
seed.is_default = True
|
||||
seed.save(update_fields=["status", "is_default", "updated_at"])
|
||||
ModelConfig.objects.filter(capability="text", is_default=True).exclude(pk=seed.pk).update(is_default=False)
|
||||
|
||||
|
||||
def noop(apps, schema_editor):
|
||||
return
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
dependencies = [
|
||||
("ai", "0035_seed_21_pro_default_text"),
|
||||
]
|
||||
operations = [
|
||||
migrations.RunPython(apply, noop),
|
||||
]
|
||||
@@ -54,11 +54,61 @@ logger = logging.getLogger(__name__)
|
||||
OFFICIAL_DIRECT_PROVIDERS = {"volcengine", "volcano", "ark", "volcano_ark", "doubao"}
|
||||
|
||||
|
||||
SEED_21_PRO_NAME = "doubao-seed-2-1-pro-260628"
|
||||
|
||||
|
||||
def is_retired_text_model(model) -> bool:
|
||||
"""DeepSeek / DP V4 Pro 已停用,解析层一律跳过。"""
|
||||
blob = f"{getattr(model, 'name', '')} {getattr(model, 'display_name', '')}".lower()
|
||||
return "deepseek" in blob or "dp v4" in blob or "dp-v4" in blob
|
||||
|
||||
|
||||
def get_seed_text_model() -> ModelConfig | None:
|
||||
qs = (
|
||||
ModelConfig.objects.select_related("provider")
|
||||
.filter(
|
||||
capability=ModelConfig.Capability.TEXT,
|
||||
status=ModelConfig.Status.ACTIVE,
|
||||
provider__status="active",
|
||||
)
|
||||
)
|
||||
return (
|
||||
qs.filter(name=SEED_21_PRO_NAME).first()
|
||||
or qs.filter(name__startswith="doubao-seed-2-1-pro").first()
|
||||
or qs.filter(name__startswith="doubao-seed-2-0-pro").first()
|
||||
)
|
||||
|
||||
|
||||
def resolve_text_model(requested: ModelConfig | None = None, requested_id=None) -> ModelConfig | None:
|
||||
"""文本模型入口:用户选了 DeepSeek 也改走 Seed 2.1 Pro。"""
|
||||
model = requested
|
||||
if model is None and requested_id:
|
||||
model = (
|
||||
ModelConfig.objects.select_related("provider")
|
||||
.filter(id=requested_id, capability=ModelConfig.Capability.TEXT, status=ModelConfig.Status.ACTIVE)
|
||||
.first()
|
||||
)
|
||||
if model is not None and not is_retired_text_model(model):
|
||||
return model
|
||||
return get_default_model(ModelConfig.Capability.TEXT)
|
||||
|
||||
|
||||
def get_default_model(capability: str) -> ModelConfig:
|
||||
qs = (
|
||||
ModelConfig.objects.select_related("provider")
|
||||
.filter(capability=capability, status=ModelConfig.Status.ACTIVE, provider__status="active")
|
||||
)
|
||||
if capability == ModelConfig.Capability.TEXT:
|
||||
for model in qs.filter(is_default=True).order_by("created_at"):
|
||||
if not is_retired_text_model(model):
|
||||
return model
|
||||
seed = get_seed_text_model()
|
||||
if seed is not None:
|
||||
return seed
|
||||
for model in qs.order_by("created_at"):
|
||||
if not is_retired_text_model(model):
|
||||
return model
|
||||
return None
|
||||
# 优先平台超管钦定的默认模型;未钦定则回落「最早创建的 active」(原行为,零回归)
|
||||
return qs.filter(is_default=True).order_by("created_at").first() or qs.order_by("created_at").first()
|
||||
|
||||
|
||||
@@ -15,6 +15,7 @@ from apps.products.models import Product, ProductImage
|
||||
|
||||
from .creation import append_message
|
||||
from .creation_agent import (
|
||||
AgentContext,
|
||||
COMPRESS_MIN_BATCH,
|
||||
DEFAULT_VIDEO_MODEL,
|
||||
KEEP_RECENT_MESSAGES,
|
||||
@@ -22,7 +23,10 @@ from .creation_agent import (
|
||||
_coerce_fields,
|
||||
_image_count,
|
||||
_merge_tool_call_deltas,
|
||||
build_messages,
|
||||
get_creation_chat_model,
|
||||
stream_creation_agent,
|
||||
submit_confirmed_image,
|
||||
submit_confirmed_video,
|
||||
video_duration,
|
||||
video_model_name,
|
||||
@@ -241,51 +245,67 @@ class GenerateImageTests(CreationAgentBaseTests):
|
||||
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):
|
||||
def _confirm_from_events(self, events):
|
||||
confirm = [e for e in events if e.get("type") == "message" and e["message"]["kind"] == "confirm"]
|
||||
self.assertEqual(len(confirm), 1)
|
||||
return CreationMessage.objects.get(id=confirm[0]["message"]["id"])
|
||||
|
||||
def test_generate_image_emits_confirm_then_submits_with_pinned_reference(self):
|
||||
events, _ = self._run(
|
||||
[_tool_chunks("generate_image", {"prompt": "干净棚拍,柔光,居中构图"})],
|
||||
refs=[{"type": "product", "id": str(self.product.id), "name": "净颜精华"}],
|
||||
)
|
||||
card = self._confirm_from_events(events)
|
||||
self.assertEqual(card.payload["kind"], "image")
|
||||
self.assertEqual(card.payload["prompt"], "干净棚拍,柔光,居中构图")
|
||||
self.assertFalse(any(e.get("type") == "task" for e in events))
|
||||
|
||||
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": "净颜精华"}],
|
||||
message, error = submit_confirmed_image(
|
||||
conversation=self.conversation, user=self.user, confirm_message=card,
|
||||
)
|
||||
kwargs = enqueue.call_args.kwargs
|
||||
|
||||
self.assertEqual(error, "")
|
||||
self.assertEqual(message.kind, CreationMessage.Kind.GENERATING)
|
||||
# @ 引用的商品图必须作为参考图带上,否则出的图跟商品长得不一样
|
||||
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"])
|
||||
events, _ = self._run(
|
||||
[_tool_chunks("generate_image", {"prompt": "白底棚拍", "count": 1})],
|
||||
)
|
||||
card = self._confirm_from_events(events)
|
||||
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})],
|
||||
submit_confirmed_image(
|
||||
conversation=self.conversation, user=self.user, confirm_message=card,
|
||||
)
|
||||
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)
|
||||
events, fake = self._run([
|
||||
_tool_chunks("generate_image", {"prompt": "第一版"}),
|
||||
_tool_chunks("generate_image", {"prompt": "第二版"}),
|
||||
])
|
||||
# 确认闸门打断循环:模型不能连出两张未确认的图
|
||||
self.assertEqual(len(fake.calls), 1)
|
||||
confirms = [e for e in events if e.get("type") == "message" and e["message"]["kind"] == "confirm"]
|
||||
self.assertEqual(len(confirms), 1)
|
||||
|
||||
def test_prompt_is_remembered_for_the_next_revision(self):
|
||||
events, _ = self._run([_tool_chunks("generate_image", {"prompt": "白底棚拍"})])
|
||||
card = self._confirm_from_events(events)
|
||||
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": "白底棚拍"})])
|
||||
submit_confirmed_image(
|
||||
conversation=self.conversation, user=self.user, confirm_message=card,
|
||||
)
|
||||
self.conversation.refresh_from_db()
|
||||
# 产物索引:下一轮「背景换夜景」要靠它知道在改哪一版
|
||||
self.assertEqual(self.conversation.memory["artifacts"][-1]["prompt"], "白底棚拍")
|
||||
@@ -342,13 +362,24 @@ class SseFramingTests(CreationAgentBaseTests):
|
||||
team=self.team, created_by=self.user, task_type=AITask.Type.PRODUCT_IMAGE,
|
||||
model_config=self.model, idempotency_key="k-sse",
|
||||
)
|
||||
events, _ = self._run([_tool_chunks("generate_image", {"prompt": "白底棚拍"})])
|
||||
card = next(e["message"] for e in events if e.get("type") == "message"
|
||||
and e["message"]["kind"] == "confirm")
|
||||
confirm = CreationMessage.objects.get(id=card["id"])
|
||||
with patch("apps.ai.services.enqueue_standalone_images", return_value=[task]):
|
||||
events, _ = self._run([_tool_chunks("generate_image", {"prompt": "白底棚拍"})])
|
||||
message, error = submit_confirmed_image(
|
||||
conversation=self.conversation, user=self.user, confirm_message=confirm,
|
||||
)
|
||||
|
||||
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))
|
||||
self.assertEqual(error, "")
|
||||
payload = {
|
||||
"id": str(message.id),
|
||||
"kind": message.kind,
|
||||
"task": str(message.task_id),
|
||||
"created_at": message.created_at,
|
||||
}
|
||||
encoded = json.dumps(payload, cls=__import__("django.core.serializers.json", fromlist=["DjangoJSONEncoder"]).DjangoJSONEncoder)
|
||||
self.assertIn(str(task.id), encoded)
|
||||
|
||||
|
||||
class SendEndpointTests(TestCase):
|
||||
@@ -464,7 +495,6 @@ class VideoPlanAndConfirmTests(CreationAgentBaseTests):
|
||||
"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秒 近景手持商品…",
|
||||
}
|
||||
@@ -616,6 +646,57 @@ class ConfirmEndpointTests(TestCase):
|
||||
# 出片没提交成功,闸门要放回去让用户改完再确认
|
||||
self.assertFalse(self.card.payload["submitted"])
|
||||
|
||||
def test_duration_change_returns_regenerate_instead_of_submitting(self):
|
||||
self.conversation.params = {
|
||||
"model": "Seedance 2.0 Fast", "resolution": "480p",
|
||||
"ratio": "1:1", "duration": "8 秒",
|
||||
}
|
||||
self.conversation.save(update_fields=["params"])
|
||||
with patch("apps.ai.free_video.submit_free_video") as submit:
|
||||
response = self.client.post(
|
||||
f"/api/ai/creations/{self.conversation.id}/send/",
|
||||
{"kind": "confirm", "reply_to": str(self.card.id),
|
||||
"params": {"duration": "15 秒", "model": "Seedance 2.0 Fast",
|
||||
"resolution": "480p", "ratio": "1:1"}},
|
||||
format="json",
|
||||
)
|
||||
self.assertEqual(response.status_code, 200)
|
||||
self.assertTrue(response.json()["regenerate"])
|
||||
submit.assert_not_called()
|
||||
self.conversation.refresh_from_db()
|
||||
self.assertEqual(self.conversation.params["duration"], "15 秒")
|
||||
self.card.refresh_from_db()
|
||||
self.assertTrue(self.card.payload["submitted"])
|
||||
|
||||
def test_confirm_applies_model_without_rewriting_script(self):
|
||||
self.conversation.params = {
|
||||
"model": "Seedance 2.0 Fast", "resolution": "480p",
|
||||
"ratio": "1:1", "duration": "8 秒",
|
||||
}
|
||||
self.conversation.save(update_fields=["params"])
|
||||
provider = ModelProvider.objects.create(name="fk-param", display_name="F", base_url="https://x")
|
||||
model = ModelConfig.objects.create(
|
||||
provider=provider, name="fk-video-param", 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-param",
|
||||
)
|
||||
with patch("apps.ai.free_video.submit_free_video", return_value=task) as submit:
|
||||
response = self.client.post(
|
||||
f"/api/ai/creations/{self.conversation.id}/send/",
|
||||
{"kind": "confirm", "reply_to": str(self.card.id),
|
||||
"params": {"duration": "8 秒", "model": "Seedance 2.5",
|
||||
"resolution": "720p", "ratio": "1:1"}},
|
||||
format="json",
|
||||
)
|
||||
self.assertEqual(response.status_code, 201)
|
||||
self.assertFalse(response.json().get("regenerate"))
|
||||
self.assertEqual(submit.call_args.kwargs["params"]["model"], "doubao-seedance-2-5-260628")
|
||||
self.assertEqual(submit.call_args.kwargs["params"]["resolution"], "720p")
|
||||
self.assertEqual(submit.call_args.kwargs["params"]["duration"], 8)
|
||||
|
||||
|
||||
class MemoryCompressionTests(CreationAgentBaseTests):
|
||||
"""长会话记忆压缩(契约 §5)。"""
|
||||
@@ -731,3 +812,60 @@ class PresetGuidanceTests(CreationAgentBaseTests):
|
||||
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()))
|
||||
|
||||
|
||||
class ChatVisionTests(CreationAgentBaseTests):
|
||||
def setUp(self):
|
||||
super().setUp()
|
||||
self.person = Asset.objects.create(
|
||||
team=self.team, created_by=self.user, name="参考人物",
|
||||
asset_type=Asset.Type.IMAGE, source=Asset.Source.UPLOAD,
|
||||
category=Asset.Category.PERSON,
|
||||
)
|
||||
AssetFile.objects.create(
|
||||
asset=self.person, object_key="k/person", bucket="b", is_primary=True,
|
||||
preview_url="https://cdn/person.jpg",
|
||||
)
|
||||
self.conversation.pinned_refs = [
|
||||
{"type": "character", "id": str(self.person.id), "name": "参考人物"},
|
||||
]
|
||||
self.conversation.save(update_fields=["pinned_refs"])
|
||||
append_message(self.conversation, role="user", text="按这个人出图")
|
||||
|
||||
def test_seed_chat_model_receives_reference_images(self):
|
||||
self.model.name = "doubao-seed-2-0-pro-260215"
|
||||
self.model.save(update_fields=["name"])
|
||||
context = AgentContext(conversation=self.conversation, user=self.user, model_config=self.model)
|
||||
messages = build_messages(context)
|
||||
user_messages = [m for m in messages if m.get("role") == "user"]
|
||||
self.assertTrue(user_messages)
|
||||
content = user_messages[-1]["content"]
|
||||
self.assertIsInstance(content, list)
|
||||
urls = [
|
||||
(item.get("image_url") or {}).get("url")
|
||||
for item in content
|
||||
if isinstance(item, dict) and item.get("type") == "image_url"
|
||||
]
|
||||
self.assertIn("https://cdn/person.jpg", urls)
|
||||
blob = " ".join(str(item.get("text") or "") for item in content if isinstance(item, dict))
|
||||
self.assertIn("性别", blob)
|
||||
|
||||
def test_plain_text_model_does_not_receive_images(self):
|
||||
context = AgentContext(conversation=self.conversation, user=self.user, model_config=self.model)
|
||||
messages = build_messages(context)
|
||||
for message in messages:
|
||||
self.assertIsInstance(message.get("content"), str)
|
||||
|
||||
|
||||
class CreationChatModelTests(CreationAgentBaseTests):
|
||||
def test_prefers_seed_21_as_chat_model(self):
|
||||
self.model.is_default = False
|
||||
self.model.save(update_fields=["is_default"])
|
||||
seed = ModelConfig.objects.create(
|
||||
provider=self.model.provider, name="doubao-seed-2-1-pro-260628",
|
||||
display_name="Doubao-Seed-2.1-Pro",
|
||||
capability=ModelConfig.Capability.TEXT, endpoint="chat/completions",
|
||||
is_default=True,
|
||||
)
|
||||
picked = get_creation_chat_model(None)
|
||||
self.assertEqual(picked.id, seed.id)
|
||||
|
||||
@@ -112,7 +112,7 @@ class ResolveRefsTests(TestCase):
|
||||
self.assertEqual(entry["review_status"], "active")
|
||||
self.assertEqual(entry["review_remote_id"], "R-1")
|
||||
|
||||
def test_model_ref_prefers_triview_over_portrait(self):
|
||||
def test_model_ref_includes_portrait_and_triview(self):
|
||||
portrait = _image_asset(self.team, self.user, "形象图", Asset.Category.MODEL_PORTRAIT, url="https://cdn/p.jpg")
|
||||
triview = _image_asset(self.team, self.user, "三视图", Asset.Category.TRI_VIEW, url="https://cdn/t.jpg")
|
||||
model = Model.objects.create(
|
||||
@@ -121,8 +121,7 @@ class ResolveRefsTests(TestCase):
|
||||
)
|
||||
|
||||
resolved = resolve_refs(self.team, [{"type": "model", "id": str(model.id)}])
|
||||
# 三视图信息量最大,锁脸优先用它
|
||||
self.assertEqual(resolved.references[0]["url"], "https://cdn/t.jpg")
|
||||
self.assertEqual([r["url"] for r in resolved.references], ["https://cdn/p.jpg", "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")
|
||||
@@ -153,6 +152,51 @@ class ResolveRefsTests(TestCase):
|
||||
# 编号错位会让 @图N 指错,必须去重
|
||||
self.assertEqual(len(resolved.references), 1)
|
||||
|
||||
def test_character_gender_in_metadata_becomes_fact(self):
|
||||
self.person.metadata = {"gender": "女", "age": "25-30"}
|
||||
self.person.save(update_fields=["metadata"])
|
||||
resolved = resolve_refs(self.team, [{"type": "character", "id": str(self.person.id)}])
|
||||
self.assertIn("性别:女", resolved.facts_text)
|
||||
self.assertIn("年龄:25-30", resolved.facts_text)
|
||||
|
||||
def test_model_gender_in_metadata_becomes_fact(self):
|
||||
portrait = _image_asset(self.team, self.user, "形象图3", Asset.Category.MODEL_PORTRAIT, url="https://cdn/p3.jpg")
|
||||
model = Model.objects.create(
|
||||
team=self.team, created_by=self.user, name="阿岚",
|
||||
portrait_asset=portrait, metadata={"gender": "男"},
|
||||
)
|
||||
resolved = resolve_refs(self.team, [{"type": "model", "id": str(model.id)}])
|
||||
self.assertIn("性别:男", resolved.facts_text)
|
||||
|
||||
def test_character_triview_is_attached_when_paired(self):
|
||||
self.person.metadata = {}
|
||||
self.person.save(update_fields=["metadata"])
|
||||
_image_asset(
|
||||
self.team, self.user, "白领三视图", Asset.Category.TRI_VIEW,
|
||||
url="https://cdn/person-tri.jpg",
|
||||
metadata={"triview_of": str(self.person.id)},
|
||||
)
|
||||
resolved = resolve_refs(self.team, [{"type": "character", "id": str(self.person.id)}])
|
||||
urls = [r["url"] for r in resolved.references]
|
||||
self.assertIn("https://cdn/person.jpg", urls)
|
||||
self.assertIn("https://cdn/person-tri.jpg", urls)
|
||||
|
||||
def test_missing_triview_does_not_mark_ref_missing(self):
|
||||
resolved = resolve_refs(self.team, [{"type": "character", "id": str(self.person.id)}])
|
||||
self.assertEqual(resolved.missing, [])
|
||||
self.assertEqual([r["url"] for r in resolved.references], ["https://cdn/person.jpg"])
|
||||
|
||||
def test_product_triview_is_attached_when_present(self):
|
||||
_image_asset(
|
||||
self.team, self.user, "商品三视图", Asset.Category.PRODUCT_IMAGE,
|
||||
url="https://cdn/prod-tri.jpg",
|
||||
metadata={"product_id": str(self.product.id), "view": "three_view"},
|
||||
)
|
||||
resolved = resolve_refs(self.team, [{"type": "product", "id": str(self.product.id)}])
|
||||
urls = [r["url"] for r in resolved.references]
|
||||
self.assertIn("https://cdn/prod.jpg", urls)
|
||||
self.assertIn("https://cdn/prod-tri.jpg", urls)
|
||||
|
||||
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))
|
||||
|
||||
@@ -242,7 +242,7 @@ class EntityExtractionRoutingTests(TestCase):
|
||||
older = self.model(self.provider("volcengine-old", 20), "doubao-seed-2-0-pro-260215")
|
||||
default = self.model(
|
||||
self.provider("volcengine-default", 10),
|
||||
"deepseek-v4-pro",
|
||||
"doubao-seed-2-1-pro-260628",
|
||||
is_default=True,
|
||||
)
|
||||
task = self.submit()
|
||||
@@ -254,5 +254,5 @@ class EntityExtractionRoutingTests(TestCase):
|
||||
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
|
||||
self.assertEqual(attempt.model_config_id, default.id)
|
||||
self.assertNotEqual(attempt.model_config_id, older.id)
|
||||
self.assertEqual(task.request_payload["model"], "deepseek-v4-pro")
|
||||
self.assertEqual(task.request_payload["model"], "doubao-seed-2-1-pro-260628")
|
||||
|
||||
|
||||
@@ -21,7 +21,7 @@ from apps.products.models import Product
|
||||
|
||||
from .generation_errors import classify_generation_error, public_error_for_task
|
||||
from .creation import append_message, sync_generating_messages
|
||||
from .creation_agent import apply_session_params, stream_creation_agent, submit_confirmed_video
|
||||
from .creation_agent import apply_confirm_params, apply_session_params, stream_creation_agent, submit_confirmed_image, 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 (
|
||||
@@ -693,12 +693,14 @@ def _free_video_list_queryset(team, *, include_replace=False):
|
||||
)
|
||||
if include_replace:
|
||||
return qs.filter(video_replace_q())
|
||||
return qs.exclude(video_replace_q())
|
||||
# 全能创作也走 FREE_VIDEO 任务类型,但不能出现在自由生成任务流里。
|
||||
return qs.exclude(video_replace_q()).exclude(request_payload__feature="omni_create")
|
||||
|
||||
|
||||
def _free_video_trash_queryset(team):
|
||||
return (
|
||||
AITask.objects.filter(team=team, task_type=AITask.Type.FREE_VIDEO, is_deleted=True, purged_at__isnull=True)
|
||||
.exclude(request_payload__feature="omni_create")
|
||||
.select_related("model_config")
|
||||
.prefetch_related("generated_assets", "generated_assets__files")
|
||||
)
|
||||
@@ -1244,9 +1246,15 @@ class FreeVideoUploadView(APIView):
|
||||
|
||||
|
||||
class ModelConfigViewSet(ReadOnlyModelViewSet):
|
||||
# 按创建序固定排序:最早创建的 active 模型排第一 = 前端选择器默认项,与 get_default_model 口径一致
|
||||
# (否则 DB 默认序不稳定,可能默认选到 Gemini 等;用户要默认 = 豆包 2.0 Pro,它最早创建)
|
||||
queryset = ModelConfig.objects.select_related("provider").filter(status=ModelConfig.Status.ACTIVE).order_by("created_at")
|
||||
# 按创建序固定排序:最早创建的 active 模型排第一 = 前端选择器默认项,与 get_default_model 口径一致。
|
||||
# DeepSeek 已停用,下拉里不再出现。
|
||||
queryset = (
|
||||
ModelConfig.objects.select_related("provider")
|
||||
.filter(status=ModelConfig.Status.ACTIVE)
|
||||
.exclude(name__icontains="deepseek")
|
||||
.exclude(display_name__icontains="deepseek")
|
||||
.order_by("created_at")
|
||||
)
|
||||
serializer_class = ModelConfigSerializer
|
||||
search_fields = ["name", "display_name", "capability"]
|
||||
ordering_fields = ["created_at", "display_name"]
|
||||
@@ -1360,9 +1368,28 @@ class CreationConversationViewSet(TeamScopedViewSetMixin, ModelViewSet):
|
||||
if (card.payload or {}).get("submitted"):
|
||||
# 确认闸门是一次性的:连点两下会出两条片、扣两次积分
|
||||
return JsonResponse({"detail": "这条方案已经确认过了"}, status=409)
|
||||
card.payload = {**(card.payload or {}), "submitted": True}
|
||||
incoming = request.data.get("params")
|
||||
if incoming is not None and not isinstance(incoming, dict):
|
||||
return JsonResponse({"detail": "params 必须是对象"}, status=400)
|
||||
latest_params, duration_changed = apply_confirm_params(
|
||||
conversation, incoming if isinstance(incoming, dict) else None
|
||||
)
|
||||
card.payload = {
|
||||
**(card.payload or {}),
|
||||
"submitted": True,
|
||||
"params": latest_params,
|
||||
}
|
||||
card.save(update_fields=["payload", "updated_at"])
|
||||
message, error = submit_confirmed_video(
|
||||
# 改时长会让旧脚本对不上(5 秒方案不能直接出 10 秒)。确认卡作废,前端再发一轮让模型重写。
|
||||
if duration_changed:
|
||||
return JsonResponse({
|
||||
"regenerate": True,
|
||||
"params": latest_params,
|
||||
"message": None,
|
||||
}, status=200)
|
||||
is_image = (card.payload or {}).get("kind") == "image" or conversation.mode == CreationConversation.Mode.IMAGE
|
||||
submitter = submit_confirmed_image if is_image else submit_confirmed_video
|
||||
message, error = submitter(
|
||||
conversation=conversation, user=request.user, confirm_message=card
|
||||
)
|
||||
if error:
|
||||
|
||||
@@ -57,13 +57,13 @@ def _tab_q(tab: str) -> Q:
|
||||
if tab == "image_creations": # 图片自由创作
|
||||
return Q(category="free_create", asset_type="image")
|
||||
if tab == "video_creations": # 视频自由创作
|
||||
return Q(category="free_create", asset_type="video")
|
||||
return Q(category="free_create", asset_type="video") & ~Q(metadata__feature="omni_create")
|
||||
if tab == "creations": # 自由创作(兼容旧入口:图片+视频)
|
||||
return Q(category="free_create")
|
||||
return Q(category="free_create") & ~Q(metadata__feature="omni_create")
|
||||
if tab == "uploads":
|
||||
return Q(category="upload")
|
||||
if tab == "materials": # 素材(期3):视频素材 = 所有视频,排除「最终成片」(final_video 隐藏不列)
|
||||
return ~Q(category="final_video") & Q(asset_type="video")
|
||||
return ~Q(category="final_video") & Q(asset_type="video") & ~Q(metadata__feature="omni_create")
|
||||
if tab == "others": # 其他(资产库成品化):我的上传 + 未归类非视频(兜底)
|
||||
return Q(category="upload") | (~Q(category__in=_KNOWN_CATS) & ~Q(asset_type="video"))
|
||||
if tab == "unclassified": # 未归类且非视频(也不含最终成片)
|
||||
|
||||
@@ -104,10 +104,7 @@ def _script_generation_inflight(project: Project) -> bool:
|
||||
|
||||
|
||||
def get_quick_script_model() -> ModelConfig | None:
|
||||
"""极速成片固定复用专业创作的豆包 Seed 2.1 Pro 脚本模型。
|
||||
|
||||
不回退到默认文本模型,避免默认配置切到 DeepSeek 后两条创作链路的成片质量不一致。
|
||||
"""
|
||||
"""极速成片固定复用专业创作的豆包 Seed 2.1 Pro 脚本模型。"""
|
||||
return (
|
||||
ModelConfig.objects.select_related("provider")
|
||||
.filter(
|
||||
|
||||
@@ -31,6 +31,7 @@ from apps.ai.services import (
|
||||
generate_base_asset,
|
||||
generate_person_triview,
|
||||
get_default_model,
|
||||
resolve_text_model,
|
||||
get_inflight_extraction,
|
||||
poll_video_segment,
|
||||
regenerate_script_segment,
|
||||
@@ -526,7 +527,7 @@ class ProjectViewSet(TeamScopedViewSetMixin, ModelViewSet):
|
||||
if total_duration not in {15, 30, 45, 60}:
|
||||
return Response({"detail": "视频时长仅支持15、30、45或60秒"}, status=status.HTTP_400_BAD_REQUEST)
|
||||
|
||||
# 极速成片和专业创作使用同一款豆包脚本模型;绝不因默认模型变化而退回 DeepSeek。
|
||||
# 极速成片和专业创作使用同一款豆包 Seed 2.1 Pro 脚本模型。
|
||||
from .services.quick_create import get_quick_script_model
|
||||
|
||||
text_model = get_quick_script_model()
|
||||
@@ -890,16 +891,7 @@ class ProjectViewSet(TeamScopedViewSetMixin, ModelViewSet):
|
||||
except (TypeError, ValueError):
|
||||
target_index = None
|
||||
|
||||
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()
|
||||
)
|
||||
if model_config is None:
|
||||
model_config = get_default_model(ModelConfig.Capability.TEXT)
|
||||
model_config = resolve_text_model(requested_id=request.data.get("model_config_id"))
|
||||
if model_config is None:
|
||||
# 纯 Django 响应:绕开 DRF 渲染(此 action 只挂了 SSE renderer)
|
||||
return JsonResponse({"detail": "没有可用的文本模型,请先在模型库配置"}, status=400)
|
||||
|
||||
Reference in New Issue
Block a user