添加全能创作功能
This commit is contained in:
@@ -2,6 +2,7 @@ import logging
|
||||
import uuid
|
||||
|
||||
from django.db import transaction
|
||||
from django.http import JsonResponse, StreamingHttpResponse
|
||||
from django.db.models import Count, Exists, OuterRef, Q
|
||||
from django.utils import timezone
|
||||
from rest_framework import status
|
||||
@@ -14,14 +15,20 @@ from rest_framework.viewsets import ModelViewSet, ReadOnlyModelViewSet
|
||||
|
||||
from apps.assets.models import Asset
|
||||
from apps.assets.serializers import AssetFileSerializer, AssetSerializer
|
||||
from apps.common.api import TeamScopedViewSetMixin, get_current_team
|
||||
from apps.common.api import ServerSentEventRenderer, TeamScopedViewSetMixin, get_current_team
|
||||
from apps.common.celery_health import require_worker, require_worker_task
|
||||
from apps.products.models import Product
|
||||
|
||||
from .generation_errors import classify_generation_error, public_error_for_task
|
||||
from .models import AITask, ImageConversation, ModelConfig
|
||||
from .creation import append_message, sync_generating_messages
|
||||
from .creation_agent import apply_session_params, stream_creation_agent, submit_confirmed_video
|
||||
from .mentions import TYPE_LABELS, VALID_TYPES, refs_from_elicit_answers, search_mentions
|
||||
from .models import AITask, CreationConversation, CreationMessage, ImageConversation, ModelConfig
|
||||
from .serializers import (
|
||||
AITaskSerializer,
|
||||
CreationConversationDetailSerializer,
|
||||
CreationConversationSerializer,
|
||||
CreationMessageSerializer,
|
||||
ImageConversationSerializer,
|
||||
ImageConversationTrashSerializer,
|
||||
ModelConfigSerializer,
|
||||
@@ -1244,3 +1251,218 @@ class ModelConfigViewSet(ReadOnlyModelViewSet):
|
||||
search_fields = ["name", "display_name", "capability"]
|
||||
ordering_fields = ["created_at", "display_name"]
|
||||
|
||||
|
||||
|
||||
class CreationConversationViewSet(TeamScopedViewSetMixin, ModelViewSet):
|
||||
"""全能创作会话 CRUD(契约 §3)。
|
||||
|
||||
list 创作历史页,按 ?status=running|completed 过滤,-last_active_at 倒序
|
||||
retrieve 进对话页,一次性带回全部消息
|
||||
create 从首页「开始创作」发起,带 mode/preset/params;首条用户消息由 messages 接口发
|
||||
partial_update 只允许改 title(重命名)
|
||||
destroy 软删(不连带删已生成的资产 —— 图还在资产库里)
|
||||
|
||||
发消息走 POST {id}/messages/(SSE),不在这里。
|
||||
"""
|
||||
|
||||
serializer_class = CreationConversationSerializer
|
||||
queryset = CreationConversation.objects.order_by("-last_active_at")
|
||||
|
||||
def get_serializer_class(self):
|
||||
if self.action == "retrieve":
|
||||
return CreationConversationDetailSerializer
|
||||
return super().get_serializer_class()
|
||||
|
||||
def get_queryset(self):
|
||||
queryset = super().get_queryset().filter(is_deleted=False, purged_at__isnull=True)
|
||||
if self.action == "retrieve":
|
||||
queryset = queryset.prefetch_related("messages")
|
||||
else:
|
||||
queryset = queryset.annotate(_message_count=Count("messages"))
|
||||
mode = self.request.query_params.get("mode", "").strip()
|
||||
if mode:
|
||||
if mode not in CreationConversation.Mode.values:
|
||||
raise ValidationError({"detail": "mode 仅支持 video / image"})
|
||||
queryset = queryset.filter(mode=mode)
|
||||
conv_status = self.request.query_params.get("status", "").strip()
|
||||
if conv_status:
|
||||
if conv_status not in CreationConversation.Status.values:
|
||||
raise ValidationError({"detail": "status 仅支持 running / completed / failed"})
|
||||
queryset = queryset.filter(status=conv_status)
|
||||
return queryset
|
||||
|
||||
def perform_destroy(self, instance):
|
||||
# 只软删会话本身。已生成的图/视频留在资产库 —— 用户删对话不等于要删素材。
|
||||
instance.is_deleted = True
|
||||
instance.save(update_fields=["is_deleted", "updated_at"])
|
||||
|
||||
def _synced_conversation(self):
|
||||
"""拉会话前先把已结束的 GENERATING 回填成 RESULT/ERROR。
|
||||
|
||||
出图在 worker 里跑,agent 只提交。前端轮询 GET 本接口拿结果。
|
||||
prefetch 缓存里是回填前的旧对象,改过必须丢掉再读。
|
||||
"""
|
||||
conversation = self.get_object()
|
||||
if sync_generating_messages(conversation):
|
||||
conversation.refresh_from_db()
|
||||
cache = getattr(conversation, "_prefetched_objects_cache", None)
|
||||
if cache is not None:
|
||||
cache.pop("messages", None)
|
||||
return conversation
|
||||
|
||||
def retrieve(self, request, *args, **kwargs):
|
||||
conversation = self._synced_conversation()
|
||||
serializer = self.get_serializer(conversation)
|
||||
return Response(serializer.data)
|
||||
|
||||
@action(detail=True, methods=["get"], url_path="messages")
|
||||
def messages(self, request, pk=None):
|
||||
"""按 ?after_seq= 增量拉消息。轮询视频结果时前端只补新的,不重拉整条会话。"""
|
||||
conversation = self._synced_conversation()
|
||||
queryset = conversation.messages.all()
|
||||
after_seq = request.query_params.get("after_seq", "").strip()
|
||||
if after_seq:
|
||||
try:
|
||||
queryset = queryset.filter(seq__gt=int(after_seq))
|
||||
except (TypeError, ValueError) as exc:
|
||||
raise ValidationError({"detail": "after_seq 必须是整数"}) from exc
|
||||
return Response(CreationMessageSerializer(queryset, many=True).data)
|
||||
|
||||
@action(
|
||||
detail=True, methods=["post"], url_path="send",
|
||||
renderer_classes=[ServerSentEventRenderer],
|
||||
)
|
||||
def send(self, request, pk=None):
|
||||
"""发一条消息 → SSE 流(契约 §3)。
|
||||
|
||||
kind=text 普通发言,text + refs
|
||||
kind=elicit_answer 回答追问卡,reply_to + answers
|
||||
kind=confirm 点确认闸门 → **不跑模型**,直接按方案卡存的 video_prompt 出片
|
||||
|
||||
响应 text/event-stream。**必须挂 ServerSentEventRenderer,否则 DRF 内容协商直接 406。**
|
||||
"""
|
||||
conversation = self.get_object()
|
||||
# 模型/比例/分辨率/时长在新建会话时锁定,发送和确认出片都按当时那套,
|
||||
# 否则 5 秒方案被改成 10 秒再出片会对不上。
|
||||
kind = str(request.data.get("kind") or "text")
|
||||
text = str(request.data.get("text") or "").strip()
|
||||
refs = request.data.get("refs") or []
|
||||
if not isinstance(refs, list):
|
||||
return JsonResponse({"detail": "refs 必须是数组"}, status=400)
|
||||
|
||||
if kind == "confirm":
|
||||
reply_to = str(request.data.get("reply_to") or "").strip()
|
||||
card = conversation.messages.filter(
|
||||
id=reply_to, kind=CreationMessage.Kind.CONFIRM
|
||||
).first() if reply_to else None
|
||||
if card is None:
|
||||
return JsonResponse({"detail": "确认卡不存在"}, status=404)
|
||||
if (card.payload or {}).get("submitted"):
|
||||
# 确认闸门是一次性的:连点两下会出两条片、扣两次积分
|
||||
return JsonResponse({"detail": "这条方案已经确认过了"}, status=409)
|
||||
card.payload = {**(card.payload or {}), "submitted": True}
|
||||
card.save(update_fields=["payload", "updated_at"])
|
||||
message, error = submit_confirmed_video(
|
||||
conversation=conversation, user=request.user, confirm_message=card
|
||||
)
|
||||
if error:
|
||||
# 出片没提交成功 → 把闸门放回去,用户可以改完再确认
|
||||
card.payload = {**(card.payload or {}), "submitted": False}
|
||||
card.save(update_fields=["payload", "updated_at"])
|
||||
failure = append_message(
|
||||
conversation, role="assistant",
|
||||
kind=CreationMessage.Kind.ERROR, text=error,
|
||||
)
|
||||
return JsonResponse(
|
||||
{"detail": error, "message": CreationMessageSerializer(failure).data}, status=400
|
||||
)
|
||||
# 纯 Django 响应:这个 action 只挂了 SSE renderer,走 DRF Response 会渲染失败
|
||||
return JsonResponse(
|
||||
{"message": CreationMessageSerializer(message).data}, status=201
|
||||
)
|
||||
|
||||
if kind == "elicit_answer":
|
||||
reply_to = str(request.data.get("reply_to") or "").strip()
|
||||
answers = request.data.get("answers")
|
||||
if not reply_to or not isinstance(answers, dict):
|
||||
return JsonResponse({"detail": "回答追问需要 reply_to 与 answers"}, status=400)
|
||||
card = conversation.messages.filter(
|
||||
id=reply_to, kind=CreationMessage.Kind.ELICIT
|
||||
).first()
|
||||
if card is None:
|
||||
return JsonResponse({"detail": "追问卡不存在"}, status=404)
|
||||
if (card.payload or {}).get("submitted"):
|
||||
# 追问卡是一次性的:重复提交会让同一个问题在上下文里出现两次答案
|
||||
return JsonResponse({"detail": "这个问题已经回答过了"}, status=409)
|
||||
payload = dict(card.payload or {})
|
||||
payload["answers"] = answers
|
||||
payload["submitted"] = True
|
||||
card.payload = payload
|
||||
card.save(update_fields=["payload", "updated_at"])
|
||||
labels = {f["key"]: f["label"] for f in payload.get("fields", [])}
|
||||
text = ";".join(
|
||||
f"{labels.get(k, k)}:{'、'.join(v) if isinstance(v, list) else v}"
|
||||
for k, v in answers.items()
|
||||
)
|
||||
if apply_session_params(conversation, payload.get("fields") or [], answers):
|
||||
text = f"{text}。请按新的会话参数重新写方案,旧方案作废"
|
||||
# 点选商品/角色必须钉成 Ref:模型常把选项做成单选文字,前端只回 answers。
|
||||
refs = list(refs)
|
||||
existing = {(item.get("type"), str(item.get("id"))) for item in refs if isinstance(item, dict)}
|
||||
for extra in refs_from_elicit_answers(conversation.team, payload.get("fields") or [], answers):
|
||||
mark = (extra.get("type"), str(extra.get("id")))
|
||||
if mark in existing:
|
||||
continue
|
||||
refs.append(extra)
|
||||
existing.add(mark)
|
||||
elif not text and not refs:
|
||||
return JsonResponse({"detail": "消息不能为空"}, status=400)
|
||||
|
||||
model_config = None
|
||||
requested = request.data.get("model_config_id")
|
||||
if requested:
|
||||
model_config = (
|
||||
ModelConfig.objects.select_related("provider")
|
||||
.filter(id=requested, capability=ModelConfig.Capability.TEXT, status=ModelConfig.Status.ACTIVE)
|
||||
.first()
|
||||
)
|
||||
|
||||
stream = stream_creation_agent(
|
||||
conversation=conversation,
|
||||
user=request.user,
|
||||
text=text,
|
||||
refs=refs,
|
||||
model_config=model_config,
|
||||
)
|
||||
response = StreamingHttpResponse(stream, content_type="text/event-stream")
|
||||
response["Cache-Control"] = "no-cache"
|
||||
response["X-Accel-Buffering"] = "no" # 关 nginx 缓冲,保证逐帧下发
|
||||
return response
|
||||
|
||||
|
||||
class MentionSearchView(APIView):
|
||||
"""@ 引用检索(契约 §3)。
|
||||
|
||||
GET /api/ai/mentions/?q=净颜&types=product,character&limit=8
|
||||
|
||||
返回 [Ref] —— 前端把它按 type 分组渲染成 @ 菜单(设计稿 .omni-mention-group)。
|
||||
**前端拿到后必须整条 Ref 存进消息的 refs 字段**,不能只把 name 拼进文本,
|
||||
否则后端取不到卖点和参考图(契约 §1)。
|
||||
"""
|
||||
|
||||
def get(self, request):
|
||||
team = get_current_team(request.user)
|
||||
q = str(request.query_params.get("q") or "").strip()
|
||||
raw_types = str(request.query_params.get("types") or "").strip()
|
||||
types = [t.strip() for t in raw_types.split(",") if t.strip()] if raw_types else None
|
||||
if types:
|
||||
unknown = [t for t in types if t not in VALID_TYPES]
|
||||
if unknown:
|
||||
raise ValidationError({"detail": f"未知引用类型:{'、'.join(unknown)}"})
|
||||
try:
|
||||
limit = int(request.query_params.get("limit") or 8)
|
||||
except (TypeError, ValueError) as exc:
|
||||
raise ValidationError({"detail": "limit 必须是整数"}) from exc
|
||||
limit = max(1, min(limit, 20))
|
||||
results = search_mentions(team, q=q, types=types, limit=limit)
|
||||
return Response({"results": results, "type_labels": TYPE_LABELS})
|
||||
|
||||
Reference in New Issue
Block a user