from pathlib import Path import uuid import requests from django.db import transaction from django.db.models import Q from django.http import StreamingHttpResponse from rest_framework import status from rest_framework.decorators import action from rest_framework.exceptions import PermissionDenied, ValidationError from rest_framework.pagination import PageNumberPagination from rest_framework.parsers import FormParser, JSONParser, MultiPartParser from rest_framework.response import Response from rest_framework.views import APIView from rest_framework.viewsets import ModelViewSet from apps.common.api import TeamScopedViewSetMixin, get_current_team from .models import Asset, AssetFile, Model from .serializers import AssetSerializer, AssetUploadSerializer, ModelLibrarySerializer from .storage import TosStorage class AssetPagination(PageNumberPagination): """允许前端用 ?page_size= 覆盖(默认 20,上限 200)。前端「取全部资产」一次大页拉完, 不再逐页串行(原来 20/页 → 资产多时每次刷新要串十几个请求,整页很慢)。""" page_size_query_param = "page_size" max_page_size = 200 # 资产库 tab → 查询条件。与前端 assetTab(library.tsx)完全一致,供服务端按页过滤/计数。 _KNOWN_CATS = ["person", "scene", "product_image", "final_video", "upload"] def _tab_q(tab: str) -> Q: if tab == "people": return Q(category="person") if tab == "scenes": return Q(category="scene") if tab == "products": return Q(category="product_image") if tab == "uploads": return Q(category="upload") if tab == "finals": # final_video,或「未归类但是视频」 return Q(category="final_video") | (~Q(category__in=_KNOWN_CATS) & Q(asset_type="video")) if tab == "unclassified": # 未归类且非视频 return ~Q(category__in=_KNOWN_CATS) & ~Q(asset_type="video") return Q() class AssetViewSet(TeamScopedViewSetMixin, ModelViewSet): # select_related 回溯链:asset → origin_task → project,供序列化器解析资产归属商品(避免 N+1)。 # ★ defer 掉 AITask 的两个巨型 JSON 列(request/response payload):列表序列化只需 project.product_id, # 不 defer 的话 select_related 会把每个资产关联 AITask 的完整 AI 请求/响应 payload 全拖出来, # 200 个资产实测 100s+(数据越多越慢);defer 后 <1s、product 仍正常解析、无额外查询。 queryset = ( Asset.objects.prefetch_related("files") .select_related("origin_task__project") .defer("origin_task__request_payload", "origin_task__response_payload") .all() ) serializer_class = AssetSerializer pagination_class = AssetPagination search_fields = ["name", "description"] ordering_fields = ["created_at", "updated_at", "name"] def get_queryset(self): """服务端过滤,支持各页按需懒加载(不再前端取全量后客户端切片)。 参数:tab(资产库分桶)/category/asset_type/source/product/q(搜索)/m_(metadata 过滤)/ordering。 默认按 -created_at 排序——无 ORDER BY 时分页会重/漏。""" qs = super().get_queryset() p = self.request.query_params if p.get("tab"): qs = qs.filter(_tab_q(p["tab"])) if p.get("category"): qs = qs.filter(category=p["category"]) if p.get("asset_type"): qs = qs.filter(asset_type=p["asset_type"]) if p.get("source"): qs = qs.filter(source=p["source"]) if p.get("product"): # 资产归属商品:独立生图写 metadata.product_id;项目内生成回溯 origin_task→project→product; # 上传的商品图经 ProductImage 关联。三路并集 = 前端 ProductDetail 的 belongsToProduct。 pid = p["product"] qs = qs.filter( Q(metadata__product_id=pid) | Q(origin_task__project__product_id=pid) | Q(product_images__product_id=pid) ).distinct() if p.get("q"): qs = qs.filter(Q(name__icontains=p["q"]) | Q(description__icontains=p["q"])) for key, val in p.items(): if key.startswith("m_") and val: qs = qs.filter(**{f"metadata__{key[2:]}": val}) ordering = p.get("ordering") or "-created_at" return qs.order_by(ordering) @action(detail=False, methods=["get"]) def summary(self, request): """资产库 tab 计数(人物/场景/商品图/成片/我的上传/未分类),供 tab 徽标——不必取全量。""" base = Asset.objects.filter(team=self.get_team()) tabs = ["people", "scenes", "products", "finals", "uploads", "unclassified"] return Response({t: base.filter(_tab_q(t)).count() for t in tabs}) @action(detail=False, methods=["get"]) def facets(self, request): """某 tab 下真实存在的筛选项(来源/类型/指定 metadata 键的取值),供下拉「只列真有的」。 参数:tab、meta_keys=gender,age,role,...(逗号分隔)。""" base = Asset.objects.filter(team=self.get_team()) if request.query_params.get("tab"): base = base.filter(_tab_q(request.query_params["tab"])) sources = sorted(s for s in base.values_list("source", flat=True).distinct() if s) kinds = sorted(k for k in base.values_list("asset_type", flat=True).distinct() if k) meta = {} for key in (k for k in request.query_params.get("meta_keys", "").split(",") if k): vals = base.values_list(f"metadata__{key}", flat=True) meta[key] = sorted({str(v) for v in vals if v not in (None, "")}) return Response({"sources": sources, "kinds": kinds, "metadata": meta}) @action(detail=True, methods=["get"], url_path="raw") def raw(self, request, pk=None): """同源流式代理资产主文件。TOS 桶未配 CORS,浏览器 JS 读不到跨域媒体数据—— 编辑器抽视频缩略图(canvas 会被 taint)/解码音频波形(fetch 被拦)都需要走这里。""" asset = self.get_object() primary = asset.files.filter(is_primary=True).first() or asset.files.first() if primary is None: return Response({"detail": "asset has no file"}, status=status.HTTP_404_NOT_FOUND) url = TosStorage().presigned_get_url(object_key=primary.object_key, expires_in=600) upstream_headers = {} if request.META.get("HTTP_RANGE"): upstream_headers["Range"] = request.META["HTTP_RANGE"] upstream = requests.get(url, headers=upstream_headers, stream=True, timeout=120) response = StreamingHttpResponse( upstream.iter_content(chunk_size=256 * 1024), status=upstream.status_code, content_type=primary.content_type or "application/octet-stream", ) for header in ("Content-Length", "Content-Range", "Accept-Ranges"): if header in upstream.headers: response[header] = upstream.headers[header] return response class ModelLibraryViewSet(ModelViewSet): """模特库(顶级实体)· 团队级 + 官方模板跨团队可见。 list 返回「本团队模特 ∪ 官方模板」;?tab=official 只看官方、?tab=mine 只看自建。 真人上传走 upload action(传图 → 建 model_portrait 资产 + Model 实体)。删除 = 软删(官方不可删)。""" serializer_class = ModelLibrarySerializer pagination_class = AssetPagination parser_classes = [JSONParser, MultiPartParser, FormParser] search_fields = ["name", "description"] def get_team(self): return get_current_team(self.request.user) def get_queryset(self): team = self.get_team() qs = ( Model.objects.filter(Q(team=team) | Q(is_official=True), is_deleted=False) .prefetch_related("portrait_asset__files", "triview_asset__files") ) tab = self.request.query_params.get("tab") if tab == "official": qs = qs.filter(is_official=True) elif tab == "mine": qs = qs.filter(team=team, is_official=False) if self.request.query_params.get("q"): q = self.request.query_params["q"] qs = qs.filter(Q(name__icontains=q) | Q(description__icontains=q)) # 官方模板靠前,再按新→旧 return qs.order_by("-is_official", "-created_at") def perform_create(self, serializer): serializer.save(team=self.get_team(), created_by=self.request.user) def perform_destroy(self, instance): if instance.is_official: raise PermissionDenied("官方模特不可删除") if instance.team_id != self.get_team().id: raise PermissionDenied("无权删除其他团队的模特") instance.is_deleted = True instance.save(update_fields=["is_deleted", "updated_at"]) @action(detail=False, methods=["post"], url_path="upload", parser_classes=[MultiPartParser, FormParser]) def upload(self, request): """真人上传:一张人像图 → 建 model_portrait 资产 + Model(source=upload)。三视图/声线后续补。""" upload = request.FILES.get("file") if upload is None: raise ValidationError({"file": "请上传一张模特形象图"}) team = self.get_team() name = (request.data.get("name") or Path(upload.name).stem or "模特")[:255] asset_id = uuid.uuid4() suffix = Path(upload.name).suffix.lower() or ".png" object_key = f"teams/{team.id}/models/{asset_id}{suffix}" with transaction.atomic(): stored = TosStorage().upload_fileobj( fileobj=upload.file, object_key=object_key, content_type=upload.content_type or "image/png", ) portrait = Asset.objects.create( id=asset_id, team=team, created_by=request.user, name=f"{name}·形象图", asset_type=Asset.Type.IMAGE, source=Asset.Source.UPLOAD, category=Asset.Category.MODEL_PORTRAIT, metadata={"kind": "model", "view": "frontal", "source": "upload"}, ) AssetFile.objects.create( asset=portrait, object_key=stored.object_key, bucket=stored.bucket, content_type=stored.content_type, size_bytes=stored.size_bytes, is_primary=True, ) model = Model.objects.create( team=team, created_by=request.user, name=name, source=Model.Source.UPLOAD, portrait_asset=portrait, metadata={"source": "upload"}, ) return Response(ModelLibrarySerializer(model).data, status=status.HTTP_201_CREATED) class AssetUploadView(APIView): parser_classes = [MultiPartParser, FormParser] @transaction.atomic def post(self, request): serializer = AssetUploadSerializer(data=request.data) serializer.is_valid(raise_exception=True) team = get_current_team(request.user) upload = serializer.validated_data["file"] suffix = Path(upload.name).suffix.lower() asset_id = uuid.uuid4() object_key = f"teams/{team.id}/uploads/{asset_id}{suffix}" stored = TosStorage().upload_fileobj( fileobj=upload.file, object_key=object_key, content_type=upload.content_type or "application/octet-stream", ) asset = Asset.objects.create( id=asset_id, team=team, created_by=request.user, name=serializer.validated_data.get("name") or upload.name, asset_type=serializer.validated_data["asset_type"], source=Asset.Source.UPLOAD, category=serializer.validated_data["category"], description=serializer.validated_data.get("description", ""), ) AssetFile.objects.create( asset=asset, object_key=stored.object_key, bucket=stored.bucket, content_type=stored.content_type, size_bytes=stored.size_bytes, is_primary=True, ) return Response(AssetSerializer(asset).data, status=status.HTTP_201_CREATED)