feat: 解耦角色与模特并完善模特库
This commit is contained in:
+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