feat: 解耦角色与模特并完善模特库
This commit is contained in:
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user