添加全能创作功能

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
+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})