from rest_framework import status from rest_framework.response import Response from rest_framework.views import APIView from rest_framework.viewsets import ReadOnlyModelViewSet from apps.assets.serializers import AssetSerializer from apps.common.api import TeamScopedViewSetMixin, get_current_team from apps.common.celery_health import require_worker from .models import AITask, ModelConfig from .serializers import AITaskSerializer, ModelConfigSerializer from .services import enqueue_standalone_images class GenerateImageView(APIView): """独立生图(不绑项目)· 图片创作/模特图/平台套图共用 —— **异步**。 POST /api/ai/generate-image/ 提交生成,秒级返回 RESERVED 任务列表(慢出图交给 worker)。 GET /api/ai/generate-image/?ids=… 轮询这些任务的状态;成功的任务带回成图 asset。 """ def post(self, request): require_worker() # 异步出图依赖 worker 兜底执行,没 worker 直接拒绝(否则任务永远 RESERVED) prompt = str(request.data.get("prompt") or "").strip() if not prompt: return Response({"detail": "prompt 不能为空"}, status=status.HTTP_400_BAD_REQUEST) mode = str(request.data.get("mode") or "image") try: count = int(request.data.get("count") or 1) except (TypeError, ValueError): count = 1 team = get_current_team(request.user) try: tasks = enqueue_standalone_images(team=team, user=request.user, prompt=prompt, mode=mode, count=count) except ValueError as exc: # 无可用模型 / 余额不足等,立即反馈 return Response({"detail": str(exc)}, status=status.HTTP_400_BAD_REQUEST) return Response( {"tasks": [{"id": str(t.id), "status": t.status} for t in tasks]}, status=status.HTTP_202_ACCEPTED, ) def get(self, request): team = get_current_team(request.user) ids = [s for s in str(request.query_params.get("ids") or "").split(",") if s] if not ids: return Response({"tasks": []}) tasks = AITask.objects.filter(team=team, id__in=ids).prefetch_related( "generated_assets", "generated_assets__files" ) data = [ { "id": str(t.id), "status": t.status, "error_message": t.error_message, "assets": AssetSerializer( [a for a in t.generated_assets.all() if not a.is_deleted], many=True ).data, } for t in tasks ] return Response({"tasks": data}) class AITaskViewSet(TeamScopedViewSetMixin, ReadOnlyModelViewSet): # 序列化器不含 request_payload/response_payload(单条可达 3MB+ base64 图),defer 掉: # 否则只为序列化 14 个小字段也会把几十 MB blob 从库里拉回(远程库实测 40 条要 30s+)。 queryset = AITask.objects.select_related("team", "project", "model_config", "model_config__provider").defer("request_payload", "response_payload").all() serializer_class = AITaskSerializer search_fields = ["idempotency_key", "provider_task_id", "project__name"] ordering_fields = ["created_at", "updated_at", "completed_at"] class ModelConfigViewSet(ReadOnlyModelViewSet): queryset = ModelConfig.objects.select_related("provider").filter(status=ModelConfig.Status.ACTIVE) serializer_class = ModelConfigSerializer search_fields = ["name", "display_name", "capability"] ordering_fields = ["created_at", "display_name"]