feat: 解耦角色与模特并完善模特库
This commit is contained in:
@@ -235,6 +235,10 @@ PROVIDER_API_VERSIONS = {
|
||||
# 自由创作视频:团队在途任务并发上限(视频长时高价,限并发同时压 web 慢调用敞口)
|
||||
FREE_VIDEO_MAX_CONCURRENT = int(env("FREE_VIDEO_MAX_CONCURRENT", "3"))
|
||||
|
||||
# 团队模特库三视图生成:必须在 API 与 Celery worker 同步部署新代码后才显式开启。
|
||||
# 9.4A 默认关闭,避免本地 API 把新任务投给共享的旧 worker。
|
||||
MODEL_TRIVIEW_GENERATION_ENABLED = env_bool("MODEL_TRIVIEW_GENERATION_ENABLED", False)
|
||||
|
||||
# 火山引擎人像素材库审核(真人资产绿/红盾)· AK/SK 暂借 AirDrama 已邀测账号,张业昌待换成 AirShelf 自有
|
||||
ASSETS_API = {
|
||||
"access_key": env("ASSETS_API_ACCESS_KEY", ""),
|
||||
|
||||
@@ -0,0 +1,32 @@
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
("ai", "0026_conversation_task_purged_at"),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.AlterField(
|
||||
model_name="aitask",
|
||||
name="task_type",
|
||||
field=models.CharField(
|
||||
choices=[
|
||||
("script_generation", "Script Generation"),
|
||||
("script_optimization", "Script Optimization"),
|
||||
("entity_extraction", "Entity Extraction"),
|
||||
("product_image", "Product Image"),
|
||||
("person_image", "Person Image"),
|
||||
("model_triview", "Model Triview"),
|
||||
("scene_image", "Scene Image"),
|
||||
("storyboard", "Storyboard"),
|
||||
("video_segment", "Video Segment"),
|
||||
("voiceover", "Voiceover"),
|
||||
("export", "Export"),
|
||||
("free_video", "Free Video"),
|
||||
],
|
||||
max_length=48,
|
||||
),
|
||||
),
|
||||
]
|
||||
@@ -94,6 +94,7 @@ class AITask(TeamOwnedModel):
|
||||
ENTITY_EXTRACTION = "entity_extraction", "Entity Extraction"
|
||||
PRODUCT_IMAGE = "product_image", "Product Image"
|
||||
PERSON_IMAGE = "person_image", "Person Image"
|
||||
MODEL_TRIVIEW = "model_triview", "Model Triview"
|
||||
SCENE_IMAGE = "scene_image", "Scene Image"
|
||||
STORYBOARD = "storyboard", "Storyboard"
|
||||
VIDEO_SEGMENT = "video_segment", "Video Segment"
|
||||
|
||||
@@ -1319,9 +1319,9 @@ def generate_base_asset(*, project, user, kind: str, prompt: str, label: str = "
|
||||
"model": model_config.name, "endpoint": model_config.endpoint, "prompt": gen_prompt,
|
||||
"kind": kind, "label": label or "", "group_id": str(group_id) if group_id else "",
|
||||
"use_edit": use_edit, "reference_image": ref_url,
|
||||
# 角色立绘出图成功后,worker 内自动接力生成它配套的三视图(ZWQ#3)。仅人物立绘生效;
|
||||
# 默认关闭(显式由调用方传 True),避免给商品/场景或不需要三视图的流程白白多扣一次费。
|
||||
"auto_triview": bool(auto_triview) and kind == BaseAssetGroup.Kind.PERSON,
|
||||
# 角色立绘不再自动接力三视图;三视图只通过角色详情里的显式按钮生成。
|
||||
# 保留字段为 False,兼容旧前端/旧任务读取,但不允许再触发自动链路。
|
||||
"auto_triview": False,
|
||||
}
|
||||
task = create_ai_task(
|
||||
project=project,
|
||||
@@ -1426,17 +1426,6 @@ def run_base_asset_task(*, task_id: str) -> None:
|
||||
from apps.assets.review import submit_asset_for_review
|
||||
|
||||
transaction.on_commit(lambda a=asset: submit_asset_for_review(a))
|
||||
# ZWQ#3 · 角色立绘出图成功后自动接力生成它配套的三视图(opt-in;商品/场景不触发)。
|
||||
# 放 on_commit:确保立绘已落库再据它发起三视图(image_edit 以该立绘为参考)。
|
||||
# best-effort:三视图发起失败(额度不足/模型不支持 image_edit 等)不影响立绘本身已成功。
|
||||
if kind == BaseAssetGroup.Kind.PERSON and payload.get("auto_triview"):
|
||||
def _kickoff_triview(portrait=asset, proj=project, usr=user):
|
||||
try:
|
||||
generate_person_triview(project=proj, user=usr, portrait_asset=portrait)
|
||||
except Exception: # noqa: BLE001 — 三视图接力是附加能力,失败绝不回头弄挂立绘流程
|
||||
logger.exception("auto triview kickoff failed for portrait %s", getattr(portrait, "id", "?"))
|
||||
|
||||
transaction.on_commit(_kickoff_triview)
|
||||
except Exception as exc: # noqa: BLE001 — 失败要退费并把错误记进 AITask 供前端轮询读取;不向上抛(避免 celery 重试二次扣费)
|
||||
task.status = AITask.Status.FAILED
|
||||
task.error_message = str(exc)
|
||||
@@ -1474,6 +1463,202 @@ def generate_person_triview(*, project, user, portrait_asset) -> AITask:
|
||||
return task
|
||||
|
||||
|
||||
_MODEL_TRIVIEW_INFLIGHT = (
|
||||
AITask.Status.CREATED,
|
||||
AITask.Status.RESERVED,
|
||||
AITask.Status.SUBMITTED,
|
||||
AITask.Status.POLLING,
|
||||
AITask.Status.POSTPROCESSING,
|
||||
)
|
||||
|
||||
|
||||
def quote_model_triview(*, model):
|
||||
"""模特三视图价格:图像模型标准价(当前 20 积分)× 团队价格系数。"""
|
||||
from apps.billing.pricing import quote_flat
|
||||
|
||||
model_config = get_default_model(ModelConfig.Capability.IMAGE)
|
||||
if model_config is None:
|
||||
raise ValueError("no active image model configured")
|
||||
return model_config, quote_flat(model_config, team=model.team)
|
||||
|
||||
|
||||
def generate_model_triview(*, model, user) -> tuple[AITask, bool]:
|
||||
"""为团队级 Model 提交三视图任务;project 保持空,同一模特只允许一个在途任务。"""
|
||||
from apps.ai.model_library import THREE_VIEW_PROMPT
|
||||
from apps.ai.tasks import generate_model_triview_task
|
||||
from apps.assets.models import Model
|
||||
|
||||
if model.is_official:
|
||||
raise ValueError("官方模特不可生成三视图")
|
||||
if model.portrait_asset_id is None:
|
||||
raise ValueError("请先设置模特形象图")
|
||||
model_config, quote = quote_model_triview(model=model)
|
||||
provider = get_image_provider(model_config)
|
||||
if not hasattr(provider, "image_edit"):
|
||||
raise ValueError(f"当前图像模型 {model_config.provider.name}:{model_config.name} 不支持参考图三视图(image_edit)")
|
||||
ref_url = _asset_preview_url(model.portrait_asset)
|
||||
if not ref_url:
|
||||
raise ValueError("模特形象图不可用,请先重新上传")
|
||||
prompt = render_prompt("person_triview", THREE_VIEW_PROMPT)
|
||||
|
||||
with transaction.atomic():
|
||||
locked = Model.objects.select_for_update().select_related("portrait_asset", "team").get(id=model.id)
|
||||
existing = (
|
||||
AITask.objects.filter(
|
||||
team=locked.team,
|
||||
task_type=AITask.Type.MODEL_TRIVIEW,
|
||||
status__in=_MODEL_TRIVIEW_INFLIGHT,
|
||||
request_payload__model_id=str(locked.id),
|
||||
)
|
||||
.order_by("-created_at")
|
||||
.first()
|
||||
)
|
||||
if existing is not None:
|
||||
return existing, False
|
||||
payload = {
|
||||
"model": model_config.name,
|
||||
"prompt": prompt,
|
||||
"kind": "model_triview",
|
||||
"model_id": str(locked.id),
|
||||
"portrait_asset_id": str(locked.portrait_asset_id),
|
||||
"reference_image": ref_url,
|
||||
"price_points": str(quote.points),
|
||||
}
|
||||
task = AITask.objects.create(
|
||||
team=locked.team,
|
||||
created_by=user,
|
||||
project=None,
|
||||
task_type=AITask.Type.MODEL_TRIVIEW,
|
||||
status=AITask.Status.CREATED,
|
||||
model_config=model_config,
|
||||
idempotency_key=f"model_triview:{locked.id}:{uuid.uuid4()}",
|
||||
request_payload=payload,
|
||||
estimated_cost=quote.points,
|
||||
base_cost=quote.base_cost_yuan,
|
||||
)
|
||||
reserve_credit(team=locked.team, user=user, task=task, amount=quote.points)
|
||||
task.status = AITask.Status.RESERVED
|
||||
task.save(update_fields=["status", "updated_at"])
|
||||
|
||||
generate_model_triview_task.delay(str(task.id))
|
||||
return task, True
|
||||
|
||||
|
||||
def run_model_triview_task(*, task_id: str) -> None:
|
||||
"""worker 执行团队模特三视图:成功扣费并切换 Model.triview_asset;失败释放预留。"""
|
||||
from apps.assets.models import Model
|
||||
|
||||
with transaction.atomic():
|
||||
task = (
|
||||
AITask.objects.select_for_update()
|
||||
.select_related("team", "created_by", "model_config")
|
||||
.filter(id=task_id)
|
||||
.first()
|
||||
)
|
||||
if task is None or task.status != AITask.Status.RESERVED:
|
||||
return
|
||||
task.status = AITask.Status.SUBMITTED
|
||||
task.submitted_at = timezone.now()
|
||||
task.save(update_fields=["status", "submitted_at", "updated_at"])
|
||||
|
||||
payload = task.request_payload or {}
|
||||
model_id = str(payload.get("model_id") or "")
|
||||
portrait_asset_id = str(payload.get("portrait_asset_id") or "")
|
||||
ref_url = str(payload.get("reference_image") or "")
|
||||
prompt = str(payload.get("prompt") or "")
|
||||
provider = get_image_provider(task.model_config)
|
||||
reservation = task.credit_reservation
|
||||
try:
|
||||
if not model_id or not portrait_asset_id or not ref_url:
|
||||
raise ValueError("模特三视图任务参数不完整")
|
||||
|
||||
def _make():
|
||||
size = prompt_ratio_size("person_triview", "1536x1024")
|
||||
response = provider.image_edit(
|
||||
model=task.model_config.name,
|
||||
prompt=prompt,
|
||||
images=[ref_url],
|
||||
size=size,
|
||||
)
|
||||
return response, provider.extract_first_media_url(response)
|
||||
|
||||
response, media = _run_image_with_retry(_make)
|
||||
with transaction.atomic():
|
||||
locked_model = (
|
||||
Model.objects.select_for_update()
|
||||
.filter(id=model_id, team=task.team, is_deleted=False, purged_at__isnull=True)
|
||||
.first()
|
||||
)
|
||||
if locked_model is None:
|
||||
raise ValueError("模特不存在或已删除")
|
||||
if str(locked_model.portrait_asset_id or "") != portrait_asset_id:
|
||||
raise ValueError("模特形象图已变化,请重新生成三视图")
|
||||
|
||||
fileobj, content_type = VolcanoArkProvider.media_to_bytes(media)
|
||||
suffix = ".jpg" if "jpeg" in content_type else ".webp" if "webp" in content_type else ".png"
|
||||
asset_id = uuid.uuid4()
|
||||
stored = TosStorage().upload_fileobj(
|
||||
fileobj=fileobj,
|
||||
object_key=f"teams/{task.team_id}/models/{locked_model.id}/triviews/{asset_id}{suffix}",
|
||||
content_type=content_type,
|
||||
)
|
||||
triview = Asset.objects.create(
|
||||
id=asset_id,
|
||||
team=task.team,
|
||||
created_by=task.created_by,
|
||||
name=f"{locked_model.name}·三视图",
|
||||
asset_type=Asset.Type.IMAGE,
|
||||
source=Asset.Source.AI_GENERATED,
|
||||
category=Asset.Category.TRI_VIEW,
|
||||
in_library=False,
|
||||
origin_task=task,
|
||||
metadata={
|
||||
"kind": "model",
|
||||
"view": "three_view",
|
||||
"model_id": str(locked_model.id),
|
||||
"triview_of": portrait_asset_id,
|
||||
},
|
||||
)
|
||||
AssetFile.objects.create(
|
||||
asset=triview,
|
||||
object_key=stored.object_key,
|
||||
bucket=stored.bucket,
|
||||
content_type=stored.content_type,
|
||||
size_bytes=stored.size_bytes,
|
||||
is_primary=True,
|
||||
)
|
||||
versions = [str(value) for value in (locked_model.metadata or {}).get("triview_versions", []) if value]
|
||||
for value in (locked_model.triview_asset_id, triview.id):
|
||||
if value and str(value) not in versions:
|
||||
versions.append(str(value))
|
||||
metadata = dict(locked_model.metadata or {})
|
||||
metadata["triview_versions"] = versions
|
||||
locked_model.triview_asset = triview
|
||||
locked_model.metadata = metadata
|
||||
locked_model.save(update_fields=["triview_asset", "metadata", "updated_at"])
|
||||
|
||||
task.status = AITask.Status.SUCCEEDED
|
||||
task.response_payload = response
|
||||
task.actual_cost = task.estimated_cost
|
||||
task.completed_at = timezone.now()
|
||||
task.save(update_fields=["status", "response_payload", "actual_cost", "completed_at", "updated_at"])
|
||||
charge_reserved_credit(reservation=reservation, actual_amount=task.actual_cost)
|
||||
|
||||
from apps.assets.review import submit_asset_for_review
|
||||
|
||||
transaction.on_commit(lambda asset=triview: submit_asset_for_review(asset))
|
||||
except Exception as exc: # noqa: BLE001
|
||||
with transaction.atomic():
|
||||
locked_task = AITask.objects.select_for_update().get(id=task.id)
|
||||
if locked_task.status == AITask.Status.SUCCEEDED:
|
||||
return
|
||||
locked_task.status = AITask.Status.FAILED
|
||||
locked_task.error_message = str(exc)
|
||||
locked_task.completed_at = timezone.now()
|
||||
locked_task.save(update_fields=["status", "error_message", "completed_at", "updated_at"])
|
||||
release_credit(reservation=reservation, reason=str(exc))
|
||||
|
||||
|
||||
def run_triview_task(*, task_id: str) -> None:
|
||||
"""Celery worker 内执行三视图慢出图(image_edit 以立绘为参考):成功落库扣费并归到立绘的三视图组 / 失败退费。
|
||||
幂等:只处理 RESERVED 任务,重复投递不会二次出图、二次扣费。"""
|
||||
@@ -1521,26 +1706,6 @@ def run_triview_task(*, task_id: str) -> None:
|
||||
group.candidate_assets.add(asset)
|
||||
group.adopted_asset = asset
|
||||
group.save(update_fields=["adopted_asset", "updated_at"])
|
||||
# B 路(期3):视频流程新生成的「角色」立绘 + 三视图成套 → 自动加入/更新模特库(团队级可复用)。
|
||||
# 幂等:同一立绘只建一次 Model;已存在该立绘的 Model(尤其「真人上传」模特,upload 时只建了 portrait、
|
||||
# 还没 triview)→ 回写 triview_asset,否则模特库永远显示「无三视图」(ZWQ#7:AI 角色能更新、真人上传不更新)。
|
||||
from apps.assets.models import Model as ModelEntity
|
||||
|
||||
portrait = Asset.objects.filter(id=asset_key).first()
|
||||
if portrait is not None:
|
||||
existing_model = ModelEntity.objects.filter(portrait_asset=portrait, is_deleted=False).first()
|
||||
if existing_model is None:
|
||||
ModelEntity.objects.create(
|
||||
team=project.team, created_by=user,
|
||||
name=(portrait.name or f"{project.name}-角色").split("·")[0].strip()[:255] or "角色",
|
||||
source=ModelEntity.Source.AI,
|
||||
portrait_asset=portrait, triview_asset=asset,
|
||||
metadata={"from_project": str(project.id), "auto_enrolled": True},
|
||||
)
|
||||
elif existing_model.triview_asset_id != asset.id:
|
||||
# 真人上传 / 已有模特补三视图:回写新三视图(应用后模特库即刷新三视图)
|
||||
existing_model.triview_asset = asset
|
||||
existing_model.save(update_fields=["triview_asset", "updated_at"])
|
||||
# 合规(期3):三视图含人脸,视频生成前必须过火山审核 → 事务提交后静默送审(best-effort)
|
||||
from apps.assets.review import submit_asset_for_review
|
||||
|
||||
@@ -2440,11 +2605,11 @@ def create_export_job(*, timeline, user) -> ExportJob:
|
||||
return ExportJob.objects.create(timeline=timeline, status=ExportJob.Status.QUEUED)
|
||||
|
||||
|
||||
# 图片趴三类(模特库+资产模型重构 期2):
|
||||
# · model + product → model_tryon(模特上身图);model 无 product → person(视频角色「生成演员」,留期3)
|
||||
# 图片趴三类(模特库+资产模型重构):
|
||||
# · model + product → model_tryon(模特上身图);model 无 product → model_portrait(新建模特候选)
|
||||
# · cover → platform_kit(平台套图);image → free_create(自由创作)
|
||||
_STANDALONE_CATEGORY = {
|
||||
"model": Asset.Category.PERSON,
|
||||
"model": Asset.Category.MODEL_PORTRAIT,
|
||||
"cover": Asset.Category.PLATFORM_KIT,
|
||||
"image": Asset.Category.FREE_CREATE,
|
||||
}
|
||||
@@ -2573,7 +2738,7 @@ def run_standalone_image_task(*, task_id: str) -> None:
|
||||
index = int(payload.get("index") or 0)
|
||||
product_id = payload.get("product_id") or None
|
||||
# 模特上身图(mode=model 且绑了商品)= 图片趴「模特上身图」(引用模特库),归 model_tryon、不送审;
|
||||
# 「生成演员」同样走 mode=model 但无 product_id = 视频角色,仍归 person(送审,留期3 收编为「角色」)。
|
||||
# 无商品的 mode=model 仅用于新建模特候选,归 model_portrait;项目角色只由项目基础资产流程创建 person。
|
||||
if mode == "model" and product_id:
|
||||
category = Asset.Category.MODEL_TRYON
|
||||
else:
|
||||
@@ -2695,11 +2860,9 @@ def run_standalone_image_task(*, task_id: str) -> None:
|
||||
id=asset_id, team=team, created_by=user, name=f"AI 生成 · {asset_label} · {index + 1}",
|
||||
asset_type=Asset.Type.IMAGE, source=Asset.Source.AI_GENERATED, category=asset_category, origin_task=task,
|
||||
metadata=asset_meta,
|
||||
# R109:图片生成(模特上身图/平台套图/图片创作/商品三视图)的成图**自动入资产库**,
|
||||
# 不再由用户点「加入资产库」手动入库。仅「生成演员」(mode=model 无商品 → person)
|
||||
# 仍暂存(in_library=False),由演员库「保存人物」时入库——否则未保存的立绘会
|
||||
# 混进「我的演员」列表(ZWQ#6)。
|
||||
in_library=asset_category != Asset.Category.PERSON,
|
||||
# 模特候选与项目角色均是功能性资料:候选保存后进入模特库,不进入 /library;
|
||||
# 商品/创作等图片趴成图保持原有自动入库行为。
|
||||
in_library=asset_category not in (Asset.Category.PERSON, Asset.Category.MODEL_PORTRAIT),
|
||||
)
|
||||
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)
|
||||
except Exception as exc: # noqa: BLE001 — 失败要退费并把错误记进 AITask 供前端轮询读取;不向上抛(避免 celery 重试二次扣费)
|
||||
|
||||
@@ -51,6 +51,15 @@ def generate_triview_task(self, task_id: str) -> str:
|
||||
return task_id
|
||||
|
||||
|
||||
@app.task(bind=True, max_retries=0)
|
||||
def generate_model_triview_task(self, task_id: str) -> str:
|
||||
"""团队模特库三视图:不绑定项目,成功更新 Model.triview_asset,失败释放预留积分。"""
|
||||
from apps.ai.services import run_model_triview_task
|
||||
|
||||
run_model_triview_task(task_id=task_id)
|
||||
return task_id
|
||||
|
||||
|
||||
@app.task(bind=True, max_retries=0)
|
||||
def poll_free_video_task(self, task_id: str, attempt: int = 0) -> str:
|
||||
"""自由创作视频·worker 兜底轮询:每 30s 一次自重排(不依赖 celery beat),
|
||||
|
||||
+213
-20
@@ -182,7 +182,7 @@ from apps.ai.models import AITask, ImageConversation, ModelConfig, ModelProvider
|
||||
from apps.ai.providers.volcano import VolcanoArkProvider
|
||||
from apps.ai.services import enqueue_standalone_images
|
||||
from apps.assets.models import Asset, AssetFile
|
||||
from apps.billing.models import CreditAccount
|
||||
from apps.billing.models import CreditAccount, CreditLedger
|
||||
from apps.products.models import Product
|
||||
|
||||
|
||||
@@ -742,17 +742,15 @@ class StandaloneCategoryTests(TestCase):
|
||||
self.assertEqual(a.category, Asset.Category.FREE_CREATE)
|
||||
self.assertTrue(a.in_library) # R109
|
||||
|
||||
def test_generate_actor_stays_person(self):
|
||||
# 生成演员:mode=model 但无 product → 视频角色,仍 person(送审范围内)
|
||||
def test_generate_model_candidate_is_model_portrait_and_not_in_asset_library(self):
|
||||
# 图片生成 mode=model 且无商品 = 新建模特候选,不再生成“演员”person 资产。
|
||||
a = self._run("model")
|
||||
self.assertEqual(a.category, Asset.Category.PERSON)
|
||||
# R109 例外:演员立绘仍暂存(不自动入库),由演员库「保存人物」时置 True——
|
||||
# 否则未保存立绘会混进「我的演员」列表(ZWQ#6)
|
||||
self.assertEqual(a.category, Asset.Category.MODEL_PORTRAIT)
|
||||
self.assertFalse(a.in_library)
|
||||
|
||||
|
||||
class TriviewAutoEnrollTests(TestCase):
|
||||
"""B 路(期3):视频流程新生成的角色立绘 + 三视图成套 → 自动入模特库;幂等不重复建。"""
|
||||
class TriviewModelDecouplingTests(TestCase):
|
||||
"""项目角色三视图只归项目,不自动创建或更新团队模特。"""
|
||||
|
||||
def setUp(self):
|
||||
from apps.projects.models import Project
|
||||
@@ -771,6 +769,7 @@ class TriviewAutoEnrollTests(TestCase):
|
||||
def _patch_provider(self):
|
||||
provider = patch("apps.ai.services.get_image_provider").start()
|
||||
prov = provider.return_value
|
||||
prov.image_generation.return_value = {"data": [{"url": "http://x/person.png"}]}
|
||||
prov.image_edit.return_value = {"data": [{"url": "http://x/tri.png"}]}
|
||||
prov.extract_first_media_url.return_value = "http://x/tri.png"
|
||||
media = patch("apps.ai.services.VolcanoArkProvider.media_to_bytes").start()
|
||||
@@ -786,27 +785,221 @@ class TriviewAutoEnrollTests(TestCase):
|
||||
task = generate_person_triview(project=self.project, user=self.user, portrait_asset=self.portrait)
|
||||
run_triview_task(task_id=str(task.id))
|
||||
|
||||
def test_triview_auto_enrolls_model(self):
|
||||
def test_triview_stays_on_project_and_does_not_enroll_model(self):
|
||||
from apps.assets.models import Model
|
||||
|
||||
self._patch_provider()
|
||||
self._gen_triview()
|
||||
m = Model.objects.filter(portrait_asset=self.portrait).first()
|
||||
self.assertIsNotNone(m)
|
||||
self.assertIsNotNone(m.triview_asset) # 成套
|
||||
self.assertTrue(m.metadata.get("auto_enrolled"))
|
||||
self.assertEqual(m.source, Model.Source.AI)
|
||||
# 合规(期3):三视图归 tri_view → 落在送审范围内
|
||||
self.assertEqual(m.triview_asset.category, Asset.Category.TRI_VIEW)
|
||||
self.assertIn(m.triview_asset.category, Asset.REVIEW_CATEGORIES)
|
||||
self.assertFalse(Model.objects.filter(portrait_asset=self.portrait).exists())
|
||||
group = next(g for g in self.project.base_asset_groups.all() if (g.metadata or {}).get("triview_of") == str(self.portrait.id))
|
||||
self.assertIsNotNone(group.adopted_asset)
|
||||
self.assertEqual(group.adopted_asset.category, Asset.Category.TRI_VIEW)
|
||||
self.assertIn(group.adopted_asset.category, Asset.REVIEW_CATEGORIES)
|
||||
|
||||
def test_auto_enroll_is_idempotent(self):
|
||||
def test_triview_does_not_overwrite_existing_model_triview(self):
|
||||
from apps.assets.models import Model
|
||||
|
||||
existing_tri = Asset.objects.create(
|
||||
team=self.team, created_by=self.user, name="模特原三视图", asset_type=Asset.Type.IMAGE,
|
||||
source=Asset.Source.AI_GENERATED, category=Asset.Category.TRI_VIEW,
|
||||
)
|
||||
model = Model.objects.create(
|
||||
team=self.team, created_by=self.user, name="已有模特", portrait_asset=self.portrait, triview_asset=existing_tri,
|
||||
)
|
||||
self._patch_provider()
|
||||
self._gen_triview()
|
||||
self._gen_triview() # 同一立绘再出一版三视图
|
||||
self.assertEqual(Model.objects.filter(portrait_asset=self.portrait).count(), 1)
|
||||
model.refresh_from_db()
|
||||
self.assertEqual(model.triview_asset_id, existing_tri.id)
|
||||
|
||||
@patch("apps.ai.services.generate_person_triview")
|
||||
def test_person_base_asset_ignores_legacy_auto_triview_flag(self, kickoff_triview):
|
||||
"""即使旧前端/旧任务传 auto_triview=True,角色立绘完成后也不能自动接力三视图。"""
|
||||
from apps.ai.services import generate_base_asset
|
||||
from apps.projects.models import BaseAssetGroup
|
||||
|
||||
ModelConfig.objects.create(
|
||||
provider=ModelProvider.objects.create(name="auto-tri-provider", display_name="Auto Tri Provider"),
|
||||
name="auto-tri-img",
|
||||
display_name="Auto Tri Img",
|
||||
capability=ModelConfig.Capability.IMAGE,
|
||||
endpoint="images/generations",
|
||||
unit_price="1.0000",
|
||||
is_default=True,
|
||||
)
|
||||
self._patch_provider()
|
||||
|
||||
task = generate_base_asset(
|
||||
project=self.project,
|
||||
user=self.user,
|
||||
kind=BaseAssetGroup.Kind.PERSON,
|
||||
prompt="项目角色立绘",
|
||||
label="女主",
|
||||
auto_triview=True,
|
||||
)
|
||||
|
||||
task.refresh_from_db()
|
||||
self.assertFalse(task.request_payload.get("auto_triview"))
|
||||
kickoff_triview.assert_not_called()
|
||||
|
||||
|
||||
class ModelLibraryTriviewTaskTests(TestCase):
|
||||
"""团队模特三视图:无项目归属、20 积分闭环、幂等防重复、只更新当前 Model。"""
|
||||
|
||||
def setUp(self):
|
||||
from apps.assets.models import Model
|
||||
|
||||
self.user = User.objects.create_user(username="model-tri", password="p")
|
||||
self.team = Team.objects.create(name="Model Tri", 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)
|
||||
self.account = CreditAccount.objects.create(team=self.team, balance="100.0000")
|
||||
provider = ModelProvider.objects.create(
|
||||
name="model-tri-provider",
|
||||
display_name="Model Tri Provider",
|
||||
base_url="http://x.test/v1",
|
||||
api_key="test-key",
|
||||
)
|
||||
self.config = ModelConfig.objects.create(
|
||||
provider=provider,
|
||||
name="model-tri-image",
|
||||
display_name="Model Tri Image",
|
||||
capability=ModelConfig.Capability.IMAGE,
|
||||
endpoint="images/edits",
|
||||
unit_price="20.0000",
|
||||
is_default=True,
|
||||
)
|
||||
self.portrait = Asset.objects.create(
|
||||
team=self.team,
|
||||
created_by=self.user,
|
||||
name="模特形象",
|
||||
asset_type=Asset.Type.IMAGE,
|
||||
source=Asset.Source.UPLOAD,
|
||||
category=Asset.Category.MODEL_PORTRAIT,
|
||||
in_library=False,
|
||||
)
|
||||
AssetFile.objects.create(
|
||||
asset=self.portrait,
|
||||
object_key="portrait.png",
|
||||
bucket="b",
|
||||
content_type="image/png",
|
||||
preview_url="http://x/portrait.png",
|
||||
is_primary=True,
|
||||
)
|
||||
self.model = Model.objects.create(
|
||||
team=self.team,
|
||||
created_by=self.user,
|
||||
name="团队模特",
|
||||
portrait_asset=self.portrait,
|
||||
)
|
||||
|
||||
def _submit(self):
|
||||
from apps.ai.services import generate_model_triview
|
||||
|
||||
with patch("apps.ai.tasks.generate_model_triview_task.delay") as delay:
|
||||
result = generate_model_triview(model=self.model, user=self.user)
|
||||
return result, delay
|
||||
|
||||
def _provider_success(self):
|
||||
provider = patch("apps.ai.services.get_image_provider").start().return_value
|
||||
provider.image_edit.return_value = {"data": [{"url": "http://x/tri.png"}]}
|
||||
provider.extract_first_media_url.return_value = "http://x/tri.png"
|
||||
patch("apps.ai.services.VolcanoArkProvider.media_to_bytes", return_value=(BytesIO(b"img"), "image/png")).start()
|
||||
storage = patch("apps.ai.services.TosStorage").start()
|
||||
stored = storage.return_value.upload_fileobj.return_value
|
||||
stored.object_key, stored.bucket, stored.content_type, stored.size_bytes = "tri.png", "b", "image/png", 3
|
||||
self.addCleanup(patch.stopall)
|
||||
|
||||
def test_submit_is_projectless_reserves_20_and_reuses_inflight(self):
|
||||
(task, created), delay = self._submit()
|
||||
self.assertTrue(created)
|
||||
self.assertIsNone(task.project_id)
|
||||
self.assertEqual(task.task_type, AITask.Type.MODEL_TRIVIEW)
|
||||
self.assertEqual(float(task.estimated_cost), 20.0)
|
||||
self.assertEqual(task.request_payload["model_id"], str(self.model.id))
|
||||
delay.assert_called_once_with(str(task.id))
|
||||
self.account.refresh_from_db()
|
||||
self.assertEqual(float(self.account.reserved_balance), 20.0)
|
||||
reserve = CreditLedger.objects.get(task=task, ledger_type=CreditLedger.Type.RESERVE)
|
||||
self.assertEqual(reserve.reason, "AI 任务预扣额度")
|
||||
|
||||
(same_task, created_again), second_delay = self._submit()
|
||||
self.assertFalse(created_again)
|
||||
self.assertEqual(same_task.id, task.id)
|
||||
second_delay.assert_not_called()
|
||||
self.assertEqual(AITask.objects.filter(task_type=AITask.Type.MODEL_TRIVIEW).count(), 1)
|
||||
|
||||
def test_success_charges_once_and_updates_only_model_triview(self):
|
||||
from apps.ai.services import run_model_triview_task
|
||||
|
||||
(task, _created), _delay = self._submit()
|
||||
self._provider_success()
|
||||
run_model_triview_task(task_id=str(task.id))
|
||||
|
||||
task.refresh_from_db()
|
||||
self.model.refresh_from_db()
|
||||
self.account.refresh_from_db()
|
||||
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
|
||||
self.assertEqual(float(task.actual_cost), 20.0)
|
||||
self.assertEqual(float(self.account.balance), 80.0)
|
||||
self.assertEqual(float(self.account.reserved_balance), 0.0)
|
||||
self.assertIsNotNone(self.model.triview_asset_id)
|
||||
self.assertEqual(self.model.triview_asset.category, Asset.Category.TRI_VIEW)
|
||||
self.assertFalse(self.model.triview_asset.in_library)
|
||||
self.assertEqual(self.model.triview_asset.metadata["model_id"], str(self.model.id))
|
||||
self.assertEqual(self.model.metadata["triview_versions"], [str(self.model.triview_asset_id)])
|
||||
charge = CreditLedger.objects.get(task=task, ledger_type=CreditLedger.Type.CHARGE)
|
||||
self.assertEqual(charge.reason, "AI 任务扣费")
|
||||
|
||||
run_model_triview_task(task_id=str(task.id))
|
||||
self.assertEqual(CreditLedger.objects.filter(task=task, ledger_type=CreditLedger.Type.CHARGE).count(), 1)
|
||||
|
||||
def test_failure_releases_credit_and_keeps_old_triview(self):
|
||||
from apps.ai.services import run_model_triview_task
|
||||
|
||||
old_tri = Asset.objects.create(
|
||||
team=self.team, created_by=self.user, name="旧三视图", asset_type=Asset.Type.IMAGE,
|
||||
source=Asset.Source.AI_GENERATED, category=Asset.Category.TRI_VIEW, in_library=False,
|
||||
)
|
||||
self.model.triview_asset = old_tri
|
||||
self.model.save(update_fields=["triview_asset"])
|
||||
(task, _created), _delay = self._submit()
|
||||
provider = patch("apps.ai.services.get_image_provider").start().return_value
|
||||
provider.image_edit.side_effect = ValueError("生成失败")
|
||||
self.addCleanup(patch.stopall)
|
||||
|
||||
run_model_triview_task(task_id=str(task.id))
|
||||
|
||||
task.refresh_from_db()
|
||||
self.model.refresh_from_db()
|
||||
self.account.refresh_from_db()
|
||||
self.assertEqual(task.status, AITask.Status.FAILED)
|
||||
self.assertEqual(self.model.triview_asset_id, old_tri.id)
|
||||
self.assertEqual(float(self.account.balance), 100.0)
|
||||
self.assertEqual(float(self.account.reserved_balance), 0.0)
|
||||
self.assertFalse(CreditLedger.objects.filter(task=task, ledger_type=CreditLedger.Type.CHARGE).exists())
|
||||
self.assertTrue(CreditLedger.objects.filter(task=task, ledger_type=CreditLedger.Type.RELEASE).exists())
|
||||
|
||||
def test_api_quote_submit_and_status_contract(self):
|
||||
quote = self.client.get(f"/api/models/{self.model.id}/triview-quote/")
|
||||
self.assertEqual(quote.status_code, 200)
|
||||
self.assertEqual(float(quote.json()["points"]), 20.0)
|
||||
|
||||
disabled = self.client.post(f"/api/models/{self.model.id}/generate-triview/")
|
||||
self.assertEqual(disabled.status_code, 409)
|
||||
self.assertFalse(AITask.objects.filter(task_type=AITask.Type.MODEL_TRIVIEW).exists())
|
||||
|
||||
with self.settings(MODEL_TRIVIEW_GENERATION_ENABLED=True):
|
||||
with patch("apps.ai.tasks.generate_model_triview_task.delay"):
|
||||
submit = self.client.post(f"/api/models/{self.model.id}/generate-triview/")
|
||||
self.assertEqual(submit.status_code, 202)
|
||||
task_id = submit.json()["task_id"]
|
||||
self.assertEqual(float(submit.json()["estimated_cost"]), 20.0)
|
||||
|
||||
status_res = self.client.get(f"/api/models/{self.model.id}/triview-status/?task_id={task_id}")
|
||||
self.assertEqual(status_res.status_code, 200)
|
||||
self.assertEqual(status_res.json()["status"], AITask.Status.RESERVED)
|
||||
self.assertIsNone(status_res.json()["model"])
|
||||
|
||||
|
||||
class _FakeStreamResp:
|
||||
|
||||
@@ -0,0 +1,86 @@
|
||||
"""审计“我的演员”收口为模特前的历史数据;只读,不写库。"""
|
||||
|
||||
from django.core.management.base import BaseCommand
|
||||
from django.db.models import Q
|
||||
|
||||
from apps.assets.models import Asset, Model
|
||||
from apps.projects.models import BaseAssetGroup
|
||||
|
||||
|
||||
class Command(BaseCommand):
|
||||
help = "只读审计旧 person 资产,输出可安全迁移为 Model 的候选和排除原因。"
|
||||
|
||||
def add_arguments(self, parser):
|
||||
parser.add_argument("--team-id", help="仅审计指定团队 UUID")
|
||||
parser.add_argument("--limit", type=int, default=50, help="最多打印多少条候选(默认 50)")
|
||||
|
||||
def handle(self, *args, **options):
|
||||
team_id = options.get("team_id")
|
||||
limit = max(0, options["limit"])
|
||||
|
||||
active_people = Asset.objects.filter(
|
||||
category=Asset.Category.PERSON,
|
||||
is_deleted=False,
|
||||
purged_at__isnull=True,
|
||||
)
|
||||
if team_id:
|
||||
active_people = active_people.filter(team_id=team_id)
|
||||
|
||||
# 项目基础资产组引用的 person(包含采用版和候选版)一律视为项目角色,不自动迁移。
|
||||
project_role_asset_ids = set(
|
||||
Asset.objects.filter(
|
||||
Q(adopted_base_groups__kind=BaseAssetGroup.Kind.PERSON)
|
||||
| Q(candidate_base_groups__kind=BaseAssetGroup.Kind.PERSON)
|
||||
)
|
||||
.values_list("id", flat=True)
|
||||
.distinct()
|
||||
)
|
||||
if team_id:
|
||||
project_role_asset_ids &= set(active_people.values_list("id", flat=True))
|
||||
|
||||
existing_model_portraits = set(
|
||||
Model.objects.filter(
|
||||
is_deleted=False,
|
||||
purged_at__isnull=True,
|
||||
portrait_asset__isnull=False,
|
||||
)
|
||||
.values_list("portrait_asset_id", flat=True)
|
||||
)
|
||||
if team_id:
|
||||
existing_model_portraits &= set(active_people.values_list("id", flat=True))
|
||||
|
||||
saved_people = active_people.filter(in_library=True)
|
||||
project_generated_ids = set(
|
||||
active_people.filter(origin_task__project__isnull=False).values_list("id", flat=True)
|
||||
)
|
||||
eligible = (
|
||||
saved_people.exclude(id__in=project_role_asset_ids)
|
||||
.exclude(id__in=project_generated_ids)
|
||||
.exclude(id__in=existing_model_portraits)
|
||||
.order_by("created_at")
|
||||
)
|
||||
|
||||
trash_people = Asset.objects.filter(
|
||||
category=Asset.Category.PERSON,
|
||||
is_deleted=True,
|
||||
purged_at__isnull=True,
|
||||
)
|
||||
purged_people = Asset.objects.filter(category=Asset.Category.PERSON, purged_at__isnull=False)
|
||||
if team_id:
|
||||
trash_people = trash_people.filter(team_id=team_id)
|
||||
purged_people = purged_people.filter(team_id=team_id)
|
||||
|
||||
self.stdout.write(self.style.WARNING("=== audit_actor_model_migration · DRY-RUN / READ-ONLY ==="))
|
||||
self.stdout.write(f"正常 person 资产: {active_people.count()}")
|
||||
self.stdout.write(f" 已有 Model 形象图: {len(existing_model_portraits)}")
|
||||
self.stdout.write(f" 项目角色引用: {len(project_role_asset_ids)}")
|
||||
self.stdout.write(f" 项目生成角色: {len(project_generated_ids)}")
|
||||
self.stdout.write(f" 旧已保存人物候选: {eligible.count()}")
|
||||
self.stdout.write(f"垃圾桶 person(不迁移): {trash_people.count()}")
|
||||
self.stdout.write(f"彻底隐藏 person(不迁移): {purged_people.count()}")
|
||||
self.stdout.write("\n可迁移候选(只打印,不创建 Model):")
|
||||
for asset in eligible[:limit]:
|
||||
self.stdout.write(f" {asset.id} | team={asset.team_id} | {asset.name}")
|
||||
if eligible.count() > limit:
|
||||
self.stdout.write(f" … 其余 {eligible.count() - limit} 条未打印(可用 --limit 调整)")
|
||||
self.stdout.write(self.style.WARNING("审计结束:未创建、更新、删除任何记录或文件。"))
|
||||
@@ -0,0 +1,80 @@
|
||||
"""把符合规则的历史“已保存演员”补为 Model;默认 dry-run。"""
|
||||
|
||||
from django.core.management.base import BaseCommand
|
||||
from django.db import transaction
|
||||
from django.db.models import Q
|
||||
|
||||
from apps.assets.models import Asset, Model
|
||||
from apps.projects.models import BaseAssetGroup
|
||||
|
||||
|
||||
def eligible_legacy_actor_assets(team_id=None):
|
||||
"""只返回正常状态、已保存、非项目角色且尚未拥有 Model 的旧 person 资产。"""
|
||||
people = Asset.objects.filter(
|
||||
category=Asset.Category.PERSON,
|
||||
in_library=True,
|
||||
is_deleted=False,
|
||||
purged_at__isnull=True,
|
||||
)
|
||||
if team_id:
|
||||
people = people.filter(team_id=team_id)
|
||||
|
||||
project_role_ids = Asset.objects.filter(
|
||||
Q(adopted_base_groups__kind=BaseAssetGroup.Kind.PERSON)
|
||||
| Q(candidate_base_groups__kind=BaseAssetGroup.Kind.PERSON)
|
||||
).values_list("id", flat=True)
|
||||
existing_model_ids = Model.objects.filter(
|
||||
is_deleted=False,
|
||||
purged_at__isnull=True,
|
||||
portrait_asset__isnull=False,
|
||||
).values_list("portrait_asset_id", flat=True)
|
||||
return (
|
||||
people.exclude(origin_task__project__isnull=False)
|
||||
.exclude(id__in=project_role_ids)
|
||||
.exclude(id__in=existing_model_ids)
|
||||
.order_by("created_at")
|
||||
)
|
||||
|
||||
|
||||
class Command(BaseCommand):
|
||||
help = "迁移正常状态的历史已保存演员为 Model;默认 dry-run,--apply 才写库。"
|
||||
|
||||
def add_arguments(self, parser):
|
||||
parser.add_argument("--team-id", help="仅迁移指定团队 UUID")
|
||||
parser.add_argument("--apply", action="store_true", help="真正创建 Model(默认只打印)")
|
||||
parser.add_argument("--limit", type=int, default=50, help="最多打印多少条候选(默认 50)")
|
||||
|
||||
def handle(self, *args, **options):
|
||||
apply = options["apply"]
|
||||
candidates = eligible_legacy_actor_assets(options.get("team_id"))
|
||||
count = candidates.count()
|
||||
limit = max(0, options["limit"])
|
||||
self.stdout.write(self.style.WARNING(f"=== migrate_saved_actors_to_models · {'APPLY' if apply else 'DRY-RUN'} ==="))
|
||||
self.stdout.write(f"可迁移旧保存人物: {count}")
|
||||
for asset in candidates[:limit]:
|
||||
self.stdout.write(f" {asset.id} | team={asset.team_id} | {asset.name}")
|
||||
|
||||
if not apply:
|
||||
self.stdout.write(self.style.WARNING("DRY-RUN 结束,未写库。加 --apply 才会创建 Model。"))
|
||||
return
|
||||
|
||||
created = 0
|
||||
with transaction.atomic():
|
||||
for asset in candidates.select_for_update():
|
||||
# select_for_update 后再次确认,保证并发/重复执行都不会造成重复 Model。
|
||||
if Model.objects.filter(
|
||||
portrait_asset=asset,
|
||||
is_deleted=False,
|
||||
purged_at__isnull=True,
|
||||
).exists():
|
||||
continue
|
||||
Model.objects.create(
|
||||
team_id=asset.team_id,
|
||||
created_by_id=asset.created_by_id,
|
||||
name=(asset.name or "模特")[:255],
|
||||
source=Model.Source.UPLOAD if asset.source == Asset.Source.UPLOAD else Model.Source.AI,
|
||||
portrait_asset=asset,
|
||||
metadata={"migrated_from_legacy_actor": True},
|
||||
)
|
||||
created += 1
|
||||
self.stdout.write(self.style.SUCCESS(f"完成:创建 {created} 个 Model;未修改任何 Asset、项目、垃圾桶或文件。"))
|
||||
@@ -1,7 +1,8 @@
|
||||
from io import BytesIO
|
||||
from io import BytesIO, StringIO
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import patch
|
||||
|
||||
from django.core.management import call_command
|
||||
from django.core.files.uploadedfile import SimpleUploadedFile
|
||||
from django.test import SimpleTestCase, TestCase
|
||||
from rest_framework.test import APIClient
|
||||
@@ -44,10 +45,60 @@ class ModelLibraryApiTests(TestCase):
|
||||
mine = {m["id"] for m in self.client.get("/api/models/?tab=mine").json()["results"]}
|
||||
self.assertEqual(mine, {str(self.mine.id)})
|
||||
|
||||
def test_filter_by_portrait_asset(self):
|
||||
portrait = Asset.objects.create(team=self.teamA, name="角色立绘", asset_type="image", source="ai_generated")
|
||||
self.mine.portrait_asset = portrait
|
||||
self.mine.save(update_fields=["portrait_asset"])
|
||||
|
||||
rows = self.client.get(f"/api/models/?portrait_asset={portrait.id}").json()["results"]
|
||||
self.assertEqual([row["id"] for row in rows], [str(self.mine.id)])
|
||||
|
||||
missing = self.client.get("/api/models/?portrait_asset=00000000-0000-0000-0000-000000000000").json()["results"]
|
||||
self.assertEqual(missing, [])
|
||||
|
||||
def test_enroll_asset_creates_model_once_and_reuses_it(self):
|
||||
portrait = Asset.objects.create(
|
||||
team=self.teamA, name="AI 模特候选", asset_type="image", source="ai_generated", category=Asset.Category.PERSON,
|
||||
)
|
||||
|
||||
first = self.client.post("/api/models/enroll-asset/", {"asset_id": str(portrait.id), "name": "通勤模特"}, format="json")
|
||||
self.assertEqual(first.status_code, 201)
|
||||
self.assertEqual(first.json()["name"], "通勤模特")
|
||||
self.assertEqual(first.json()["portrait_asset"], str(portrait.id))
|
||||
|
||||
second = self.client.post("/api/models/enroll-asset/", {"asset_id": str(portrait.id), "name": "重复名称不应覆盖"}, format="json")
|
||||
self.assertEqual(second.status_code, 200)
|
||||
self.assertEqual(second.json()["id"], first.json()["id"])
|
||||
self.assertEqual(Model.objects.filter(team=self.teamA, portrait_asset=portrait, is_deleted=False).count(), 1)
|
||||
|
||||
def test_official_flag_in_payload(self):
|
||||
row = next(m for m in self.client.get("/api/models/").json()["results"] if m["id"] == str(self.official.id))
|
||||
self.assertTrue(row["is_official"])
|
||||
|
||||
def test_update_own_model_profile_does_not_rewrite_portrait_asset(self):
|
||||
portrait = Asset.objects.create(team=self.teamA, name="原始形象资产名", asset_type="image", source="upload")
|
||||
self.mine.portrait_asset = portrait
|
||||
self.mine.save(update_fields=["portrait_asset"])
|
||||
|
||||
res = self.client.patch(
|
||||
f"/api/models/{self.mine.id}/",
|
||||
{"name": "更新后的模特名", "description": "通勤风格"},
|
||||
format="json",
|
||||
)
|
||||
|
||||
self.assertEqual(res.status_code, 200)
|
||||
self.mine.refresh_from_db()
|
||||
portrait.refresh_from_db()
|
||||
self.assertEqual(self.mine.name, "更新后的模特名")
|
||||
self.assertEqual(self.mine.description, "通勤风格")
|
||||
self.assertEqual(portrait.name, "原始形象资产名")
|
||||
|
||||
def test_official_model_cannot_be_edited(self):
|
||||
res = self.client.patch(f"/api/models/{self.official.id}/", {"name": "不应更新"}, format="json")
|
||||
self.assertEqual(res.status_code, 403)
|
||||
self.official.refresh_from_db()
|
||||
self.assertEqual(self.official.name, "官方小美")
|
||||
|
||||
def test_delete_is_soft(self):
|
||||
res = self.client.delete(f"/api/models/{self.mine.id}/")
|
||||
self.assertEqual(res.status_code, 204)
|
||||
@@ -129,6 +180,103 @@ class ModelLibraryApiTests(TestCase):
|
||||
self.assertIsNotNone(model.portrait_asset)
|
||||
self.assertEqual(model.portrait_asset.category, Asset.Category.MODEL_PORTRAIT)
|
||||
|
||||
def test_upload_portrait_replaces_current_model_only_and_stays_out_of_library(self):
|
||||
old_portrait = Asset.objects.create(
|
||||
team=self.teamA, name="旧形象", asset_type="image", source="upload", category=Asset.Category.MODEL_PORTRAIT,
|
||||
)
|
||||
self.mine.portrait_asset = old_portrait
|
||||
self.mine.save(update_fields=["portrait_asset"])
|
||||
stored = SimpleNamespace(object_key="teams/x/models/new.png", bucket="b", content_type="image/png", size_bytes=10)
|
||||
|
||||
with patch("apps.assets.views.TosStorage") as Tos:
|
||||
Tos.return_value.upload_fileobj.return_value = stored
|
||||
image = SimpleUploadedFile("new.png", b"img", content_type="image/png")
|
||||
res = self.client.post(
|
||||
f"/api/models/{self.mine.id}/upload-portrait/",
|
||||
{"file": image, "name": "一次保存的新模特名", "description": "一次保存的新描述"},
|
||||
format="multipart",
|
||||
)
|
||||
|
||||
self.assertEqual(res.status_code, 200)
|
||||
self.mine.refresh_from_db()
|
||||
self.assertNotEqual(self.mine.portrait_asset_id, old_portrait.id)
|
||||
self.assertEqual(self.mine.name, "一次保存的新模特名")
|
||||
self.assertEqual(self.mine.description, "一次保存的新描述")
|
||||
new_portrait = self.mine.portrait_asset
|
||||
self.assertEqual(new_portrait.category, Asset.Category.MODEL_PORTRAIT)
|
||||
self.assertFalse(new_portrait.in_library)
|
||||
self.assertTrue(Asset.objects.filter(id=old_portrait.id).exists())
|
||||
self.assertEqual(Model.objects.filter(team=self.teamA).count(), 1)
|
||||
self.assertEqual(
|
||||
self.mine.metadata["portrait_versions"],
|
||||
[str(old_portrait.id), str(new_portrait.id)],
|
||||
)
|
||||
library_ids = {row["id"] for row in self.client.get("/api/assets/?page_size=200").json()["results"]}
|
||||
self.assertNotIn(str(new_portrait.id), library_ids)
|
||||
|
||||
def test_official_model_cannot_replace_portrait(self):
|
||||
with patch("apps.assets.views.TosStorage") as Tos:
|
||||
image = SimpleUploadedFile("new.png", b"img", content_type="image/png")
|
||||
res = self.client.post(f"/api/models/{self.official.id}/upload-portrait/", {"file": image}, format="multipart")
|
||||
self.assertEqual(res.status_code, 403)
|
||||
Tos.assert_not_called()
|
||||
|
||||
|
||||
class ActorModelMigrationAuditTests(TestCase):
|
||||
"""审计命令只能分类和输出,绝不能改写旧资产、项目或模特记录。"""
|
||||
|
||||
def setUp(self):
|
||||
from apps.products.models import Product
|
||||
from apps.projects.models import BaseAssetGroup, Project
|
||||
|
||||
self.user, self.team = _mk_team("audit-u", "Audit")
|
||||
self.eligible = Asset.objects.create(
|
||||
team=self.team, name="旧保存演员", asset_type="image", source="upload", category=Asset.Category.PERSON, in_library=True,
|
||||
)
|
||||
existing = Asset.objects.create(
|
||||
team=self.team, name="已有模特形象", asset_type="image", source="upload", category=Asset.Category.PERSON, in_library=True,
|
||||
)
|
||||
Model.objects.create(team=self.team, name="已有模特", portrait_asset=existing)
|
||||
role = Asset.objects.create(
|
||||
team=self.team, name="项目角色", asset_type="image", source="ai_generated", category=Asset.Category.PERSON, in_library=True,
|
||||
)
|
||||
product = Product.objects.create(team=self.team, title="测试商品")
|
||||
project = Project.objects.create(team=self.team, name="测试项目", product=product)
|
||||
BaseAssetGroup.objects.create(project=project, kind=BaseAssetGroup.Kind.PERSON, adopted_asset=role)
|
||||
Asset.objects.create(
|
||||
team=self.team, name="垃圾桶演员", asset_type="image", source="upload", category=Asset.Category.PERSON, is_deleted=True,
|
||||
)
|
||||
|
||||
def test_audit_is_read_only_and_excludes_project_and_existing_models(self):
|
||||
before_assets = Asset.objects.count()
|
||||
before_models = Model.objects.count()
|
||||
out = StringIO()
|
||||
|
||||
call_command("audit_actor_model_migration", "--team-id", str(self.team.id), stdout=out)
|
||||
|
||||
report = out.getvalue()
|
||||
self.assertIn("旧已保存人物候选: 1", report)
|
||||
self.assertIn(str(self.eligible.id), report)
|
||||
self.assertIn("垃圾桶 person(不迁移): 1", report)
|
||||
self.assertIn("未创建、更新、删除任何记录或文件", report)
|
||||
self.assertEqual(Asset.objects.count(), before_assets)
|
||||
self.assertEqual(Model.objects.count(), before_models)
|
||||
|
||||
def test_apply_migrates_only_eligible_saved_actor_and_is_idempotent(self):
|
||||
call_command("migrate_saved_actors_to_models", "--team-id", str(self.team.id), "--apply")
|
||||
|
||||
model = Model.objects.get(team=self.team, portrait_asset=self.eligible)
|
||||
self.assertEqual(model.name, self.eligible.name)
|
||||
self.assertEqual(model.metadata.get("migrated_from_legacy_actor"), True)
|
||||
self.eligible.refresh_from_db()
|
||||
self.assertEqual(self.eligible.category, Asset.Category.PERSON) # 历史分类不改
|
||||
self.assertFalse(self.eligible.is_deleted)
|
||||
self.assertEqual(Model.objects.filter(team=self.team, portrait_asset=self.eligible).count(), 1)
|
||||
self.assertFalse(Model.objects.filter(team=self.team, portrait_asset__name="项目角色").exists())
|
||||
|
||||
call_command("migrate_saved_actors_to_models", "--team-id", str(self.team.id), "--apply")
|
||||
self.assertEqual(Model.objects.filter(team=self.team, portrait_asset=self.eligible).count(), 1)
|
||||
|
||||
|
||||
class AssetSoftDeleteTests(TestCase):
|
||||
"""资产库不展示软删资产(is_deleted=True):列表 + tab 计数都排除。"""
|
||||
@@ -148,6 +296,22 @@ class AssetSoftDeleteTests(TestCase):
|
||||
def test_summary_excludes_soft_deleted(self):
|
||||
self.assertEqual(self.client.get("/api/assets/summary/").json()["tryon"], 1)
|
||||
|
||||
def test_functional_model_and_project_assets_do_not_appear_in_others(self):
|
||||
model_portrait = Asset.objects.create(
|
||||
team=self.team, name="模特形象", asset_type="image", source="upload", category=Asset.Category.MODEL_PORTRAIT,
|
||||
)
|
||||
tri_view = Asset.objects.create(
|
||||
team=self.team, name="角色三视图", asset_type="image", source="ai_generated", category=Asset.Category.TRI_VIEW,
|
||||
)
|
||||
other = Asset.objects.create(
|
||||
team=self.team, name="其他图片", asset_type="image", source="upload", category=Asset.Category.UNCATEGORIZED,
|
||||
)
|
||||
|
||||
ids = {row["id"] for row in self.client.get("/api/assets/?tab=others&page_size=50").json()["results"]}
|
||||
self.assertIn(str(other.id), ids)
|
||||
self.assertNotIn(str(model_portrait.id), ids)
|
||||
self.assertNotIn(str(tri_view.id), ids)
|
||||
|
||||
|
||||
class AssetTrashApiTests(TestCase):
|
||||
"""R108 资产垃圾桶:trash 只列本团队软删资产;restore 后资产库重新可见;purge 二级软删。"""
|
||||
|
||||
@@ -3,6 +3,7 @@ import re
|
||||
import uuid
|
||||
|
||||
import requests
|
||||
from django.conf import settings
|
||||
from django.db import transaction
|
||||
from django.db.models import Q
|
||||
from django.http import StreamingHttpResponse
|
||||
@@ -36,6 +37,9 @@ class AssetPagination(PageNumberPagination):
|
||||
_KNOWN_CATS = [
|
||||
"person", "scene", "product_image", "final_video", "upload",
|
||||
"model_tryon", "platform_kit", "free_create",
|
||||
# 模特资料、项目角色三视图/分镜是功能性引用资产,只在模特库或项目中管理,
|
||||
# 不应落入 /library 的「其他」形成第二个人物入口。
|
||||
"model_portrait", "tri_view", "storyboard", "voice",
|
||||
]
|
||||
|
||||
|
||||
@@ -445,6 +449,10 @@ class ModelLibraryViewSet(ModelViewSet):
|
||||
qs = qs.filter(is_official=True)
|
||||
elif tab == "mine":
|
||||
qs = qs.filter(team=team, is_official=False)
|
||||
# 项目角色详情用形象资产精确判断「是否已升级为模特」;不靠前端拉全量模特列表猜测,
|
||||
# 避免模特数量超过单页上限时把已升级角色误显示成可重复升级。
|
||||
if portrait_asset_id := self.request.query_params.get("portrait_asset"):
|
||||
qs = qs.filter(portrait_asset_id=portrait_asset_id)
|
||||
if self.request.query_params.get("q"):
|
||||
q = self.request.query_params["q"]
|
||||
qs = qs.filter(Q(name__icontains=q) | Q(description__icontains=q))
|
||||
@@ -454,6 +462,14 @@ class ModelLibraryViewSet(ModelViewSet):
|
||||
def perform_create(self, serializer):
|
||||
serializer.save(team=self.get_team(), created_by=self.request.user)
|
||||
|
||||
def perform_update(self, serializer):
|
||||
instance = serializer.instance
|
||||
if instance.is_official:
|
||||
raise PermissionDenied("官方模特不可编辑")
|
||||
if instance.team_id != self.get_team().id:
|
||||
raise PermissionDenied("无权编辑其他团队的模特")
|
||||
serializer.save()
|
||||
|
||||
def perform_destroy(self, instance):
|
||||
if instance.is_official:
|
||||
raise PermissionDenied("官方模特不可删除")
|
||||
@@ -462,6 +478,77 @@ class ModelLibraryViewSet(ModelViewSet):
|
||||
instance.is_deleted = True
|
||||
instance.save(update_fields=["is_deleted", "updated_at"])
|
||||
|
||||
@action(detail=True, methods=["get"], url_path="triview-quote")
|
||||
def triview_quote(self, request, pk=None):
|
||||
"""返回当前团队实际三视图价格;标准价 20 积分,含团队价格系数。"""
|
||||
model = self.get_object()
|
||||
if model.is_official:
|
||||
raise PermissionDenied("官方模特不可生成三视图")
|
||||
from apps.ai.services import quote_model_triview
|
||||
|
||||
try:
|
||||
_config, quote = quote_model_triview(model=model)
|
||||
except ValueError as exc:
|
||||
raise ValidationError({"detail": str(exc)}) from exc
|
||||
return Response({"points": str(quote.points)}, status=status.HTTP_200_OK)
|
||||
|
||||
@action(detail=True, methods=["post"], url_path="generate-triview")
|
||||
def generate_triview(self, request, pk=None):
|
||||
"""提交团队模特三视图任务;9.4A 仅供接口与自动化测试,前端按钮尚未启用。"""
|
||||
if not getattr(settings, "MODEL_TRIVIEW_GENERATION_ENABLED", False):
|
||||
return Response(
|
||||
{"detail": "模特三视图生成尚未启用,请先同步部署 worker"},
|
||||
status=status.HTTP_409_CONFLICT,
|
||||
)
|
||||
model = self.get_object()
|
||||
if model.is_official:
|
||||
raise PermissionDenied("官方模特不可生成三视图")
|
||||
from apps.ai.services import generate_model_triview
|
||||
|
||||
try:
|
||||
task, created = generate_model_triview(model=model, user=request.user)
|
||||
except ValueError as exc:
|
||||
raise ValidationError({"detail": str(exc)}) from exc
|
||||
return Response(
|
||||
{
|
||||
"task_id": str(task.id),
|
||||
"status": task.status,
|
||||
"estimated_cost": str(task.estimated_cost),
|
||||
"reused": not created,
|
||||
},
|
||||
status=status.HTTP_202_ACCEPTED if created else status.HTTP_200_OK,
|
||||
)
|
||||
|
||||
@action(detail=True, methods=["get"], url_path="triview-status")
|
||||
def triview_status(self, request, pk=None):
|
||||
"""轮询指定模特三视图任务,只允许读取本团队且 payload 归属当前模特的任务。"""
|
||||
model = self.get_object()
|
||||
task_id = request.query_params.get("task_id")
|
||||
if not task_id:
|
||||
raise ValidationError({"task_id": "缺少任务 id"})
|
||||
from apps.ai.models import AITask
|
||||
|
||||
task = AITask.objects.filter(
|
||||
id=task_id,
|
||||
team=self.get_team(),
|
||||
task_type=AITask.Type.MODEL_TRIVIEW,
|
||||
request_payload__model_id=str(model.id),
|
||||
).first()
|
||||
if task is None:
|
||||
raise ValidationError({"task_id": "任务不存在或不属于当前模特"})
|
||||
model.refresh_from_db()
|
||||
return Response(
|
||||
{
|
||||
"task_id": str(task.id),
|
||||
"status": task.status,
|
||||
"estimated_cost": str(task.estimated_cost),
|
||||
"actual_cost": str(task.actual_cost),
|
||||
"error_message": task.error_message,
|
||||
"model": ModelLibrarySerializer(model).data if task.status == AITask.Status.SUCCEEDED else None,
|
||||
},
|
||||
status=status.HTTP_200_OK,
|
||||
)
|
||||
|
||||
@action(detail=False, methods=["get"], url_path="trash")
|
||||
def trash(self, request):
|
||||
"""垃圾桶:列出本团队已软删且未彻底隐藏的自建模特。"""
|
||||
@@ -533,6 +620,64 @@ class ModelLibraryViewSet(ModelViewSet):
|
||||
)
|
||||
return Response(ModelLibrarySerializer(model).data, status=status.HTTP_201_CREATED)
|
||||
|
||||
@action(detail=True, methods=["post"], url_path="upload-portrait", parser_classes=[MultiPartParser, FormParser])
|
||||
def upload_portrait(self, request, pk=None):
|
||||
"""显式维护已有模特的形象图:上传新 model_portrait 并切换当前引用;不新建模特、不回写项目角色。"""
|
||||
model = self.get_object()
|
||||
if model.is_official:
|
||||
raise PermissionDenied("官方模特不可编辑")
|
||||
team = self.get_team()
|
||||
if model.team_id != team.id:
|
||||
raise PermissionDenied("无权编辑其他团队的模特")
|
||||
upload = request.FILES.get("file")
|
||||
if upload is None:
|
||||
raise ValidationError({"file": "请上传一张模特形象图"})
|
||||
next_name = str(request.data.get("name", model.name)).strip()[:255]
|
||||
if not next_name:
|
||||
raise ValidationError({"name": "模特名称不能为空"})
|
||||
next_description = str(request.data.get("description", model.description or "")).strip()
|
||||
|
||||
asset_id = uuid.uuid4()
|
||||
suffix = Path(upload.name).suffix.lower() or ".png"
|
||||
object_key = f"teams/{team.id}/models/{model.id}/portraits/{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"{next_name}·形象图",
|
||||
asset_type=Asset.Type.IMAGE,
|
||||
source=Asset.Source.UPLOAD,
|
||||
category=Asset.Category.MODEL_PORTRAIT,
|
||||
in_library=False,
|
||||
metadata={"kind": "model", "view": "frontal", "source": "upload", "model_id": str(model.id)},
|
||||
)
|
||||
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,
|
||||
)
|
||||
versions = [str(value) for value in (model.metadata or {}).get("portrait_versions", []) if value]
|
||||
for value in (model.portrait_asset_id, portrait.id):
|
||||
if value and str(value) not in versions:
|
||||
versions.append(str(value))
|
||||
metadata = dict(model.metadata or {})
|
||||
metadata["portrait_versions"] = versions
|
||||
model.name = next_name
|
||||
model.description = next_description
|
||||
model.portrait_asset = portrait
|
||||
model.metadata = metadata
|
||||
model.save(update_fields=["name", "description", "portrait_asset", "metadata", "updated_at"])
|
||||
return Response(ModelLibrarySerializer(model).data, status=status.HTTP_200_OK)
|
||||
|
||||
@action(detail=False, methods=["post"], url_path="enroll-asset")
|
||||
def enroll_asset(self, request):
|
||||
"""图片创作「加入模特库」:把一张已生成的图复用为模特形象图建 Model 条目。
|
||||
|
||||
@@ -24,6 +24,7 @@ _STAGE_BUCKET = {
|
||||
AITask.Type.ENTITY_EXTRACTION: "script",
|
||||
AITask.Type.PRODUCT_IMAGE: "base",
|
||||
AITask.Type.PERSON_IMAGE: "base",
|
||||
AITask.Type.MODEL_TRIVIEW: "base",
|
||||
AITask.Type.SCENE_IMAGE: "base",
|
||||
AITask.Type.STORYBOARD: "storyboard",
|
||||
AITask.Type.VIDEO_SEGMENT: "video",
|
||||
|
||||
@@ -349,6 +349,33 @@ class ProjectApiTests(TestCase):
|
||||
group = BaseAssetGroup.objects.get(project=project, kind=BaseAssetGroup.Kind.PERSON)
|
||||
self.assertEqual(group.metadata.get("label"), "女主")
|
||||
|
||||
@patch("apps.ai.services._store_generated_media")
|
||||
@patch("apps.ai.services.get_image_provider")
|
||||
def test_generate_base_asset_ignores_auto_triview_request(self, get_provider, store_media):
|
||||
ModelConfig.objects.create(
|
||||
provider=self.provider, name="img-model-auto-tri", display_name="Img Auto Tri",
|
||||
capability=ModelConfig.Capability.IMAGE, endpoint="images/generations", unit_price="1.0000",
|
||||
)
|
||||
asset = Asset.objects.create(
|
||||
team=self.team, created_by=self.user, name="person",
|
||||
asset_type=Asset.Type.IMAGE, source=Asset.Source.AI_GENERATED, category=Asset.Category.PERSON,
|
||||
)
|
||||
store_media.return_value = asset
|
||||
provider = get_provider.return_value
|
||||
provider.image_generation.return_value = {"data": [{"url": "http://x/img.png"}]}
|
||||
provider.extract_first_media_url.return_value = "http://x/img.png"
|
||||
project = Project.objects.create(team=self.team, created_by=self.user, product=self.product, name="P-auto-tri")
|
||||
|
||||
response = self.client.post(
|
||||
f"/api/projects/{project.id}/generate-base-asset/",
|
||||
{"kind": "person", "prompt": "portrait", "label": "hero", "auto_triview": True},
|
||||
format="json",
|
||||
)
|
||||
|
||||
self.assertEqual(response.status_code, 202)
|
||||
task = AITask.objects.get(id=response.data["task"]["id"])
|
||||
self.assertFalse(task.request_payload.get("auto_triview"))
|
||||
|
||||
@patch("apps.ai.services._store_generated_media")
|
||||
@patch("apps.ai.services.get_image_provider")
|
||||
def test_product_base_asset_uses_cover_image_as_reference(self, get_provider, store_media):
|
||||
|
||||
@@ -542,9 +542,8 @@ class ProjectViewSet(TeamScopedViewSetMixin, ModelViewSet):
|
||||
prompt=request.data.get("prompt", ""),
|
||||
label=request.data.get("label", ""),
|
||||
reference_asset_id=request.data.get("reference_asset_id") or None,
|
||||
# ZWQ#3:角色立绘可让 worker 出图后自动接力生成三视图。前端勾选时传 auto_triview=true;
|
||||
# 不传则维持原行为(只出立绘,用户再手点「AI 生成三视图」)。
|
||||
auto_triview=bool(request.data.get("auto_triview")),
|
||||
# 角色立绘不再自动接力三视图;三视图只由角色详情里的显式按钮生成。
|
||||
auto_triview=False,
|
||||
)
|
||||
except ValueError as exc: # 无可用模型 / 余额不足等,立即反馈
|
||||
return Response({"detail": str(exc)}, status=status.HTTP_400_BAD_REQUEST)
|
||||
|
||||
Reference in New Issue
Block a user