生图模型可选(火山/gpt-image)+ 模特上身图提示词强化 + 多会话等改动
本轮(生图模型选择 + 火山接入): - 工作室新增「生图模型」选择器(模特上身图/平台套图头部 chip + 图片创作底部 Pill), 默认火山 Seedream,可切 gpt-image-2;选择写入 localStorage,下次进页面读回 - 后端 resolve_image_model 解析所选模型;enqueue_standalone_images 接 image_model - worker 按模型能力分流:有 image_edit(gpt-image)走多图编辑;无(火山)走 image_generation(image=参考图);新增 _ratio_to_volcano_size 让火山按比例出图 → 内衣等敏感品类用火山可绕开 gpt-image 的 sexual 内容审核 模特上身图提示词: - 穿戴/非穿戴分流、按 index 变化动作场景镜头、负面词尾接、多图参考序号自适应 - _product_reference_urls:商品参考真实上传图优先、排除 AI 生成图、可多张 其他(并入此前各会话未提交改动): - 图片创作多会话(ImageConversation + migration 0019)、任务中心按类型过滤 - accounts/projects/assets/team/auth 等零散调整、相关测试 - 测试脚本(tryon_*.py)、测试清单(core/bug/*.xlsx)
This commit is contained in:
@@ -126,6 +126,65 @@ class InvitationFlowTests(TestCase):
|
||||
self.assertEqual(gen2.status_code, 403)
|
||||
|
||||
|
||||
class MemberQuotaAndRoleTests(TestCase):
|
||||
"""PMC#3:① 新建子账号不填限额时默认 100(不再 0=不限,避免子账号无上限花团队池);
|
||||
② login/me 返回当前用户在团队的 role,前端据此只给主账号(owner)显示「团队」「消费」页。"""
|
||||
|
||||
def _register(self, client, username, **extra):
|
||||
return client.post(
|
||||
"/api/auth/register/",
|
||||
{"username": username, "password": "strong-password", **extra},
|
||||
format="json",
|
||||
)
|
||||
|
||||
def setUp(self):
|
||||
self.owner_client = APIClient()
|
||||
r = self._register(self.owner_client, "q-owner", team_name="Q Team", invite_code=make_create_team_code())
|
||||
self.assertEqual(r.status_code, 201)
|
||||
self.token = r.data["token"]
|
||||
self.team = Team.objects.get(name="Q Team")
|
||||
self.owner_client.credentials(HTTP_AUTHORIZATION=f"Token {self.token}")
|
||||
|
||||
def test_register_owner_payload_role_is_owner(self):
|
||||
# 注册即登录的返回体带 role=owner(主账号)
|
||||
r = self._register(APIClient(), "solo-owner", team_name="Solo", invite_code=make_create_team_code())
|
||||
self.assertEqual(r.data.get("role"), "owner")
|
||||
|
||||
def test_create_member_defaults_limit_100(self):
|
||||
res = self.owner_client.post(
|
||||
"/api/auth/team/members/",
|
||||
{"username": "sub-a", "password": "strong-password", "role": "member"},
|
||||
format="json",
|
||||
)
|
||||
self.assertEqual(res.status_code, 201, res.content)
|
||||
member = TeamMember.objects.get(user__username="sub-a", team=self.team)
|
||||
self.assertEqual(member.monthly_credit_limit, Decimal("100"))
|
||||
|
||||
def test_create_member_explicit_limit_respected(self):
|
||||
res = self.owner_client.post(
|
||||
"/api/auth/team/members/",
|
||||
{"username": "sub-b", "password": "strong-password", "role": "member", "monthly_credit_limit": 800},
|
||||
format="json",
|
||||
)
|
||||
self.assertEqual(res.status_code, 201, res.content)
|
||||
member = TeamMember.objects.get(user__username="sub-b", team=self.team)
|
||||
self.assertEqual(member.monthly_credit_limit, Decimal("800"))
|
||||
|
||||
def test_member_login_returns_member_role(self):
|
||||
self.owner_client.post(
|
||||
"/api/auth/team/members/",
|
||||
{"username": "sub-c", "password": "strong-password", "role": "member"},
|
||||
format="json",
|
||||
)
|
||||
login = APIClient().post(
|
||||
"/api/auth/login/",
|
||||
{"username": "sub-c", "password": "strong-password"},
|
||||
format="json",
|
||||
)
|
||||
self.assertEqual(login.status_code, 200)
|
||||
self.assertEqual(login.data.get("role"), "member")
|
||||
|
||||
|
||||
class PlatformAdminTests(TestCase):
|
||||
"""Phase 0:平台超管基础 —— 管理命令 / 权限类 / 登录 me 返回标志且团队为空。"""
|
||||
|
||||
|
||||
@@ -25,11 +25,25 @@ from .serializers import (
|
||||
)
|
||||
|
||||
|
||||
# 子账号默认月度限额(0=不限,只留给主账号/超管)
|
||||
DEFAULT_MEMBER_MONTHLY_LIMIT = 100
|
||||
|
||||
|
||||
def member_role(user, team):
|
||||
"""当前用户在当前团队的成员角色(owner/admin/member);无团队成员关系返回 ""。
|
||||
前端据此决定「团队」「消费」等仅主账号(超管)可见的页面是否展示(PMC#3)。"""
|
||||
if team is None:
|
||||
return ""
|
||||
m = TeamMember.objects.filter(team=team, user=user, status=TeamMember.Status.ACTIVE).first()
|
||||
return m.role if m else ""
|
||||
|
||||
|
||||
def auth_payload(user, team, token):
|
||||
return {
|
||||
"token": token.key,
|
||||
"user": UserSerializer(user).data,
|
||||
"team": TeamSerializer(team).data if team is not None else None,
|
||||
"role": member_role(user, team),
|
||||
}
|
||||
|
||||
|
||||
@@ -157,6 +171,7 @@ def me(request):
|
||||
{
|
||||
"user": UserSerializer(user).data,
|
||||
"team": TeamSerializer(team).data if team is not None else None,
|
||||
"role": member_role(user, team),
|
||||
}
|
||||
)
|
||||
|
||||
@@ -276,7 +291,8 @@ def team_members(request):
|
||||
team=team,
|
||||
user=user,
|
||||
role=role,
|
||||
monthly_credit_limit=request.data.get("monthly_credit_limit") or request.data.get("monthly") or 0,
|
||||
# 子账号默认月度限额 100(0=不限只留给主账号/超管)—— 避免新成员不设限就无上限地花主账号团队池的钱(PMC#3)。
|
||||
monthly_credit_limit=request.data.get("monthly_credit_limit") or request.data.get("monthly") or DEFAULT_MEMBER_MONTHLY_LIMIT,
|
||||
)
|
||||
return Response(TeamMemberSerializer(member).data, status=status.HTTP_201_CREATED)
|
||||
|
||||
@@ -299,7 +315,7 @@ def team_invitations(request):
|
||||
kind=Invitation.Kind.JOIN_TEAM, # 团管发的码恒为「加入本团队」
|
||||
role=role,
|
||||
email=str(request.data.get("email") or "").strip(),
|
||||
monthly_credit_limit=request.data.get("monthly_credit_limit") or request.data.get("monthly") or 0,
|
||||
monthly_credit_limit=request.data.get("monthly_credit_limit") or request.data.get("monthly") or DEFAULT_MEMBER_MONTHLY_LIMIT,
|
||||
created_by=request.user,
|
||||
)
|
||||
return Response(InvitationSerializer(invite).data, status=status.HTTP_201_CREATED)
|
||||
|
||||
+104
@@ -0,0 +1,104 @@
|
||||
# Generated by Django 5.1.15 on 2026-06-27 03:14
|
||||
|
||||
import django.db.models.deletion
|
||||
import uuid
|
||||
from django.conf import settings
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
("accounts", "0005_invitation_kind_alter_invitation_team"),
|
||||
("ai", "0018_default_script_gemini"),
|
||||
("products", "0002_product_products_pr_team_id_ebc2f5_idx"),
|
||||
("projects", "0006_migrate_storyboard_to_shots"),
|
||||
migrations.swappable_dependency(settings.AUTH_USER_MODEL),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.CreateModel(
|
||||
name="ImageConversation",
|
||||
fields=[
|
||||
(
|
||||
"id",
|
||||
models.UUIDField(
|
||||
default=uuid.uuid4,
|
||||
editable=False,
|
||||
primary_key=True,
|
||||
serialize=False,
|
||||
),
|
||||
),
|
||||
("created_at", models.DateTimeField(auto_now_add=True)),
|
||||
("updated_at", models.DateTimeField(auto_now=True)),
|
||||
("title", models.CharField(default="默认创作", max_length=120)),
|
||||
(
|
||||
"mode",
|
||||
models.CharField(
|
||||
choices=[
|
||||
("image", "图片创作"),
|
||||
("model", "模特上身图"),
|
||||
("cover", "平台套图"),
|
||||
],
|
||||
default="image",
|
||||
max_length=16,
|
||||
),
|
||||
),
|
||||
("is_deleted", models.BooleanField(default=False)),
|
||||
("last_active_at", models.DateTimeField(auto_now_add=True)),
|
||||
(
|
||||
"created_by",
|
||||
models.ForeignKey(
|
||||
blank=True,
|
||||
null=True,
|
||||
on_delete=django.db.models.deletion.SET_NULL,
|
||||
related_name="created_%(class)s_set",
|
||||
to=settings.AUTH_USER_MODEL,
|
||||
),
|
||||
),
|
||||
(
|
||||
"product",
|
||||
models.ForeignKey(
|
||||
blank=True,
|
||||
null=True,
|
||||
on_delete=django.db.models.deletion.SET_NULL,
|
||||
related_name="image_conversations",
|
||||
to="products.product",
|
||||
),
|
||||
),
|
||||
(
|
||||
"team",
|
||||
models.ForeignKey(
|
||||
on_delete=django.db.models.deletion.CASCADE,
|
||||
related_name="%(class)s_set",
|
||||
to="accounts.team",
|
||||
),
|
||||
),
|
||||
],
|
||||
),
|
||||
migrations.AddField(
|
||||
model_name="aitask",
|
||||
name="conversation",
|
||||
field=models.ForeignKey(
|
||||
blank=True,
|
||||
null=True,
|
||||
on_delete=django.db.models.deletion.SET_NULL,
|
||||
related_name="tasks",
|
||||
to="ai.imageconversation",
|
||||
),
|
||||
),
|
||||
migrations.AddIndex(
|
||||
model_name="aitask",
|
||||
index=models.Index(
|
||||
fields=["conversation", "created_at"],
|
||||
name="ai_aitask_convers_bbca59_idx",
|
||||
),
|
||||
),
|
||||
migrations.AddIndex(
|
||||
model_name="imageconversation",
|
||||
index=models.Index(
|
||||
fields=["team", "mode", "-last_active_at"],
|
||||
name="ai_imagecon_team_id_5dd98e_idx",
|
||||
),
|
||||
),
|
||||
]
|
||||
@@ -53,6 +53,38 @@ class ModelConfig(TimeStampedModel):
|
||||
return f"{self.provider.name}:{self.name}:{self.capability}"
|
||||
|
||||
|
||||
class ImageConversation(TeamOwnedModel):
|
||||
"""图片创作工作室的「对话」实体:一条对话 = 一个生图会话线程,把多次生成(AITask)串起来。
|
||||
左栏会话列表、切换、重命名、历史都基于本表。删除走软删(is_deleted),不连带删图——
|
||||
成图仍在资产库里。mode 区分图片创作 / 模特上身图 / 平台套图三种工作台。"""
|
||||
|
||||
class Mode(models.TextChoices):
|
||||
IMAGE = "image", "图片创作"
|
||||
MODEL = "model", "模特上身图"
|
||||
COVER = "cover", "平台套图"
|
||||
|
||||
title = models.CharField(max_length=120, default="默认创作")
|
||||
mode = models.CharField(max_length=16, choices=Mode.choices, default=Mode.IMAGE)
|
||||
product = models.ForeignKey(
|
||||
"products.Product",
|
||||
on_delete=models.SET_NULL,
|
||||
null=True,
|
||||
blank=True,
|
||||
related_name="image_conversations",
|
||||
)
|
||||
is_deleted = models.BooleanField(default=False)
|
||||
# 每次在该对话里发起生成都刷新,左栏「最近」按它倒序
|
||||
last_active_at = models.DateTimeField(auto_now_add=True)
|
||||
|
||||
class Meta:
|
||||
indexes = [
|
||||
models.Index(fields=["team", "mode", "-last_active_at"]),
|
||||
]
|
||||
|
||||
def __str__(self) -> str:
|
||||
return f"conv:{self.mode}:{self.title}"
|
||||
|
||||
|
||||
class AITask(TeamOwnedModel):
|
||||
class Type(models.TextChoices):
|
||||
SCRIPT_GENERATION = "script_generation", "Script Generation"
|
||||
@@ -84,6 +116,14 @@ class AITask(TeamOwnedModel):
|
||||
blank=True,
|
||||
related_name="ai_tasks",
|
||||
)
|
||||
# 图片创作工作室的对话归属(独立生图任务才挂;流水线内部任务为空)。删对话不删任务。
|
||||
conversation = models.ForeignKey(
|
||||
ImageConversation,
|
||||
on_delete=models.SET_NULL,
|
||||
null=True,
|
||||
blank=True,
|
||||
related_name="tasks",
|
||||
)
|
||||
task_type = models.CharField(max_length=48, choices=Type.choices)
|
||||
status = models.CharField(max_length=32, choices=Status.choices, default=Status.CREATED)
|
||||
model_config = models.ForeignKey(ModelConfig, on_delete=models.PROTECT, related_name="tasks")
|
||||
@@ -105,6 +145,8 @@ class AITask(TeamOwnedModel):
|
||||
models.Index(fields=["provider_task_id"]),
|
||||
# 任务历史默认按团队 + 创建时间倒序(AI 工具页 / asset-factory)
|
||||
models.Index(fields=["team", "-created_at"]),
|
||||
# 按对话拉历史(图片创作工作室切换会话时回填批次)
|
||||
models.Index(fields=["conversation", "created_at"]),
|
||||
]
|
||||
|
||||
def __str__(self) -> str:
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
from rest_framework import serializers
|
||||
|
||||
from .models import AITask, ModelConfig, ModelProvider
|
||||
from .models import AITask, ImageConversation, ModelConfig, ModelProvider
|
||||
|
||||
|
||||
class ModelProviderSerializer(serializers.ModelSerializer):
|
||||
@@ -19,8 +19,34 @@ class ModelConfigSerializer(serializers.ModelSerializer):
|
||||
read_only_fields = fields
|
||||
|
||||
|
||||
class ImageConversationSerializer(serializers.ModelSerializer):
|
||||
"""图片创作对话:左栏列表用。title 可写(重命名),mode/product 创建时可指定。"""
|
||||
|
||||
task_count = serializers.SerializerMethodField()
|
||||
|
||||
class Meta:
|
||||
model = ImageConversation
|
||||
fields = ["id", "title", "mode", "product", "task_count", "last_active_at", "created_at", "updated_at"]
|
||||
read_only_fields = ["id", "task_count", "last_active_at", "created_at", "updated_at"]
|
||||
|
||||
def get_task_count(self, obj) -> int:
|
||||
# list 接口已 annotate;无 annotate 时回落实时 count(详情/创建场景)
|
||||
cached = getattr(obj, "_task_count", None)
|
||||
return cached if cached is not None else obj.tasks.count()
|
||||
|
||||
|
||||
class AITaskSerializer(serializers.ModelSerializer):
|
||||
model_config = ModelConfigSerializer(read_only=True)
|
||||
# 从 request_payload 抽出的分组/标签信息(由 AITaskViewSet annotate 提供;其它调用处无此注解则为 None)。
|
||||
# 不直接读 obj.request_payload —— 那列已 defer,读了会触发懒加载把几 MB payload 整列拉回。
|
||||
batch_id = serializers.SerializerMethodField()
|
||||
mode = serializers.SerializerMethodField()
|
||||
|
||||
def get_batch_id(self, obj):
|
||||
return getattr(obj, "rp_batch_id", None)
|
||||
|
||||
def get_mode(self, obj):
|
||||
return getattr(obj, "rp_mode", None)
|
||||
|
||||
class Meta:
|
||||
model = AITask
|
||||
@@ -28,6 +54,8 @@ class AITaskSerializer(serializers.ModelSerializer):
|
||||
"id",
|
||||
"project",
|
||||
"task_type",
|
||||
"batch_id",
|
||||
"mode",
|
||||
"status",
|
||||
"model_config",
|
||||
"provider_task_id",
|
||||
@@ -40,5 +68,6 @@ class AITaskSerializer(serializers.ModelSerializer):
|
||||
"created_at",
|
||||
"updated_at",
|
||||
]
|
||||
read_only_fields = fields
|
||||
# batch_id / mode 是显式声明的 SerializerMethodField(本就只读),不能再列进 read_only_fields(DRF 会报错)
|
||||
read_only_fields = [f for f in fields if f not in ("batch_id", "mode")]
|
||||
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
import json
|
||||
import math
|
||||
import re
|
||||
import subprocess
|
||||
import tempfile
|
||||
@@ -760,6 +761,33 @@ def _ratio_to_image_size(ratio: str) -> str:
|
||||
}.get((ratio or "").strip(), "1024x1024")
|
||||
|
||||
|
||||
def _ratio_to_volcano_size(ratio: str) -> str:
|
||||
"""前端比例 → 火山 Seedream 尺寸(~2K 面积,各边夹在 [1024,4096] 且取 16 的倍数)。
|
||||
预设比例直接给好尺寸;自定义 W:H 按 2K 面积换算;解析不到回落 '2K'。"""
|
||||
presets = {
|
||||
"1:1": "2048x2048",
|
||||
"3:4": "1728x2304",
|
||||
"4:3": "2304x1728",
|
||||
"9:16": "1440x2560",
|
||||
"16:9": "2560x1440",
|
||||
"4:5": "1664x2080",
|
||||
}
|
||||
r = (ratio or "").strip()
|
||||
if r in presets:
|
||||
return presets[r]
|
||||
if ":" in r:
|
||||
try:
|
||||
w_str, h_str = r.split(":", 1)
|
||||
w, h = float(w_str), float(h_str)
|
||||
if w > 0 and h > 0:
|
||||
scale = math.sqrt((2048 * 2048) / (w * h))
|
||||
side = lambda v: min(max(int(round(v * scale / 16) * 16), 1024), 4096) # noqa: E731
|
||||
return f"{side(w)}x{side(h)}"
|
||||
except (ValueError, ZeroDivisionError):
|
||||
pass
|
||||
return "2K"
|
||||
|
||||
|
||||
def _product_cover_url(product) -> str:
|
||||
"""商品主图 URL:优先 cover_asset,其次标记为主图的商品图,再次首张商品图。无图返回 ''。"""
|
||||
if product is None:
|
||||
@@ -953,6 +981,32 @@ def build_platform_cover_prompt_refs(product, has_model: bool, base_prompt: str
|
||||
return " ".join(lines)
|
||||
|
||||
|
||||
def build_free_reference_prompt(base_prompt: str, n_refs: int = 1) -> str:
|
||||
"""图片创作自由模式 · 带用户上传参考图时的提示词。
|
||||
纯把用户原话丢给图生图,模型只会松散借个色调、不会真的保留参考图里的主体(背心/商品/人物),
|
||||
这正是「没参考我上传的素材」的根因。这里显式把参考图钉成「画面主体的唯一依据」,要求严格保留其
|
||||
外形/款式/配色/材质/品牌文字/图案,再在此基础上按用户要求创作。"""
|
||||
base = (base_prompt or "").strip()
|
||||
if n_refs <= 1:
|
||||
ref_intro = "参考图是用户提供的素材,是本次画面主体(商品 / 人物 / 物体)的唯一依据。"
|
||||
ref_word = "参考图"
|
||||
else:
|
||||
rng = f"1-{n_refs}" if n_refs > 2 else "1、2"
|
||||
ref_intro = f"参考图{rng}是用户提供的素材(同一主体的不同角度 / 多个主体),是本次画面主体的唯一依据。"
|
||||
ref_word = f"参考图{rng}"
|
||||
lines = [
|
||||
ref_intro,
|
||||
f"请严格保留{ref_word}中主体的外形、款式、配色、材质、纹样、品牌文字与 Logo,"
|
||||
"不要重新设计、不要换款、不要生成相似但不同的物体;主体须与参考图高度一致。",
|
||||
]
|
||||
if base:
|
||||
lines.append(f"在此基础上,按用户要求创作:{base}")
|
||||
else:
|
||||
lines.append("在此基础上,生成干净、专业的电商视觉画面。")
|
||||
lines.append("画面真实、构图协调、细节清晰。")
|
||||
return " ".join(lines)
|
||||
|
||||
|
||||
def build_person_frontal_prompt(description: str = "") -> str:
|
||||
"""人物正面氛围图提示词:把脚本提取(或用户输入)的人物描述包成统一模板。
|
||||
用户钦定格式:电商真人模特,氛围正面全身照,<描述>,自然妆容,柔和影棚光,真实质感,单人,纯色背景。"""
|
||||
@@ -1998,7 +2052,7 @@ def _reap_stale_standalone_image_tasks(*, team) -> None:
|
||||
continue
|
||||
|
||||
|
||||
def enqueue_standalone_images(*, team, user, prompt: str, mode: str = "image", count: int = 1, product_id: str | None = None, reference_product: bool = False, model_id: str | None = None, model_entity_id: str | None = None, ratio: str | None = None, image_model: str | None = None) -> list[AITask]:
|
||||
def enqueue_standalone_images(*, team, user, prompt: str, mode: str = "image", count: int = 1, product_id: str | None = None, reference_product: bool = False, model_id: str | None = None, model_entity_id: str | None = None, ratio: str | None = None, image_model: str | None = None, conversation=None, reference_image_ids: list[str] | None = None) -> list[AITask]:
|
||||
"""独立生图(图片创作 / 模特上身图 / 平台套图)改为**异步**:本函数在 Web 请求里只做「建任务 +
|
||||
预留额度」这种秒级的活,真正 ~30s 的 ARK 出图交给 Celery worker(generate_standalone_image_task)。
|
||||
|
||||
@@ -2014,6 +2068,7 @@ def enqueue_standalone_images(*, team, user, prompt: str, mode: str = "image", c
|
||||
raise ValueError("no active image model configured")
|
||||
task_type = _STANDALONE_TASK_TYPE.get(mode, AITask.Type.PRODUCT_IMAGE)
|
||||
count = max(1, min(int(count or 1), 12))
|
||||
ref_ids = [str(r) for r in (reference_image_ids or []) if r]
|
||||
# 本次提交 = 一组(模特上身图组 / 平台套图组):同一 batch_id 串起这批图,前端可成组展示。
|
||||
batch_id = str(uuid.uuid4())
|
||||
tasks: list[AITask] = []
|
||||
@@ -2023,11 +2078,12 @@ def enqueue_standalone_images(*, team, user, prompt: str, mode: str = "image", c
|
||||
team=team,
|
||||
created_by=user,
|
||||
project=None,
|
||||
conversation=conversation,
|
||||
task_type=task_type,
|
||||
status=AITask.Status.CREATED,
|
||||
model_config=model_config,
|
||||
idempotency_key=f"standalone-image:{team.id}:{uuid.uuid4()}",
|
||||
request_payload={"model": model_config.name, "endpoint": model_config.endpoint, "prompt": prompt, "mode": mode, "index": index, "product_id": str(product_id) if product_id else None, "reference_product": bool(reference_product), "model_id": str(model_id) if model_id else None, "model_entity_id": str(model_entity_id) if model_entity_id else None, "batch_id": batch_id, "ratio": str(ratio) if ratio else None},
|
||||
request_payload={"model": model_config.name, "endpoint": model_config.endpoint, "prompt": prompt, "mode": mode, "index": index, "product_id": str(product_id) if product_id else None, "reference_product": bool(reference_product), "model_id": str(model_id) if model_id else None, "model_entity_id": str(model_entity_id) if model_entity_id else None, "batch_id": batch_id, "ratio": str(ratio) if ratio else None, "reference_image_ids": ref_ids},
|
||||
estimated_cost=cost,
|
||||
)
|
||||
# 预留额度若余额不足会抛 ValueError,在同步的 Web 请求里立刻反馈给前端(不会先建半套任务)
|
||||
@@ -2082,6 +2138,14 @@ def run_standalone_image_task(*, task_id: str) -> None:
|
||||
product_url = _product_cover_url(product) if product is not None else ""
|
||||
# 模特上身图:真实上传图优先、排除 AI 生成图,可多张(多角度更易锁外形/品牌)
|
||||
product_urls = _product_reference_urls(product, limit=3) if product is not None else []
|
||||
# 图片创作自由模式:用户上传的参考图(已先传成 Asset)→ 取直链,作多图参考(image_edit / 图生图)
|
||||
ref_urls: list[str] = []
|
||||
for rid in (payload.get("reference_image_ids") or []):
|
||||
ref_asset = Asset.objects.filter(id=rid).first()
|
||||
if ref_asset is not None:
|
||||
u = _asset_preview_url(ref_asset)
|
||||
if u:
|
||||
ref_urls.append(u)
|
||||
|
||||
# 参考图收集与 provider 无关:先把「该用哪些参考图 + 哪条提示词」定下来,再按模型能力选调用方式。
|
||||
edit_images: list[str] = []
|
||||
@@ -2100,6 +2164,11 @@ def run_standalone_image_task(*, task_id: str) -> None:
|
||||
elif bool(payload.get("reference_product")) and product_url:
|
||||
edit_images = [product_url]
|
||||
edit_prompt = build_product_triview_prompt_refs(product, "")
|
||||
elif ref_urls:
|
||||
# 图片创作自由模式:以用户上传的参考图为基底出图。提示词必须显式要求「保留参考图主体」,
|
||||
# 否则图生图只会松散借个色调、不真的还原上传的素材(=用户反馈的「没参考我的图」)。
|
||||
edit_images = ref_urls
|
||||
edit_prompt = build_free_reference_prompt(prompt, n_refs=len(ref_urls))
|
||||
use_edit = bool(edit_images)
|
||||
try:
|
||||
if use_edit and can_edit:
|
||||
@@ -2110,8 +2179,10 @@ def run_standalone_image_task(*, task_id: str) -> None:
|
||||
size = _ratio_to_image_size(str(payload.get("ratio") or "")) # 模特图按选中比例
|
||||
response = provider.image_edit(model=model_config.name, prompt=edit_prompt, images=edit_images, size=size)
|
||||
elif use_edit:
|
||||
# 火山 Seedream 无 image_edit:走 image_generation 带 image=参考图(图生图多参考),size 用 2K
|
||||
response = provider.image_generation(model=model_config.name, endpoint=model_config.endpoint, prompt=edit_prompt, image=edit_images, size="2K")
|
||||
# 火山 Seedream 无 image_edit:走 image_generation 带 image=参考图(图生图多参考);
|
||||
# 尺寸按选中比例换算成火山可接受的 ~2K 尺寸(三视图固定横向)
|
||||
vsize = "2304x1728" if payload.get("reference_product") else _ratio_to_volcano_size(str(payload.get("ratio") or ""))
|
||||
response = provider.image_generation(model=model_config.name, endpoint=model_config.endpoint, prompt=edit_prompt, image=edit_images, size=vsize)
|
||||
else:
|
||||
response = provider.image_generation(model=model_config.name, endpoint=model_config.endpoint, prompt=prompt)
|
||||
media = provider.extract_first_media_url(response)
|
||||
|
||||
@@ -175,7 +175,10 @@ from unittest.mock import patch
|
||||
|
||||
from django.test import TestCase
|
||||
|
||||
from apps.accounts.models import Team, User
|
||||
from rest_framework.test import APIClient
|
||||
|
||||
from apps.accounts.models import Team, TeamMember, User
|
||||
from apps.ai.models import AITask, ImageConversation, ModelConfig
|
||||
from apps.ai.providers.volcano import VolcanoArkProvider
|
||||
from apps.ai.services import enqueue_standalone_images
|
||||
from apps.assets.models import Asset, AssetFile
|
||||
@@ -245,6 +248,116 @@ class StandaloneImageReferenceTests(TestCase):
|
||||
prov.image_generation.assert_called_once()
|
||||
prov.image_edit.assert_not_called()
|
||||
|
||||
def test_image_mode_uses_uploaded_reference_images(self):
|
||||
"""图片创作自由模式:用户上传的参考图必须作为 image_edit 的参考图传进去,
|
||||
而不是被忽略走纯文生图——回归保护 #图片创作没参考我上传的素材# 这个最致命的 bug。"""
|
||||
prov = self._patch_provider()
|
||||
ref = Asset.objects.create(
|
||||
team=self.team, created_by=self.user, name="背心参考", asset_type=Asset.Type.IMAGE,
|
||||
source=Asset.Source.UPLOAD, category=Asset.Category.UPLOAD,
|
||||
)
|
||||
AssetFile.objects.create(asset=ref, object_key="r.png", bucket="b", content_type="image/png", preview_url="http://x/ref.png", is_primary=True)
|
||||
enqueue_standalone_images(
|
||||
team=self.team, user=self.user, prompt="按这件背心生成不同场景的穿搭",
|
||||
mode="image", count=1, ratio="1:1", reference_image_ids=[str(ref.id)],
|
||||
)
|
||||
prov.image_edit.assert_called_once()
|
||||
self.assertEqual(prov.image_edit.call_args.kwargs["images"], ["http://x/ref.png"])
|
||||
# 提示词必须显式钉住「保留参考图主体」+ 带上用户原话(否则图生图不真的还原上传素材)
|
||||
used_prompt = prov.image_edit.call_args.kwargs["prompt"]
|
||||
self.assertIn("按这件背心生成不同场景的穿搭", used_prompt)
|
||||
self.assertIn("参考图", used_prompt)
|
||||
self.assertIn("严格保留", used_prompt)
|
||||
prov.image_generation.assert_not_called()
|
||||
|
||||
|
||||
class ImageConversationTests(TestCase):
|
||||
"""图片创作「对话」实体:CRUD + 团队隔离 + 软删 + 生图自动归属对话 + 任务回填。"""
|
||||
|
||||
def setUp(self):
|
||||
self.user = User.objects.create_user(username="convowner", password="pass")
|
||||
self.team = Team.objects.create(name="ConvT", owner=self.user)
|
||||
TeamMember.objects.create(team=self.team, user=self.user, role="owner", status="active")
|
||||
self.client = APIClient()
|
||||
self.client.force_authenticate(self.user)
|
||||
|
||||
def test_crud_and_listing(self):
|
||||
# 新建
|
||||
r = self.client.post("/api/ai/image-conversations/", {"mode": "image", "title": "测试创作"}, format="json")
|
||||
self.assertEqual(r.status_code, 201, r.content)
|
||||
conv_id = r.json()["id"]
|
||||
# 列出(只 image 模式)
|
||||
r = self.client.get("/api/ai/image-conversations/?mode=image")
|
||||
self.assertEqual(r.status_code, 200)
|
||||
self.assertEqual(len(r.json()["results"]), 1)
|
||||
# 重命名
|
||||
r = self.client.patch(f"/api/ai/image-conversations/{conv_id}/", {"title": "改后"}, format="json")
|
||||
self.assertEqual(r.status_code, 200)
|
||||
self.assertEqual(r.json()["title"], "改后")
|
||||
# 软删:列表消失,但 DB 记录仍在(is_deleted=True)
|
||||
r = self.client.delete(f"/api/ai/image-conversations/{conv_id}/")
|
||||
self.assertIn(r.status_code, (204, 200))
|
||||
self.assertEqual(self.client.get("/api/ai/image-conversations/?mode=image").json()["results"], [])
|
||||
self.assertTrue(ImageConversation.objects.get(id=conv_id).is_deleted)
|
||||
|
||||
def test_team_isolation(self):
|
||||
other = User.objects.create_user(username="other", password="pass")
|
||||
other_team = Team.objects.create(name="Other", owner=other)
|
||||
ImageConversation.objects.create(team=other_team, created_by=other, title="别人的")
|
||||
r = self.client.get("/api/ai/image-conversations/?mode=image")
|
||||
self.assertEqual(r.json()["results"], [])
|
||||
|
||||
def test_generate_auto_creates_and_binds_conversation(self):
|
||||
# 余额 + 默认图像模型(迁移已 seed),provider/worker 全 mock,聚焦验证「任务挂到对话」
|
||||
CreditAccount.objects.create(team=self.team, balance="100.0000")
|
||||
patch("apps.ai.tasks.generate_standalone_image_task.delay").start()
|
||||
self.addCleanup(patch.stopall)
|
||||
conv = ImageConversation.objects.create(team=self.team, created_by=self.user, title="目标对话")
|
||||
tasks = enqueue_standalone_images(team=self.team, user=self.user, prompt="一只猫", mode="image", count=2, conversation=conv)
|
||||
self.assertEqual(len(tasks), 2)
|
||||
self.assertTrue(all(t.conversation_id == conv.id for t in tasks))
|
||||
self.assertEqual(conv.tasks.count(), 2)
|
||||
|
||||
def test_tasks_endpoint_returns_grouped_history(self):
|
||||
conv = ImageConversation.objects.create(team=self.team, created_by=self.user, title="历史")
|
||||
mc = ModelConfig.objects.filter(capability=ModelConfig.Capability.IMAGE).first()
|
||||
AITask.objects.create(
|
||||
team=self.team, created_by=self.user, conversation=conv, task_type=AITask.Type.PRODUCT_IMAGE,
|
||||
status=AITask.Status.SUCCEEDED, model_config=mc, idempotency_key="conv-test-1",
|
||||
request_payload={"prompt": "猫", "batch_id": "bx", "ratio": "1:1"},
|
||||
)
|
||||
r = self.client.get(f"/api/ai/image-conversations/{conv.id}/tasks/")
|
||||
self.assertEqual(r.status_code, 200, r.content)
|
||||
data = r.json()["tasks"]
|
||||
self.assertEqual(len(data), 1)
|
||||
self.assertEqual(data[0]["prompt"], "猫")
|
||||
self.assertEqual(data[0]["batch_id"], "bx")
|
||||
|
||||
def test_generate_with_refs_persists_and_tasks_endpoint_returns_them(self):
|
||||
"""HTTP 全链路:带 reference_image_ids 提交 → 任务 payload 落 ids → tasks 接口把参考图解析回 {name,url}。
|
||||
守护「上传的参考图被带进生成 + 切换/刷新后批次头仍能回显参考图」。"""
|
||||
CreditAccount.objects.create(team=self.team, balance="100.0000")
|
||||
patch("apps.ai.tasks.generate_standalone_image_task.delay").start()
|
||||
self.addCleanup(patch.stopall)
|
||||
ref = Asset.objects.create(
|
||||
team=self.team, created_by=self.user, name="背心参考", asset_type=Asset.Type.IMAGE,
|
||||
source=Asset.Source.UPLOAD, category=Asset.Category.UPLOAD,
|
||||
)
|
||||
AssetFile.objects.create(asset=ref, object_key="r.png", bucket="b", content_type="image/png", preview_url="http://x/ref.png", is_primary=True)
|
||||
r = self.client.post(
|
||||
"/api/ai/generate-image/",
|
||||
{"prompt": "按这件背心生成不同场景", "mode": "image", "count": 1, "reference_image_ids": [str(ref.id)]},
|
||||
format="json",
|
||||
)
|
||||
self.assertEqual(r.status_code, 202, r.content)
|
||||
conv_id = r.json()["conversation_id"]
|
||||
# 任务 payload 落了参考图 id
|
||||
task = AITask.objects.filter(conversation_id=conv_id).first()
|
||||
self.assertEqual(task.request_payload.get("reference_image_ids"), [str(ref.id)])
|
||||
# tasks 接口把参考图解析成 {name,url} 回显
|
||||
data = self.client.get(f"/api/ai/image-conversations/{conv_id}/tasks/").json()["tasks"]
|
||||
self.assertEqual(data[0]["reference_images"], [{"name": "背心参考", "url": "http://x/ref.png"}])
|
||||
|
||||
|
||||
class StandaloneCategoryTests(TestCase):
|
||||
"""图片趴三类归类(期2):模特上身图→model_tryon / 平台套图→platform_kit / 自由创作→free_create;
|
||||
|
||||
@@ -1,11 +1,12 @@
|
||||
from django.urls import path
|
||||
from rest_framework.routers import DefaultRouter
|
||||
|
||||
from .views import AITaskViewSet, GenerateImageView, ModelConfigViewSet
|
||||
from .views import AITaskViewSet, GenerateImageView, ImageConversationViewSet, ModelConfigViewSet
|
||||
|
||||
router = DefaultRouter()
|
||||
router.register("tasks", AITaskViewSet, basename="ai-task")
|
||||
router.register("models", ModelConfigViewSet, basename="model-config")
|
||||
router.register("image-conversations", ImageConversationViewSet, basename="image-conversation")
|
||||
|
||||
urlpatterns = [
|
||||
path("generate-image/", GenerateImageView.as_view(), name="ai-generate-image"),
|
||||
|
||||
@@ -1,14 +1,17 @@
|
||||
from django.db.models import Count
|
||||
from django.utils import timezone
|
||||
from rest_framework import status
|
||||
from rest_framework.decorators import action
|
||||
from rest_framework.response import Response
|
||||
from rest_framework.views import APIView
|
||||
from rest_framework.viewsets import ReadOnlyModelViewSet
|
||||
from rest_framework.viewsets import ModelViewSet, 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 .models import AITask, ImageConversation, ModelConfig
|
||||
from .serializers import AITaskSerializer, ImageConversationSerializer, ModelConfigSerializer
|
||||
from .services import enqueue_standalone_images
|
||||
|
||||
|
||||
@@ -35,13 +38,33 @@ class GenerateImageView(APIView):
|
||||
model_entity_id = str(request.data.get("model_entity_id") or "").strip() or None
|
||||
ratio = str(request.data.get("ratio") or "").strip() or None
|
||||
image_model = str(request.data.get("image_model") or "").strip() or None
|
||||
conversation_id = str(request.data.get("conversation_id") or "").strip() or None
|
||||
# 用户在图片创作里上传的参考图(已先传成 Asset),按 id 列表带入 → 生成时作多图参考(image_edit)
|
||||
raw_refs = request.data.get("reference_image_ids") or []
|
||||
if isinstance(raw_refs, str):
|
||||
raw_refs = [s for s in raw_refs.split(",") if s.strip()]
|
||||
reference_image_ids = [str(r).strip() for r in raw_refs if str(r).strip()]
|
||||
team = get_current_team(request.user)
|
||||
# 对话归属:传了 id 就用现成对话(限本团队);没传则自动开一条新对话,标题取 prompt 前 24 字。
|
||||
conversation = None
|
||||
if conversation_id:
|
||||
conversation = ImageConversation.objects.filter(team=team, id=conversation_id, is_deleted=False).first()
|
||||
if conversation is None:
|
||||
conversation = ImageConversation.objects.create(
|
||||
team=team,
|
||||
created_by=request.user,
|
||||
mode=mode if mode in dict(ImageConversation.Mode.choices) else ImageConversation.Mode.IMAGE,
|
||||
product_id=product_id,
|
||||
title=(prompt[:24] or "默认创作"),
|
||||
)
|
||||
try:
|
||||
tasks = enqueue_standalone_images(team=team, user=request.user, prompt=prompt, mode=mode, count=count, product_id=product_id, reference_product=reference_product, model_id=model_id, model_entity_id=model_entity_id, ratio=ratio, image_model=image_model)
|
||||
tasks = enqueue_standalone_images(team=team, user=request.user, prompt=prompt, mode=mode, count=count, product_id=product_id, reference_product=reference_product, model_id=model_id, model_entity_id=model_entity_id, ratio=ratio, image_model=image_model, conversation=conversation, reference_image_ids=reference_image_ids)
|
||||
except ValueError as exc: # 无可用模型 / 余额不足等,立即反馈
|
||||
return Response({"detail": str(exc)}, status=status.HTTP_400_BAD_REQUEST)
|
||||
# 本次提交即刷新对话活跃时间,左栏「最近」据此置顶
|
||||
ImageConversation.objects.filter(id=conversation.id).update(last_active_at=timezone.now())
|
||||
return Response(
|
||||
{"tasks": [{"id": str(t.id), "status": t.status} for t in tasks]},
|
||||
{"conversation_id": str(conversation.id), "tasks": [{"id": str(t.id), "status": t.status} for t in tasks]},
|
||||
status=status.HTTP_202_ACCEPTED,
|
||||
)
|
||||
|
||||
@@ -82,6 +105,13 @@ class AITaskViewSet(TeamScopedViewSetMixin, ReadOnlyModelViewSet):
|
||||
# 可选 ?task_type=a,b,c 过滤:生图工作室的任务中心只想看生图任务(模特上身图/平台套图/
|
||||
# 图片创作 = person_image / product_image),不掺脚本/实体抽取/故事板等流水线内部任务。
|
||||
queryset = super().get_queryset()
|
||||
# 从 request_payload(已 defer)里只抽 batch_id / mode 两个 JSON 标量供前端「按批分组 + 标签」用:
|
||||
# KeyTextTransform 在 SQL 层 JSON_EXTRACT,不会把几 MB 的 payload 整列拉回(避开 payload 性能坑)。
|
||||
from django.db.models.fields.json import KeyTextTransform
|
||||
queryset = queryset.annotate(
|
||||
rp_batch_id=KeyTextTransform("batch_id", "request_payload"),
|
||||
rp_mode=KeyTextTransform("mode", "request_payload"),
|
||||
)
|
||||
raw = self.request.query_params.get("task_type", "").strip()
|
||||
if raw:
|
||||
types = [t.strip() for t in raw.split(",") if t.strip()]
|
||||
@@ -90,6 +120,73 @@ class AITaskViewSet(TeamScopedViewSetMixin, ReadOnlyModelViewSet):
|
||||
return queryset
|
||||
|
||||
|
||||
class ImageConversationViewSet(TeamScopedViewSetMixin, ModelViewSet):
|
||||
"""图片创作工作室的「对话」CRUD。
|
||||
list 按 ?mode= 过滤、排除软删、按 last_active_at 倒序(左栏「最近」);
|
||||
create 开新对话;partial_update 重命名;destroy 软删(不连带删图)。
|
||||
detail action `tasks` 返回该对话下的生图任务 + 成图 asset,供切换对话时回填批次流。
|
||||
"""
|
||||
|
||||
serializer_class = ImageConversationSerializer
|
||||
queryset = ImageConversation.objects.filter(is_deleted=False).select_related("product").order_by("-last_active_at")
|
||||
|
||||
def get_queryset(self):
|
||||
queryset = super().get_queryset().annotate(_task_count=Count("tasks"))
|
||||
mode = self.request.query_params.get("mode", "").strip()
|
||||
if mode:
|
||||
queryset = queryset.filter(mode=mode)
|
||||
return queryset
|
||||
|
||||
def perform_destroy(self, instance):
|
||||
# 软删:对话从列表消失,但其 AITask.conversation 置空由 DB on_delete=SET_NULL 不触发(我们没真删),
|
||||
# 成图始终留在资产库。仅打标记。
|
||||
instance.is_deleted = True
|
||||
instance.save(update_fields=["is_deleted", "updated_at"])
|
||||
|
||||
@action(detail=True, methods=["get"])
|
||||
def tasks(self, request, pk=None):
|
||||
from apps.assets.models import Asset
|
||||
from apps.assets.serializers import _asset_preview
|
||||
|
||||
conversation = self.get_object()
|
||||
tasks = (
|
||||
AITask.objects.filter(conversation=conversation)
|
||||
.prefetch_related("generated_assets", "generated_assets__files")
|
||||
.order_by("created_at")
|
||||
)
|
||||
# 参考图 id → {name,url}:跨任务可能重复,缓存一次解析,供切换/刷新后批次头回显「参考了哪些图」
|
||||
ref_cache: dict[str, dict] = {}
|
||||
|
||||
def resolve_refs(ids):
|
||||
out = []
|
||||
for rid in ids or []:
|
||||
rid = str(rid)
|
||||
if rid not in ref_cache:
|
||||
a = Asset.objects.filter(id=rid).prefetch_related("files").first()
|
||||
ref_cache[rid] = {"name": a.name, "url": _asset_preview(a)} if a else None
|
||||
if ref_cache[rid]:
|
||||
out.append(ref_cache[rid])
|
||||
return out
|
||||
|
||||
data = [
|
||||
{
|
||||
"id": str(t.id),
|
||||
"status": t.status,
|
||||
"error_message": t.error_message,
|
||||
"prompt": (t.request_payload or {}).get("prompt", ""),
|
||||
"batch_id": (t.request_payload or {}).get("batch_id", ""),
|
||||
"ratio": (t.request_payload or {}).get("ratio") or "",
|
||||
"reference_images": resolve_refs((t.request_payload or {}).get("reference_image_ids")),
|
||||
"created_at": t.created_at,
|
||||
"assets": AssetSerializer(
|
||||
[a for a in t.generated_assets.all() if not a.is_deleted], many=True
|
||||
).data,
|
||||
}
|
||||
for t in tasks
|
||||
]
|
||||
return Response({"conversation_id": str(conversation.id), "tasks": data})
|
||||
|
||||
|
||||
class ModelConfigViewSet(ReadOnlyModelViewSet):
|
||||
# 按创建序固定排序:最早创建的 active 模型排第一 = 前端选择器默认项,与 get_default_model 口径一致
|
||||
# (否则 DB 默认序不稳定,可能默认选到 Gemini 等;用户要默认 = 豆包 2.0 Pro,它最早创建)
|
||||
|
||||
@@ -104,11 +104,18 @@ class AssetViewSet(TeamScopedViewSetMixin, ModelViewSet):
|
||||
参数:tab(资产库分桶)/category/asset_type/source/product/q(搜索)/m_<key>(metadata 过滤)/ordering。
|
||||
默认按 -created_at 排序——无 ORDER BY 时分页会重/漏。"""
|
||||
qs = super().get_queryset().filter(is_deleted=False) # 软删资产不出现在资产库
|
||||
# 资产库列表/批次只展示「已加入资产库」的资产(in_library=True);未加入的工作台生成图不出现在这里。
|
||||
# 仅对列表型 action 过滤——retrieve / set-library / submit-review 等仍要能取到未加入的资产。
|
||||
if self.action in ("list", "batches"):
|
||||
qs = qs.filter(in_library=True)
|
||||
p = self.request.query_params
|
||||
# 资产库列表/批次默认只展示「已加入资产库」的资产(in_library=True);未加入的工作台生成图不出现在这里。
|
||||
# 仅对列表型 action 过滤——retrieve / set-library / submit-review 等仍要能取到未加入的资产。
|
||||
# 任务中心要看「全部生成图」(含未入库),传 ?in_library=all 旁路;?in_library=false 只看未入库。
|
||||
if self.action in ("list", "batches"):
|
||||
lib = (p.get("in_library") or "").lower()
|
||||
if lib == "all":
|
||||
pass
|
||||
elif lib in ("false", "0"):
|
||||
qs = qs.filter(in_library=False)
|
||||
else:
|
||||
qs = qs.filter(in_library=True)
|
||||
if p.get("tab"):
|
||||
qs = qs.filter(_tab_q(p["tab"]))
|
||||
if p.get("category"):
|
||||
@@ -117,6 +124,8 @@ class AssetViewSet(TeamScopedViewSetMixin, ModelViewSet):
|
||||
qs = qs.filter(asset_type=p["asset_type"])
|
||||
if p.get("source"):
|
||||
qs = qs.filter(source=p["source"])
|
||||
if p.get("origin_task"):
|
||||
qs = qs.filter(origin_task_id=p["origin_task"])
|
||||
if p.get("product"):
|
||||
# 资产归属商品:独立生图写 metadata.product_id;项目内生成回溯 origin_task→project→product;
|
||||
# 上传的商品图经 ProductImage 关联。三路并集 = 前端 ProductDetail 的 belongsToProduct。
|
||||
|
||||
@@ -872,3 +872,82 @@ class BaseAssetAdoptDeleteTests(TestCase):
|
||||
res = self.client.post(f"/api/projects/{self.proj.id}/delete-base-asset/", {"group_id": str(self.group.id)}, format="json")
|
||||
self.assertEqual(res.status_code, 200)
|
||||
self.assertFalse(BaseAssetGroup.objects.filter(id=self.group.id).exists())
|
||||
|
||||
|
||||
class AttachOfficialModelTests(TestCase):
|
||||
"""ZWQ#4/#5 回归:从演员库挑「官方模特」(跨团队)新增/替换角色,旧版因 team 隔离报 asset not found
|
||||
且替换失败、还留下空角色组。修复后:官方模特形象图自动克隆进本团队并真正挂上,不再 404。"""
|
||||
|
||||
def setUp(self):
|
||||
from apps.assets.models import Asset, AssetFile, Model
|
||||
from apps.projects.models import BaseAssetGroup
|
||||
|
||||
self.user = User.objects.create_user(username="biz", password="x")
|
||||
self.team = Team.objects.create(name="商家团队", owner=self.user)
|
||||
TeamMember.objects.create(team=self.team, user=self.user, role=TeamMember.Role.OWNER)
|
||||
self.product = Product.objects.create(team=self.team, created_by=self.user, title="口红")
|
||||
self.proj = Project.objects.create(team=self.team, name="带货视频", product=self.product, created_by=self.user)
|
||||
|
||||
# 官方模特属于「官方团队」(跨团队可见,但 portrait 资产不属本商家团队)——正是触发 bug 的根因。
|
||||
self.official_owner = User.objects.create_user(username="ops", password="x")
|
||||
self.official_team = Team.objects.create(name="官方", owner=self.official_owner)
|
||||
self.portrait = Asset.objects.create(
|
||||
team=self.official_team, created_by=self.official_owner, name="都市白领 林晚",
|
||||
asset_type=Asset.Type.IMAGE, source=Asset.Source.AI_GENERATED, category=Asset.Category.PERSON,
|
||||
metadata={"kind": "model"},
|
||||
)
|
||||
AssetFile.objects.create(asset=self.portrait, object_key="m/linwan.png", bucket="b", content_type="image/png", preview_url="http://x/linwan.png", is_primary=True)
|
||||
self.model = Model.objects.create(
|
||||
team=self.official_team, created_by=self.official_owner, name="都市白领 林晚",
|
||||
is_official=True, portrait_asset=self.portrait,
|
||||
)
|
||||
self.BaseAssetGroup = BaseAssetGroup
|
||||
self.client = APIClient()
|
||||
self.client.force_authenticate(self.user)
|
||||
|
||||
def test_add_actor_from_official_model_clones_and_attaches(self):
|
||||
"""新增角色:无 group,传 kind+label,选官方模特 → 200,克隆进本团队并采用,不再 asset not found。"""
|
||||
from apps.assets.models import Asset
|
||||
|
||||
res = self.client.post(
|
||||
f"/api/projects/{self.proj.id}/attach-base-asset/",
|
||||
{"kind": "person", "label": "林晚", "asset_id": str(self.portrait.id)},
|
||||
format="json",
|
||||
)
|
||||
self.assertEqual(res.status_code, 200, res.content)
|
||||
group = self.BaseAssetGroup.objects.get(id=res.json()["id"])
|
||||
self.assertIsNotNone(group.adopted_asset_id)
|
||||
clone = Asset.objects.get(id=group.adopted_asset_id)
|
||||
# 克隆体真正落进本商家团队(不是官方团队),并共享同一份 TOS 对象。
|
||||
self.assertEqual(clone.team_id, self.team.id)
|
||||
self.assertNotEqual(clone.id, self.portrait.id)
|
||||
self.assertEqual(clone.files.first().object_key, "m/linwan.png")
|
||||
self.assertEqual(clone.metadata.get("cloned_from_asset"), str(self.portrait.id))
|
||||
|
||||
def test_replace_existing_group_with_official_model(self):
|
||||
"""替换:已有角色组,传 group_id + 官方模特 asset_id → 200,采用克隆体。"""
|
||||
from apps.assets.models import Asset
|
||||
|
||||
group = self.BaseAssetGroup.objects.create(project=self.proj, kind=self.BaseAssetGroup.Kind.PERSON, metadata={"label": "主播"})
|
||||
res = self.client.post(
|
||||
f"/api/projects/{self.proj.id}/attach-base-asset/",
|
||||
{"group_id": str(group.id), "asset_id": str(self.portrait.id)},
|
||||
format="json",
|
||||
)
|
||||
self.assertEqual(res.status_code, 200, res.content)
|
||||
group.refresh_from_db()
|
||||
clone = Asset.objects.get(id=group.adopted_asset_id)
|
||||
self.assertEqual(clone.team_id, self.team.id)
|
||||
|
||||
def test_unknown_asset_still_404_without_ghost_group(self):
|
||||
"""彻底不存在的 asset_id → 仍 404,且不留下空角色组(不再在 404 前提交建组)。"""
|
||||
import uuid as _uuid
|
||||
|
||||
before = self.BaseAssetGroup.objects.filter(project=self.proj).count()
|
||||
res = self.client.post(
|
||||
f"/api/projects/{self.proj.id}/attach-base-asset/",
|
||||
{"kind": "person", "label": "幽灵", "asset_id": str(_uuid.uuid4())},
|
||||
format="json",
|
||||
)
|
||||
self.assertEqual(res.status_code, 404)
|
||||
self.assertEqual(self.BaseAssetGroup.objects.filter(project=self.proj).count(), before)
|
||||
|
||||
@@ -3,7 +3,7 @@ from pathlib import Path
|
||||
import uuid
|
||||
|
||||
from django.db import transaction
|
||||
from django.db.models import Count
|
||||
from django.db.models import Count, Q
|
||||
from django.http import HttpResponse, JsonResponse, StreamingHttpResponse
|
||||
from rest_framework import status
|
||||
from rest_framework.decorators import action
|
||||
@@ -105,6 +105,39 @@ def _store_uploaded_asset(*, team, user, upload, asset_type: str, category: str,
|
||||
return asset
|
||||
|
||||
|
||||
def _clone_asset_into_team(src: Asset, *, team, user) -> Asset:
|
||||
"""把跨团队可见的「官方模特」形象图克隆一份进目标团队:共享同一份 TOS 对象(只新建 Asset/AssetFile DB 行,
|
||||
不重新上传),使其成为本团队真实资产 —— 可被基础资产组采用、在「我的演员」里列出,不再因 team 隔离报 asset not found。
|
||||
去掉 metadata 里的 kind=model 标记,让克隆体作为普通「我的演员」person 资产,而不是被误判成模特库预设。"""
|
||||
src_meta = {k: v for k, v in (src.metadata or {}).items() if k != "kind"}
|
||||
clone = Asset.objects.create(
|
||||
team=team,
|
||||
created_by=user,
|
||||
name=src.name,
|
||||
asset_type=src.asset_type,
|
||||
source=src.source,
|
||||
category=src.category,
|
||||
description=src.description,
|
||||
metadata={**src_meta, "cloned_from_asset": str(src.id), "from_official_model": True},
|
||||
in_library=src.in_library,
|
||||
)
|
||||
for f in src.files.all():
|
||||
AssetFile.objects.create(
|
||||
asset=clone,
|
||||
object_key=f.object_key,
|
||||
bucket=f.bucket,
|
||||
content_type=f.content_type,
|
||||
size_bytes=f.size_bytes,
|
||||
checksum=f.checksum,
|
||||
width=f.width,
|
||||
height=f.height,
|
||||
duration_ms=f.duration_ms,
|
||||
preview_url=f.preview_url,
|
||||
is_primary=f.is_primary,
|
||||
)
|
||||
return clone
|
||||
|
||||
|
||||
def promote_base_asset_stage_if_ready(project: Project) -> bool:
|
||||
adopted_kind_count = (
|
||||
project.base_asset_groups.filter(adopted_asset__isnull=False).values("kind").distinct().count()
|
||||
@@ -492,8 +525,26 @@ class ProjectViewSet(TeamScopedViewSetMixin, ModelViewSet):
|
||||
@transaction.atomic
|
||||
def attach_base_asset(self, request, pk=None):
|
||||
"""流程步骤4 · 用演员库现有资产替换某基础资产卡:把团队内现有 Asset 挂为该组候选并采用。
|
||||
seed 占位卡还没有 group,此时不传 group_id,改传 kind+label:按 label 命中同实体组,没有则据 label 建组再挂(不出图)。"""
|
||||
seed 占位卡还没有 group,此时不传 group_id,改传 kind+label:按 label 命中同实体组,没有则据 label 建组再挂(不出图)。
|
||||
所选资产可能是「官方模特」形象图(跨团队可见但不属本团队)→ 克隆一份进本团队再挂:避免 team 隔离导致
|
||||
明明选了官方模特却报 asset not found、替换失败。"""
|
||||
project = self.get_object()
|
||||
# 先把要挂的资产定下来(本团队资产,或官方模特形象图克隆),拿不到立即 404 —— 必须在建/找组「之前」,
|
||||
# 否则会留下一个采用资产为空的幽灵角色组(旧版本在 404 return 前就 create 了组,且 @atomic 正常 return 仍会提交)。
|
||||
asset_id = request.data.get("asset_id")
|
||||
asset = Asset.objects.filter(team=project.team, id=asset_id).first()
|
||||
if asset is None and asset_id:
|
||||
# 跨团队可见的官方模特形象图 / 三视图 → 克隆进本团队再挂
|
||||
official = (
|
||||
Asset.objects.filter(id=asset_id, is_deleted=False)
|
||||
.filter(Q(model_as_portrait__is_official=True) | Q(model_as_triview__is_official=True))
|
||||
.first()
|
||||
)
|
||||
if official is not None:
|
||||
asset = _clone_asset_into_team(official, team=project.team, user=request.user)
|
||||
if asset is None:
|
||||
return Response({"detail": "asset not found"}, status=status.HTTP_404_NOT_FOUND)
|
||||
|
||||
group = None
|
||||
group_id = request.data.get("group_id")
|
||||
if group_id:
|
||||
@@ -517,9 +568,6 @@ class ProjectViewSet(TeamScopedViewSetMixin, ModelViewSet):
|
||||
group_meta = {"label": label} if label else {}
|
||||
group = BaseAssetGroup.objects.create(project=project, kind=kind, metadata=group_meta)
|
||||
group = BaseAssetGroup.objects.select_for_update().get(id=group.id)
|
||||
asset = Asset.objects.filter(team=project.team, id=request.data.get("asset_id")).first()
|
||||
if asset is None:
|
||||
return Response({"detail": "asset not found"}, status=status.HTTP_404_NOT_FOUND)
|
||||
group.candidate_assets.add(asset)
|
||||
group.adopted_asset_id = asset.id
|
||||
group.save(update_fields=["adopted_asset", "updated_at"])
|
||||
|
||||
Reference in New Issue
Block a user