1476 lines
75 KiB
Python
1476 lines
75 KiB
Python
import json
|
|
from pathlib import Path
|
|
|
|
from django.conf import settings
|
|
from django.test import SimpleTestCase
|
|
|
|
from apps.ai.script_agent import normalize_draft
|
|
|
|
|
|
class SkillPromptBundlingTests(SimpleTestCase):
|
|
"""守护「skills 必须随后端打进镜像」这条线 —— 历史上 skills 放仓库根、不在 Docker 构建上下文
|
|
(`./core/backend`)→ 镜像里没有 → 提示词加载为空 → 提取模型收不到「只输出 JSON」铁律 → 吐散文
|
|
→ 「提取结果解析失败」。这两条断言锁死:① skills 落在 BASE_DIR 内(随镜像走);② 提示词非空且带铁律。"""
|
|
|
|
def test_skills_dir_bundled_under_base_dir(self):
|
|
"""skills 必须在 BASE_DIR(=core/backend)内,才会被 Dockerfile 的 `COPY . .` 打进镜像。"""
|
|
for name in ("ecommerce-entity-extract", "ecommerce-video-script"):
|
|
self.assertTrue(
|
|
(Path(settings.BASE_DIR) / "skills" / name / "SKILL.md").exists(),
|
|
f"skills/{name}/SKILL.md 不在 BASE_DIR 内 —— 镜像将丢失该提示词(详见 services._skills_root)",
|
|
)
|
|
|
|
def test_entity_extract_prompt_loads_nonempty_with_json_rule(self):
|
|
"""提取系统提示词必须加载到非空内容,且含「只输出 JSON」铁律(空 → 模型不吐 JSON → 解析失败)。"""
|
|
from apps.ai.services import _load_skill_system_prompt
|
|
|
|
prompt = _load_skill_system_prompt("ecommerce-entity-extract")
|
|
self.assertGreater(len(prompt), 1000, "提取提示词为空/过短 —— skills 没被正确加载")
|
|
self.assertIn("JSON", prompt)
|
|
self.assertIn("entities", prompt)
|
|
|
|
def test_extract_output_contract_is_hardcoded_safety_net(self):
|
|
"""写死的提取输出契约(随代码进镜像)必须存在且含 JSON 铁律 —— 这是 skill 丢失时的兜底,
|
|
保证提取永远收到「只输出 JSON」指令(对齐脚本 agent 的 _OUTPUT_PROTOCOL,根治"系统提示词为空"事故)。"""
|
|
from apps.ai.services import _EXTRACT_OUTPUT_CONTRACT
|
|
|
|
self.assertGreater(len(_EXTRACT_OUTPUT_CONTRACT), 200)
|
|
for marker in ("entities", "segments", "character", "scene", "JSON"):
|
|
self.assertIn(marker, _EXTRACT_OUTPUT_CONTRACT)
|
|
|
|
|
|
class NormalizeDraftTests(SimpleTestCase):
|
|
"""normalize_draft 对模型不按契约输出的容错(防「旁白/画面全空」回归)。"""
|
|
|
|
def test_scriptdraft_shots_variant_fills_narration_and_visual(self):
|
|
"""模型常见变体:{"ScriptDraft":{"basicInfo":..,"shots":[{scene,dialogue,subtitle}]}}。
|
|
必须解开外壳 + 把 shots→segments、scene→画面、dialogue/subtitle→旁白,而非全填空占位镜。"""
|
|
raw = json.dumps({
|
|
"ScriptDraft": {
|
|
"basicInfo": {"totalDuration": 60, "aspectRatio": "9:16", "totalShots": 4},
|
|
"shots": [
|
|
{"shotNo": 1, "duration": 15, "scene": "更衣室扯卡裆旧裤", "dialogue": "卡裆太社死!", "subtitle": "还在卡裆?"},
|
|
{"shotNo": 2, "duration": 15, "scene": "特写拉扯面料回弹", "dialogue": "高弹不变形!", "subtitle": "裸感面料"},
|
|
{"shotNo": 3, "duration": 15, "scene": "健身切通勤", "dialogue": "都能穿!", "subtitle": "一裤多穿"},
|
|
{"shotNo": 4, "duration": 15, "scene": "对镜弹小黄车", "dialogue": "点小黄车抢!", "subtitle": "点击入手"},
|
|
],
|
|
}
|
|
}, ensure_ascii=False)
|
|
draft = normalize_draft(raw, aspect_ratio="9:16", total_duration=60)
|
|
self.assertEqual(len(draft["segments"]), 4)
|
|
self.assertEqual([s["role"] for s in draft["segments"]], ["钩子", "痛点", "卖点", "CTA"])
|
|
for seg in draft["segments"]:
|
|
self.assertTrue(seg["narration"], "旁白不应为空")
|
|
self.assertTrue(seg["visual"], "画面不应为空")
|
|
self.assertEqual(draft["segments"][0]["narration"], "卡裆太社死!")
|
|
self.assertEqual(draft["segments"][0]["visual"], "更衣室扯卡裆旧裤")
|
|
|
|
def test_lines_and_caption_variant_fills_narration(self):
|
|
"""另一种变体(Doubao-Seed-2.0-P):shots 用 lines:[{role,content}] 放口播、caption 放字幕。
|
|
必须把 lines 的 content 当对白/旁白,caption 兜底,而不是只填画面留旁白空。"""
|
|
raw = json.dumps({
|
|
"ScriptDraft": {
|
|
"basicInfo": {"totalDuration": 30, "aspectRatio": "9:16"},
|
|
"shots": [
|
|
{"shotNo": 1, "scene": "化妆台两闺蜜", "lines": [{"role": "女主", "content": "防晒泛白太尴尬!"}, {"role": "闺蜜", "content": "试试这个!"}], "caption": "防晒踩雷?"},
|
|
{"shotNo": 2, "scene": "手背挤膏体", "lines": [{"role": "闺蜜", "content": "质地清透不泛白"}], "caption": "清透质地"},
|
|
],
|
|
}
|
|
}, ensure_ascii=False)
|
|
draft = normalize_draft(raw, aspect_ratio="9:16", total_duration=30)
|
|
self.assertEqual(len(draft["segments"]), 2)
|
|
for seg in draft["segments"]:
|
|
self.assertTrue(seg["narration"], "旁白不应为空")
|
|
self.assertTrue(seg["visual"], "画面不应为空")
|
|
self.assertIn("防晒泛白太尴尬", draft["segments"][0]["narration"])
|
|
self.assertEqual(len(draft["segments"][0]["dialogue"]), 2) # lines→结构化对白
|
|
|
|
def test_generic_resolver_covers_unseen_field_names(self):
|
|
"""模型每次换字段名(scene/screenDescription/画面…、dialogue/lines/caption…)。
|
|
通用解析应「优先键 + 关键词模糊匹配」都能填,且跳过 bgMusic/note 等非内容键。"""
|
|
variants = {
|
|
"canonical": {"total_duration": 15, "segments": [{"role": "钩子", "narration": "口播A", "visual": "画面A"}]},
|
|
"scene+dialogue_str+subtitle": {"total_duration": 15, "shots": [{"scene": "画面B", "dialogue": "口播B", "subtitle": "字幕B"}]},
|
|
"scene+lines+caption": {"total_duration": 15, "shots": [{"scene": "画面C", "lines": [{"role": "女主", "content": "口播C"}], "caption": "字幕C"}]},
|
|
"screenDescription+dialogue+bgMusic+note": {"total_duration": 15, "shots": [{"screenDescription": "画面D", "dialogue": "口播D", "bgMusic": "音乐D", "note": "备注D"}]},
|
|
}
|
|
for name, raw in variants.items():
|
|
draft = normalize_draft(json.dumps(raw, ensure_ascii=False), aspect_ratio="9:16", total_duration=15)
|
|
seg = draft["segments"][0]
|
|
self.assertTrue(seg["narration"], f"{name}: 旁白为空")
|
|
self.assertTrue(seg["visual"], f"{name}: 画面为空")
|
|
# bgMusic/note 不得被误当画面/旁白
|
|
d = normalize_draft(json.dumps(variants["screenDescription+dialogue+bgMusic+note"], ensure_ascii=False), aspect_ratio="9:16", total_duration=15)
|
|
self.assertEqual(d["segments"][0]["visual"], "画面D")
|
|
self.assertEqual(d["segments"][0]["narration"], "口播D")
|
|
|
|
def test_string_dialogue_not_iterated_as_chars(self):
|
|
"""dialogue 为整句字符串时,要当作旁白而不是逐字符遍历。"""
|
|
raw = json.dumps({
|
|
"segments": [{"role": "钩子", "dialogue": "一句完整口播", "visual": "画面"}],
|
|
"total_duration": 15,
|
|
}, ensure_ascii=False)
|
|
draft = normalize_draft(raw, aspect_ratio="9:16", total_duration=15)
|
|
self.assertEqual(draft["segments"][0]["narration"], "一句完整口播")
|
|
self.assertEqual(draft["segments"][0]["dialogue"], []) # 字符串不进结构化对白
|
|
|
|
def test_empty_segments_skeleton_yields_to_richer_scenes(self):
|
|
"""真实回归(全空根因):模型把内容放 scenes/voiceover/visual,却另给一个**空 segments 骨架**。
|
|
必须挑内容最丰富的数组(scenes),而不是见 segments 是 list 就用→落 4 个空镜。"""
|
|
raw = json.dumps({
|
|
"scenes": [
|
|
{"sceneIndex": 1, "sceneTitle": "通勤", "voiceover": "手机一卡效率掉线", "visual": "地铁口拿出亮银色机身"},
|
|
{"sceneIndex": 2, "sceneTitle": "办公", "voiceover": "A19多任务很跟手", "visual": "办公桌俯拍切换应用"},
|
|
],
|
|
"segments": [{"index": 0, "role": "钩子", "narration": "", "visual": ""}, {"index": 1, "role": "痛点", "narration": "", "visual": ""}],
|
|
"total_duration": 30,
|
|
}, ensure_ascii=False)
|
|
draft = normalize_draft(raw, aspect_ratio="9:16", total_duration=30)
|
|
self.assertEqual(draft["segments"][0]["narration"], "手机一卡效率掉线")
|
|
self.assertEqual(draft["segments"][0]["visual"], "地铁口拿出亮银色机身")
|
|
|
|
def test_script_key_with_visual_object_is_flattened(self):
|
|
"""真实回归:数组键叫 script、visual 写成 {setting,camera,key_shots} 对象。
|
|
必须认出 script 数组并把 visual 对象拍平成一句,而非留空。"""
|
|
raw = json.dumps({
|
|
"script": [
|
|
{"scene_id": 1, "scene_title": "通勤", "voiceover": "选手机看重流畅好看",
|
|
"visual": {"setting": "地铁口", "camera": "竖屏手持跟拍", "key_shots": ["拿出亮银色机身"]}},
|
|
],
|
|
"total_duration": 15,
|
|
}, ensure_ascii=False)
|
|
draft = normalize_draft(raw, aspect_ratio="9:16", total_duration=15)
|
|
seg = draft["segments"][0]
|
|
self.assertEqual(seg["narration"], "选手机看重流畅好看")
|
|
self.assertIn("地铁口", seg["visual"])
|
|
self.assertIn("竖屏手持跟拍", seg["visual"])
|
|
|
|
def test_audio_field_fills_narration(self):
|
|
"""真实回归:shots 用 audio 放口播(GPT 变体),旁白曾因 audio 不在词典而留空。"""
|
|
raw = json.dumps({
|
|
"shots": [{"scene": 1, "visual": "咖啡厅办公把玩机身", "audio": "这款亮银色是我的高光决定"}],
|
|
"total_duration": 15,
|
|
}, ensure_ascii=False)
|
|
draft = normalize_draft(raw, aspect_ratio="9:16", total_duration=15)
|
|
self.assertEqual(draft["segments"][0]["narration"], "这款亮银色是我的高光决定")
|
|
self.assertEqual(draft["segments"][0]["visual"], "咖啡厅办公把玩机身")
|
|
|
|
def test_canonical_flat_schema_still_works(self):
|
|
"""契约内的扁平 schema(segments/narration/visual)不受兼容改动影响。"""
|
|
raw = json.dumps({
|
|
"total_duration": 30,
|
|
"segments": [
|
|
{"role": "钩子", "narration": "口播1", "visual": "画面1"},
|
|
{"role": "CTA", "narration": "口播2", "visual": "画面2"},
|
|
],
|
|
}, ensure_ascii=False)
|
|
draft = normalize_draft(raw, aspect_ratio="9:16", total_duration=30)
|
|
self.assertEqual(len(draft["segments"]), 2)
|
|
self.assertEqual(draft["segments"][0]["narration"], "口播1")
|
|
self.assertEqual(draft["segments"][1]["visual"], "画面2")
|
|
|
|
|
|
from io import BytesIO
|
|
from unittest.mock import patch
|
|
|
|
from django.test import TestCase, override_settings
|
|
|
|
from rest_framework.test import APIClient
|
|
|
|
from apps.accounts.models import Team, TeamMember, User
|
|
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, CreditLedger
|
|
from apps.products.models import Product
|
|
|
|
|
|
class FreeVideoAssetDeletionTests(TestCase):
|
|
"""自由创作视频资产删除后,任务流应降级为无可播成品,不能回退到旧 fallback URL。"""
|
|
|
|
def setUp(self):
|
|
self.user = User.objects.create_user(username="fvdel", password="pass")
|
|
self.team = Team.objects.create(name="FreeVideoDel", owner=self.user)
|
|
TeamMember.objects.create(team=self.team, user=self.user, role="owner", status="active")
|
|
provider = ModelProvider.objects.create(name="fv-provider", display_name="FV Provider")
|
|
self.model = ModelConfig.objects.create(
|
|
provider=provider,
|
|
name="fv-model",
|
|
display_name="FV Model",
|
|
capability=ModelConfig.Capability.VIDEO,
|
|
)
|
|
self.task = AITask.objects.create(
|
|
team=self.team,
|
|
created_by=self.user,
|
|
task_type=AITask.Type.FREE_VIDEO,
|
|
status=AITask.Status.SUCCEEDED,
|
|
model_config=self.model,
|
|
idempotency_key="free-video-delete-test",
|
|
request_payload={
|
|
"prompt": "生成一个广告视频",
|
|
"fallback_video_url": "http://fallback.example/video.mp4",
|
|
"fallback_note": "fallback",
|
|
},
|
|
)
|
|
self.asset = Asset.objects.create(
|
|
team=self.team,
|
|
created_by=self.user,
|
|
name="自由创作视频",
|
|
asset_type=Asset.Type.VIDEO,
|
|
source=Asset.Source.AI_GENERATED,
|
|
category=Asset.Category.FREE_CREATE,
|
|
origin_task=self.task,
|
|
)
|
|
AssetFile.objects.create(
|
|
asset=self.asset,
|
|
object_key="free-video.mp4",
|
|
bucket="b",
|
|
content_type="video/mp4",
|
|
preview_url="http://tos.example/video.mp4",
|
|
is_primary=True,
|
|
)
|
|
AssetFile.objects.create(
|
|
asset=self.asset,
|
|
object_key="free-video-poster.jpg",
|
|
bucket="b",
|
|
content_type="image/jpeg",
|
|
preview_url="http://tos.example/free-video-poster.jpg",
|
|
is_primary=False,
|
|
)
|
|
|
|
def test_deleted_persisted_asset_does_not_fall_back_to_raw_video_url(self):
|
|
from apps.ai.free_video import serialize_free_video_task
|
|
|
|
self.assertEqual(serialize_free_video_task(self.task)["video_url"], "http://tos.example/video.mp4")
|
|
self.asset.is_deleted = True
|
|
self.asset.save(update_fields=["is_deleted"])
|
|
task = AITask.objects.prefetch_related("generated_assets", "generated_assets__files").get(id=self.task.id)
|
|
data = serialize_free_video_task(task)
|
|
self.assertEqual(data["video_url"], "")
|
|
self.assertEqual(data["thumbnail_url"], "")
|
|
self.assertEqual(data["fallback_note"], "fallback")
|
|
trash_data = serialize_free_video_task(task, include_deleted_assets=True)
|
|
self.assertEqual(trash_data["video_url"], "http://tos.example/video.mp4")
|
|
self.assertEqual(trash_data["thumbnail_url"], "http://tos.example/free-video-poster.jpg")
|
|
|
|
def test_asset_library_uses_full_free_video_title_and_searches_it(self):
|
|
"""资产库应从关联任务恢复历史截断标题,并支持搜索完整标题的后半段。"""
|
|
full_prompt = "@O1CN01anEq9u245Jn8Ccedn_!!4611686018427385691-0-s-完整后缀"
|
|
self.task.request_payload["prompt"] = full_prompt
|
|
self.task.save(update_fields=["request_payload"])
|
|
self.asset.name = full_prompt[:50] # 模拟旧版本写库的截断数据
|
|
self.asset.save(update_fields=["name"])
|
|
|
|
client = APIClient()
|
|
client.force_authenticate(self.user)
|
|
listing = client.get("/api/assets/?tab=video_creations&page_size=10").json()
|
|
|
|
self.assertEqual(listing["count"], 1)
|
|
self.assertEqual(listing["results"][0]["display_name"], full_prompt)
|
|
|
|
search = client.get(f"/api/assets/?tab=video_creations&q={full_prompt[-4:]}&page_size=10").json()
|
|
self.assertEqual([row["id"] for row in search["results"]], [str(self.asset.id)])
|
|
|
|
def test_direct_asset_delete_stays_in_asset_trash(self):
|
|
client = APIClient()
|
|
client.force_authenticate(self.user)
|
|
|
|
delete = client.delete(f"/api/assets/{self.asset.id}/")
|
|
self.assertEqual(delete.status_code, 204)
|
|
self.task.refresh_from_db()
|
|
self.asset.refresh_from_db()
|
|
self.assertFalse(self.task.is_deleted)
|
|
self.assertTrue(self.asset.is_deleted)
|
|
|
|
trash_ids = {a["id"] for a in client.get("/api/assets/trash/").json()["results"]}
|
|
self.assertIn(str(self.asset.id), trash_ids)
|
|
self.assertEqual(client.get("/api/ai/free-video/").json()["total"], 0)
|
|
|
|
restore = client.post(f"/api/assets/{self.asset.id}/restore/")
|
|
self.assertEqual(restore.status_code, 200)
|
|
live_ids = {t["id"] for t in client.get("/api/ai/free-video/").json()["results"]}
|
|
self.assertIn(str(self.task.id), live_ids)
|
|
|
|
def test_trash_restore_and_purge_free_video_keeps_record(self):
|
|
self.client = APIClient()
|
|
self.client.force_authenticate(self.user)
|
|
|
|
delete = self.client.delete(f"/api/ai/free-video/{self.task.id}/")
|
|
self.assertEqual(delete.status_code, 204)
|
|
self.task.refresh_from_db()
|
|
self.asset.refresh_from_db()
|
|
self.assertTrue(self.task.is_deleted)
|
|
self.assertTrue(self.asset.is_deleted)
|
|
self.assertIsNone(self.task.purged_at)
|
|
self.assertEqual(self.client.get("/api/ai/free-video/").json()["total"], 0)
|
|
trash_items = self.client.get("/api/ai/free-video/trash/").json()["results"]
|
|
trash_ids = {t["id"] for t in trash_items}
|
|
self.assertIn(str(self.task.id), trash_ids)
|
|
asset_trash_ids = {a["id"] for a in self.client.get("/api/assets/trash/").json()["results"]}
|
|
self.assertNotIn(str(self.asset.id), asset_trash_ids)
|
|
trash_item = next(t for t in trash_items if t["id"] == str(self.task.id))
|
|
self.assertEqual(trash_item["thumbnail_url"], "http://tos.example/free-video-poster.jpg")
|
|
|
|
restore = self.client.post(f"/api/ai/free-video/{self.task.id}/restore/")
|
|
self.assertEqual(restore.status_code, 200)
|
|
self.task.refresh_from_db()
|
|
self.asset.refresh_from_db()
|
|
self.assertFalse(self.task.is_deleted)
|
|
self.assertFalse(self.asset.is_deleted)
|
|
self.assertIsNone(self.task.purged_at)
|
|
live_ids = {t["id"] for t in self.client.get("/api/ai/free-video/").json()["results"]}
|
|
self.assertIn(str(self.task.id), live_ids)
|
|
|
|
self.client.delete(f"/api/ai/free-video/{self.task.id}/")
|
|
purge = self.client.delete(f"/api/ai/free-video/{self.task.id}/purge/")
|
|
self.assertEqual(purge.status_code, 204)
|
|
self.task.refresh_from_db()
|
|
self.asset.refresh_from_db()
|
|
self.assertTrue(self.task.is_deleted)
|
|
self.assertTrue(self.asset.is_deleted)
|
|
self.assertIsNotNone(self.task.purged_at)
|
|
self.assertIsNotNone(self.asset.purged_at)
|
|
self.assertTrue(AITask.objects.filter(id=self.task.id).exists())
|
|
self.assertTrue(Asset.objects.filter(id=self.asset.id).exists())
|
|
self.assertEqual(self.asset.files.count(), 2)
|
|
trash_ids = {t["id"] for t in self.client.get("/api/ai/free-video/trash/").json()["results"]}
|
|
self.assertNotIn(str(self.task.id), trash_ids)
|
|
|
|
|
|
class StandaloneImageReferenceTests(TestCase):
|
|
"""独立生图(平台套图 / 模特上身图)必须把商品真实主图(+ 模特图)作为参考图走 image_edit,
|
|
而不是纯文生图——回归保护 #图片生成没参考商品主图# 这个 bug。
|
|
(图像默认模型由迁移 seed 的 tokenssr:gpt-image-2 提供,get_image_provider 全程 mock。)"""
|
|
|
|
def setUp(self):
|
|
self.user = User.objects.create_user(username="owner", password="pass")
|
|
self.team = Team.objects.create(name="T", owner=self.user)
|
|
CreditAccount.objects.create(team=self.team, balance="100.0000")
|
|
# 商品 + 主图(带可访问 preview_url)
|
|
self.product = Product.objects.create(team=self.team, created_by=self.user, title="南卡 Lite Pro")
|
|
cover = Asset.objects.create(
|
|
team=self.team, created_by=self.user, name="主图", asset_type=Asset.Type.IMAGE,
|
|
source=Asset.Source.UPLOAD, category=Asset.Category.PRODUCT_IMAGE,
|
|
)
|
|
AssetFile.objects.create(asset=cover, object_key="c.png", bucket="b", content_type="image/png", preview_url="http://x/cover.png", is_primary=True)
|
|
self.product.cover_asset = cover
|
|
self.product.save(update_fields=["cover_asset"])
|
|
|
|
def _patch_provider(self):
|
|
# image_edit 返回值 + 媒体落库链路全部 mock,聚焦验证「传了哪些参考图」
|
|
provider = patch("apps.ai.services.get_image_provider").start()
|
|
prov = provider.return_value
|
|
prov.image_edit.return_value = {"data": [{"url": "http://x/out.png"}]}
|
|
prov.image_generation.return_value = {"data": [{"url": "http://x/out.png"}]}
|
|
prov.extract_first_media_url.return_value = "http://x/out.png"
|
|
media = patch("apps.ai.services.VolcanoArkProvider.media_to_bytes").start()
|
|
media.return_value = (BytesIO(b"img"), "image/png")
|
|
store = patch("apps.ai.services.TosStorage").start()
|
|
stored = store.return_value.upload_fileobj.return_value
|
|
stored.object_key, stored.bucket, stored.content_type, stored.size_bytes = "o.png", "b", "image/png", 3
|
|
self.addCleanup(patch.stopall)
|
|
return prov
|
|
|
|
def test_cover_mode_references_product_main_image(self):
|
|
prov = self._patch_provider()
|
|
enqueue_standalone_images(team=self.team, user=self.user, prompt="平台套图", mode="cover", count=1, product_id=str(self.product.id), ratio="4:5")
|
|
prov.image_edit.assert_called_once()
|
|
self.assertEqual(prov.image_edit.call_args.kwargs["images"], ["http://x/cover.png"])
|
|
prov.image_generation.assert_not_called()
|
|
|
|
def test_model_tryon_combines_product_and_model_images(self):
|
|
prov = self._patch_provider()
|
|
# 模特资产(PERSON)带 preview_url
|
|
model_asset = Asset.objects.create(
|
|
team=self.team, created_by=self.user, name="模特", asset_type=Asset.Type.IMAGE,
|
|
source=Asset.Source.AI_GENERATED, category=Asset.Category.PERSON,
|
|
)
|
|
AssetFile.objects.create(asset=model_asset, object_key="m.png", bucket="b", content_type="image/png", preview_url="http://x/model.png", is_primary=True)
|
|
enqueue_standalone_images(team=self.team, user=self.user, prompt="上身图", mode="model", count=1, product_id=str(self.product.id), model_id=str(model_asset.id), ratio="4:5")
|
|
prov.image_edit.assert_called_once()
|
|
self.assertEqual(prov.image_edit.call_args.kwargs["images"], ["http://x/cover.png", "http://x/model.png"])
|
|
prov.image_generation.assert_not_called()
|
|
|
|
@override_settings(MODEL_TRYON_PROMPT_V2_ENABLED=True)
|
|
def test_model_tryon_v2_uses_structured_prompt_and_persists_trace(self):
|
|
prov = self._patch_provider()
|
|
self.product.title = "中蓝色高腰宽腿牛仔裤"
|
|
self.product.category = "服饰内衣"
|
|
self.product.save(update_fields=["title", "category"])
|
|
self.product.selling_points.create(title="显高显瘦", sort_order=0)
|
|
model_asset = Asset.objects.create(
|
|
team=self.team, created_by=self.user, name="模特", asset_type=Asset.Type.IMAGE,
|
|
source=Asset.Source.AI_GENERATED, category=Asset.Category.PERSON,
|
|
)
|
|
AssetFile.objects.create(
|
|
asset=model_asset, object_key="m-v2.png", bucket="b", content_type="image/png",
|
|
preview_url="http://x/model-v2.png", is_primary=True,
|
|
)
|
|
user_prompt = "参考图1的女生穿上图2的牛仔裤,全身站姿,双腿完整露出,浅灰纯色背景"
|
|
|
|
submitted = enqueue_standalone_images(
|
|
team=self.team,
|
|
user=self.user,
|
|
prompt=user_prompt,
|
|
mode="model",
|
|
count=2,
|
|
product_id=str(self.product.id),
|
|
model_id=str(model_asset.id),
|
|
ratio="7:10",
|
|
)
|
|
|
|
tasks = list(AITask.objects.filter(id__in=[task.id for task in submitted]).order_by("request_payload__index"))
|
|
self.assertEqual(len(prov.image_edit.call_args_list), 2)
|
|
used_prompts = [call.kwargs["prompt"] for call in prov.image_edit.call_args_list]
|
|
used_sizes = [call.kwargs["size"] for call in prov.image_edit.call_args_list]
|
|
self.assertTrue(all(user_prompt in prompt for prompt in used_prompts))
|
|
self.assertTrue(all("参考图角色(最高优先级)" in prompt for prompt in used_prompts))
|
|
self.assertTrue(all("严格使用 7:10 比例" in prompt for prompt in used_prompts))
|
|
self.assertTrue(all("半身近景" not in prompt and "手持或局部" not in prompt for prompt in used_prompts))
|
|
self.assertNotEqual(used_prompts[0], used_prompts[1])
|
|
self.assertEqual(used_sizes, ["1024x1536", "1024x1536"])
|
|
|
|
for index, task in enumerate(tasks):
|
|
trace = task.request_payload["tryon_prompt"]
|
|
self.assertTrue(trace["applied"])
|
|
self.assertEqual(trace["version"], "v2.2")
|
|
self.assertEqual(trace["rollout_source"], "global")
|
|
self.assertEqual(trace["effective_prompt"], used_prompts[index])
|
|
self.assertEqual(trace["shot_index"], index)
|
|
self.assertEqual(trace["batch_count"], 2)
|
|
self.assertEqual(trace["requested_ratio"], "7:10")
|
|
self.assertEqual(trace["resolved_ratio"], "7:10")
|
|
self.assertEqual(trace["reference_roles"]["product_numbers"], [1])
|
|
self.assertEqual(trace["reference_roles"]["model_portrait_number"], 2)
|
|
self.assertIsNone(trace["reference_roles"]["model_triview_number"])
|
|
self.assertEqual(task.request_payload["tryon_classification"]["kind"], "lower")
|
|
|
|
@override_settings(
|
|
MODEL_TRYON_PROMPT_V2_ENABLED=False,
|
|
MODEL_TRYON_PROMPT_V2_CANARY_TEAM_IDS=frozenset(),
|
|
)
|
|
def test_model_tryon_flag_off_keeps_legacy_prompt(self):
|
|
prov = self._patch_provider()
|
|
self.product.title = "女装裤子"
|
|
self.product.save(update_fields=["title"])
|
|
model_asset = Asset.objects.create(
|
|
team=self.team, created_by=self.user, name="模特", asset_type=Asset.Type.IMAGE,
|
|
source=Asset.Source.AI_GENERATED, category=Asset.Category.PERSON,
|
|
)
|
|
AssetFile.objects.create(
|
|
asset=model_asset, object_key="m-legacy.png", bucket="b", content_type="image/png",
|
|
preview_url="http://x/model-legacy.png", is_primary=True,
|
|
)
|
|
submitted = enqueue_standalone_images(
|
|
team=self.team, user=self.user, prompt="模特上身图", mode="model", count=1,
|
|
product_id=str(self.product.id), model_id=str(model_asset.id), ratio="3:4",
|
|
)
|
|
used_prompt = prov.image_edit.call_args.kwargs["prompt"]
|
|
self.assertIn("半身近景", used_prompt)
|
|
task = AITask.objects.get(id=submitted[0].id)
|
|
self.assertNotIn("tryon_prompt", task.request_payload)
|
|
|
|
@override_settings(MODEL_TRYON_PROMPT_V2_ENABLED=False)
|
|
def test_model_tryon_canary_team_uses_v22_and_records_source(self):
|
|
prov = self._patch_provider()
|
|
self.product.title = "女装裤子"
|
|
self.product.save(update_fields=["title"])
|
|
with self.settings(MODEL_TRYON_PROMPT_V2_CANARY_TEAM_IDS={str(self.team.id).upper()}):
|
|
submitted = enqueue_standalone_images(
|
|
team=self.team,
|
|
user=self.user,
|
|
prompt="全身站姿,双腿完整露出",
|
|
mode="model",
|
|
count=1,
|
|
product_id=str(self.product.id),
|
|
ratio="3:4",
|
|
)
|
|
|
|
used_prompt = prov.image_edit.call_args.kwargs["prompt"]
|
|
self.assertIn("商品关键结构优先", used_prompt)
|
|
task = AITask.objects.get(id=submitted[0].id)
|
|
trace = task.request_payload["tryon_prompt"]
|
|
self.assertTrue(trace["applied"])
|
|
self.assertEqual(trace["version"], "v2.2")
|
|
self.assertEqual(trace["rollout_source"], "canary")
|
|
|
|
@override_settings(
|
|
MODEL_TRYON_PROMPT_V2_ENABLED=False,
|
|
MODEL_TRYON_PROMPT_V2_CANARY_TEAM_IDS={"not-a-team-uuid"},
|
|
)
|
|
def test_model_tryon_non_canary_team_fails_closed_to_legacy_prompt(self):
|
|
prov = self._patch_provider()
|
|
self.product.title = "女装裤子"
|
|
self.product.save(update_fields=["title"])
|
|
submitted = enqueue_standalone_images(
|
|
team=self.team,
|
|
user=self.user,
|
|
prompt="模特上身图",
|
|
mode="model",
|
|
count=1,
|
|
product_id=str(self.product.id),
|
|
ratio="3:4",
|
|
)
|
|
|
|
used_prompt = prov.image_edit.call_args.kwargs["prompt"]
|
|
self.assertIn("半身近景", used_prompt)
|
|
task = AITask.objects.get(id=submitted[0].id)
|
|
self.assertNotIn("tryon_prompt", task.request_payload)
|
|
|
|
@override_settings(MODEL_TRYON_PROMPT_V2_ENABLED=False)
|
|
def test_model_tryon_canary_without_product_reference_records_safe_fallback(self):
|
|
prov = self._patch_provider()
|
|
self.product.cover_asset = None
|
|
self.product.save(update_fields=["cover_asset"])
|
|
with self.settings(MODEL_TRYON_PROMPT_V2_CANARY_TEAM_IDS={str(self.team.id)}):
|
|
submitted = enqueue_standalone_images(
|
|
team=self.team,
|
|
user=self.user,
|
|
prompt="模特上身图",
|
|
mode="model",
|
|
count=1,
|
|
product_id=str(self.product.id),
|
|
ratio="3:4",
|
|
)
|
|
|
|
prov.image_generation.assert_called_once()
|
|
task = AITask.objects.get(id=submitted[0].id)
|
|
trace = task.request_payload["tryon_prompt"]
|
|
self.assertFalse(trace["applied"])
|
|
self.assertEqual(trace["rollout_source"], "canary")
|
|
self.assertEqual(trace["fallback_reason"], "no_product_reference")
|
|
|
|
@override_settings(MODEL_TRYON_PROMPT_V2_ENABLED=True)
|
|
def test_model_tryon_v22_uses_category_default_ratio_when_user_omits_ratio(self):
|
|
prov = self._patch_provider()
|
|
self.product.title = "女装裤子"
|
|
self.product.save(update_fields=["title"])
|
|
submitted = enqueue_standalone_images(
|
|
team=self.team,
|
|
user=self.user,
|
|
prompt="模特上身图",
|
|
mode="model",
|
|
count=1,
|
|
product_id=str(self.product.id),
|
|
)
|
|
|
|
self.assertEqual(prov.image_edit.call_args.kwargs["size"], "1024x1536")
|
|
task = AITask.objects.get(id=submitted[0].id)
|
|
trace = task.request_payload["tryon_prompt"]
|
|
self.assertEqual(trace["requested_ratio"], None)
|
|
self.assertEqual(trace["resolved_ratio"], "3:4")
|
|
self.assertEqual(trace["ratio_source"], "default")
|
|
self.assertEqual(trace["rollout_source"], "global")
|
|
|
|
@override_settings(MODEL_TRYON_PROMPT_V2_ENABLED=True)
|
|
def test_model_tryon_old_queued_task_without_v2_metadata_falls_back_safely(self):
|
|
from apps.ai.services import run_standalone_image_task
|
|
|
|
prov = self._patch_provider()
|
|
self.product.title = "女装裤子"
|
|
self.product.save(update_fields=["title"])
|
|
with patch("apps.ai.tasks.generate_standalone_image_task.delay"):
|
|
task = enqueue_standalone_images(
|
|
team=self.team,
|
|
user=self.user,
|
|
prompt="模特上身图",
|
|
mode="model",
|
|
count=1,
|
|
product_id=str(self.product.id),
|
|
ratio="3:4",
|
|
)[0]
|
|
|
|
legacy_payload = dict(task.request_payload)
|
|
legacy_payload.pop("tryon_batch_count", None)
|
|
legacy_payload.pop("tryon_classification", None)
|
|
task.request_payload = legacy_payload
|
|
task.save(update_fields=["request_payload", "updated_at"])
|
|
|
|
run_standalone_image_task(task_id=str(task.id))
|
|
|
|
task.refresh_from_db()
|
|
used_prompt = prov.image_edit.call_args.kwargs["prompt"]
|
|
self.assertIn("半身近景", used_prompt)
|
|
self.assertEqual(task.status, AITask.Status.SUCCEEDED)
|
|
trace = task.request_payload["tryon_prompt"]
|
|
self.assertFalse(trace["applied"])
|
|
self.assertEqual(trace["rollout_source"], "global")
|
|
self.assertEqual(trace["fallback_reason"], "unsupported_tryon_batch_count")
|
|
|
|
def test_cover_mode_falls_back_to_t2i_without_main_image(self):
|
|
prov = self._patch_provider()
|
|
self.product.cover_asset = None
|
|
self.product.save(update_fields=["cover_asset"])
|
|
enqueue_standalone_images(team=self.team, user=self.user, prompt="平台套图", mode="cover", count=1, product_id=str(self.product.id), ratio="4:5")
|
|
prov.image_generation.assert_called_once()
|
|
prov.image_edit.assert_not_called()
|
|
|
|
def test_image_mode_uses_uploaded_reference_images(self):
|
|
"""图片创作自由模式:用户上传的参考图必须作为 image_edit 的参考图传进去,
|
|
而不是被忽略走纯文生图——回归保护 #图片创作没参考我上传的素材# 这个最致命的 bug。"""
|
|
prov = self._patch_provider()
|
|
ref = Asset.objects.create(
|
|
team=self.team, created_by=self.user, name="背心参考", asset_type=Asset.Type.IMAGE,
|
|
source=Asset.Source.UPLOAD, category=Asset.Category.UPLOAD,
|
|
)
|
|
AssetFile.objects.create(asset=ref, object_key="r.png", bucket="b", content_type="image/png", preview_url="http://x/ref.png", is_primary=True)
|
|
enqueue_standalone_images(
|
|
team=self.team, user=self.user, prompt="按这件背心生成不同场景的穿搭",
|
|
mode="image", count=1, ratio="1:1", reference_image_ids=[str(ref.id)],
|
|
)
|
|
prov.image_edit.assert_called_once()
|
|
self.assertEqual(prov.image_edit.call_args.kwargs["images"], ["http://x/ref.png"])
|
|
# 提示词必须显式钉住「参考图=主体唯一依据」+ 带上用户原话(否则图生图不真的还原上传素材);
|
|
# 但不能再无差别「严格保留款式」——用户要求(如换现代服装)必须是最高优先级,否则同批一半不换装
|
|
used_prompt = prov.image_edit.call_args.kwargs["prompt"]
|
|
self.assertIn("按这件背心生成不同场景的穿搭", used_prompt)
|
|
self.assertIn("唯一依据", used_prompt)
|
|
self.assertIn("用户要求是最高优先级", used_prompt)
|
|
prov.image_generation.assert_not_called()
|
|
|
|
def test_refs_force_image_edit_model_when_volcano_selected(self):
|
|
"""用户选了火山 Seedream(只有图生图、不会真锁主体)却传了参考图时,必须自动切到支持
|
|
image_edit 的 gpt-image-2——否则参考图形同虚设。回归保护「选火山+传参考图=完全不参考」。"""
|
|
from apps.ai.models import ModelProvider
|
|
|
|
vp, _ = ModelProvider.objects.get_or_create(name="volcengine", defaults={"display_name": "火山"})
|
|
ModelConfig.objects.create(provider=vp, name="seedream-4", display_name="Seedream", capability=ModelConfig.Capability.IMAGE)
|
|
self._patch_provider()
|
|
ref = Asset.objects.create(
|
|
team=self.team, created_by=self.user, name="背心", asset_type=Asset.Type.IMAGE,
|
|
source=Asset.Source.UPLOAD, category=Asset.Category.UPLOAD,
|
|
)
|
|
AssetFile.objects.create(asset=ref, object_key="r.png", bucket="b", content_type="image/png", preview_url="http://x/ref.png", is_primary=True)
|
|
tasks = enqueue_standalone_images(
|
|
team=self.team, user=self.user, prompt="不同场景", mode="image", count=1,
|
|
image_model="volcano", reference_image_ids=[str(ref.id)],
|
|
)
|
|
# 任务最终落在支持 image_edit 的模型(gpt-image),而不是用户选的火山 Seedream
|
|
self.assertIn("gpt-image", tasks[0].model_config.name)
|
|
|
|
def test_model_tryon_keeps_each_user_selected_model_even_with_extra_refs(self):
|
|
"""模特上身图由专用 Worker 向两类模型传商品/人物参考图,不能沿用自由创作的自动换模逻辑。"""
|
|
vp, _ = ModelProvider.objects.get_or_create(
|
|
name="volcengine",
|
|
defaults={"display_name": "火山"},
|
|
)
|
|
seedream = ModelConfig.objects.create(
|
|
provider=vp,
|
|
name="seedream-v22-user-choice",
|
|
display_name="Seedream V2.2 user choice",
|
|
capability=ModelConfig.Capability.IMAGE,
|
|
)
|
|
gpt = ModelConfig.objects.filter(
|
|
capability=ModelConfig.Capability.IMAGE,
|
|
name__icontains="gpt-image",
|
|
).first()
|
|
self.assertIsNotNone(gpt)
|
|
extra_ref = Asset.objects.create(
|
|
team=self.team,
|
|
created_by=self.user,
|
|
name="额外参考",
|
|
asset_type=Asset.Type.IMAGE,
|
|
source=Asset.Source.UPLOAD,
|
|
category=Asset.Category.UPLOAD,
|
|
)
|
|
AssetFile.objects.create(
|
|
asset=extra_ref,
|
|
object_key="extra.png",
|
|
bucket="b",
|
|
content_type="image/png",
|
|
preview_url="http://x/extra.png",
|
|
is_primary=True,
|
|
)
|
|
|
|
seedream_task = enqueue_standalone_images(
|
|
team=self.team,
|
|
user=self.user,
|
|
prompt="模特上身图",
|
|
mode="model",
|
|
count=1,
|
|
product_id=str(self.product.id),
|
|
image_model=f"{vp.name}:{seedream.name}",
|
|
reference_image_ids=[str(extra_ref.id)],
|
|
dispatch=False,
|
|
)[0]
|
|
gpt_task = enqueue_standalone_images(
|
|
team=self.team,
|
|
user=self.user,
|
|
prompt="模特上身图",
|
|
mode="model",
|
|
count=1,
|
|
product_id=str(self.product.id),
|
|
image_model=f"{gpt.provider.name}:{gpt.name}",
|
|
reference_image_ids=[str(extra_ref.id)],
|
|
dispatch=False,
|
|
)[0]
|
|
|
|
self.assertEqual(seedream_task.model_config_id, seedream.id)
|
|
self.assertEqual(gpt_task.model_config_id, gpt.id)
|
|
|
|
|
|
class ImageConversationTests(TestCase):
|
|
"""图片创作「对话」实体:CRUD + 团队隔离 + 软删 + 生图自动归属对话 + 任务回填。"""
|
|
|
|
def setUp(self):
|
|
self.user = User.objects.create_user(username="convowner", password="pass")
|
|
self.team = Team.objects.create(name="ConvT", owner=self.user)
|
|
TeamMember.objects.create(team=self.team, user=self.user, role="owner", status="active")
|
|
self.client = APIClient()
|
|
self.client.force_authenticate(self.user)
|
|
|
|
def test_crud_and_listing(self):
|
|
# 新建
|
|
r = self.client.post("/api/ai/image-conversations/", {"mode": "image", "title": "测试创作"}, format="json")
|
|
self.assertEqual(r.status_code, 201, r.content)
|
|
conv_id = r.json()["id"]
|
|
# 列出(只 image 模式)
|
|
r = self.client.get("/api/ai/image-conversations/?mode=image")
|
|
self.assertEqual(r.status_code, 200)
|
|
self.assertEqual(len(r.json()["results"]), 1)
|
|
# 重命名
|
|
r = self.client.patch(f"/api/ai/image-conversations/{conv_id}/", {"title": "改后"}, format="json")
|
|
self.assertEqual(r.status_code, 200)
|
|
self.assertEqual(r.json()["title"], "改后")
|
|
# 软删:列表消失,但 DB 记录仍在(is_deleted=True)
|
|
r = self.client.delete(f"/api/ai/image-conversations/{conv_id}/")
|
|
self.assertIn(r.status_code, (204, 200))
|
|
self.assertEqual(self.client.get("/api/ai/image-conversations/?mode=image").json()["results"], [])
|
|
self.assertTrue(ImageConversation.objects.get(id=conv_id).is_deleted)
|
|
|
|
def test_trash_restore_and_purge_conversation_keeps_record(self):
|
|
conv = ImageConversation.objects.create(team=self.team, created_by=self.user, title="待删", mode=ImageConversation.Mode.IMAGE)
|
|
provider = ModelProvider.objects.create(name="conv-trash-provider", display_name="Conv Trash Provider")
|
|
model = ModelConfig.objects.create(
|
|
provider=provider,
|
|
name="conv-trash-model",
|
|
display_name="Conv Trash Model",
|
|
capability=ModelConfig.Capability.IMAGE,
|
|
)
|
|
task = AITask.objects.create(
|
|
team=self.team,
|
|
created_by=self.user,
|
|
conversation=conv,
|
|
task_type=AITask.Type.PRODUCT_IMAGE,
|
|
status=AITask.Status.SUCCEEDED,
|
|
model_config=model,
|
|
idempotency_key="conv-trash-task",
|
|
)
|
|
asset = Asset.objects.create(
|
|
team=self.team,
|
|
created_by=self.user,
|
|
name="会话生成图",
|
|
asset_type=Asset.Type.IMAGE,
|
|
source=Asset.Source.AI_GENERATED,
|
|
category=Asset.Category.FREE_CREATE,
|
|
origin_task=task,
|
|
)
|
|
asset_file = AssetFile.objects.create(
|
|
asset=asset,
|
|
object_key="free-create/image.png",
|
|
bucket="test",
|
|
content_type="image/png",
|
|
is_primary=True,
|
|
)
|
|
|
|
delete = self.client.delete(f"/api/ai/image-conversations/{conv.id}/")
|
|
self.assertIn(delete.status_code, (204, 200))
|
|
conv.refresh_from_db()
|
|
task.refresh_from_db()
|
|
asset.refresh_from_db()
|
|
self.assertTrue(conv.is_deleted)
|
|
self.assertTrue(task.is_deleted)
|
|
self.assertTrue(asset.is_deleted)
|
|
self.assertIsNone(conv.purged_at)
|
|
trash_ids = {c["id"] for c in self.client.get("/api/ai/image-conversations/trash/?mode=image").json()["results"]}
|
|
self.assertIn(str(conv.id), trash_ids)
|
|
asset_trash_ids = {a["id"] for a in self.client.get("/api/assets/trash/").json()["results"]}
|
|
self.assertNotIn(str(asset.id), asset_trash_ids)
|
|
|
|
restore = self.client.post(f"/api/ai/image-conversations/{conv.id}/restore/")
|
|
self.assertEqual(restore.status_code, 200)
|
|
conv.refresh_from_db()
|
|
task.refresh_from_db()
|
|
asset.refresh_from_db()
|
|
self.assertFalse(conv.is_deleted)
|
|
self.assertFalse(task.is_deleted)
|
|
self.assertFalse(asset.is_deleted)
|
|
self.assertIsNone(conv.purged_at)
|
|
live_ids = {c["id"] for c in self.client.get("/api/ai/image-conversations/?mode=image").json()["results"]}
|
|
self.assertIn(str(conv.id), live_ids)
|
|
|
|
self.client.delete(f"/api/ai/image-conversations/{conv.id}/")
|
|
purge = self.client.delete(f"/api/ai/image-conversations/{conv.id}/purge/")
|
|
self.assertEqual(purge.status_code, 204)
|
|
conv.refresh_from_db()
|
|
task.refresh_from_db()
|
|
asset.refresh_from_db()
|
|
self.assertTrue(conv.is_deleted)
|
|
self.assertTrue(task.is_deleted)
|
|
self.assertTrue(asset.is_deleted)
|
|
self.assertIsNotNone(conv.purged_at)
|
|
self.assertIsNotNone(task.purged_at)
|
|
self.assertIsNotNone(asset.purged_at)
|
|
self.assertTrue(ImageConversation.objects.filter(id=conv.id).exists())
|
|
self.assertTrue(AITask.objects.filter(id=task.id).exists())
|
|
self.assertTrue(Asset.objects.filter(id=asset.id).exists())
|
|
self.assertTrue(AssetFile.objects.filter(id=asset_file.id).exists())
|
|
trash_ids = {c["id"] for c in self.client.get("/api/ai/image-conversations/trash/?mode=image").json()["results"]}
|
|
self.assertNotIn(str(conv.id), trash_ids)
|
|
|
|
def test_direct_asset_delete_hides_image_but_keeps_conversation(self):
|
|
conv = ImageConversation.objects.create(
|
|
team=self.team,
|
|
created_by=self.user,
|
|
title="资产独立删除",
|
|
mode=ImageConversation.Mode.IMAGE,
|
|
)
|
|
provider = ModelProvider.objects.create(name="asset-delete-provider", display_name="Asset Delete Provider")
|
|
model = ModelConfig.objects.create(
|
|
provider=provider,
|
|
name="asset-delete-model",
|
|
display_name="Asset Delete Model",
|
|
capability=ModelConfig.Capability.IMAGE,
|
|
)
|
|
task = AITask.objects.create(
|
|
team=self.team,
|
|
created_by=self.user,
|
|
conversation=conv,
|
|
task_type=AITask.Type.PRODUCT_IMAGE,
|
|
status=AITask.Status.SUCCEEDED,
|
|
model_config=model,
|
|
idempotency_key="direct-asset-delete-image",
|
|
)
|
|
asset = Asset.objects.create(
|
|
team=self.team,
|
|
created_by=self.user,
|
|
name="自由创作图片",
|
|
asset_type=Asset.Type.IMAGE,
|
|
source=Asset.Source.AI_GENERATED,
|
|
category=Asset.Category.FREE_CREATE,
|
|
origin_task=task,
|
|
)
|
|
|
|
delete = self.client.delete(f"/api/assets/{asset.id}/")
|
|
self.assertEqual(delete.status_code, 204)
|
|
conv.refresh_from_db()
|
|
task.refresh_from_db()
|
|
self.assertFalse(conv.is_deleted)
|
|
self.assertFalse(task.is_deleted)
|
|
|
|
conversation_ids = {
|
|
c["id"] for c in self.client.get("/api/ai/image-conversations/?mode=image").json()["results"]
|
|
}
|
|
self.assertIn(str(conv.id), conversation_ids)
|
|
deleted_tasks = self.client.get(f"/api/ai/image-conversations/{conv.id}/tasks/").json()["tasks"]
|
|
self.assertEqual(deleted_tasks[0]["assets"], [])
|
|
|
|
restore = self.client.post(f"/api/assets/{asset.id}/restore/")
|
|
self.assertEqual(restore.status_code, 200)
|
|
restored_tasks = self.client.get(f"/api/ai/image-conversations/{conv.id}/tasks/").json()["tasks"]
|
|
self.assertEqual([a["id"] for a in restored_tasks[0]["assets"]], [str(asset.id)])
|
|
|
|
def test_team_isolation(self):
|
|
other = User.objects.create_user(username="other", password="pass")
|
|
other_team = Team.objects.create(name="Other", owner=other)
|
|
ImageConversation.objects.create(team=other_team, created_by=other, title="别人的")
|
|
r = self.client.get("/api/ai/image-conversations/?mode=image")
|
|
self.assertEqual(r.json()["results"], [])
|
|
|
|
def test_generate_auto_creates_and_binds_conversation(self):
|
|
# 余额 + 默认图像模型(迁移已 seed),provider/worker 全 mock,聚焦验证「任务挂到对话」
|
|
CreditAccount.objects.create(team=self.team, balance="100.0000")
|
|
patch("apps.ai.tasks.generate_standalone_image_task.delay").start()
|
|
self.addCleanup(patch.stopall)
|
|
conv = ImageConversation.objects.create(team=self.team, created_by=self.user, title="目标对话")
|
|
tasks = enqueue_standalone_images(team=self.team, user=self.user, prompt="一只猫", mode="image", count=2, conversation=conv)
|
|
self.assertEqual(len(tasks), 2)
|
|
self.assertTrue(all(t.conversation_id == conv.id for t in tasks))
|
|
self.assertEqual(conv.tasks.count(), 2)
|
|
|
|
def test_tasks_endpoint_returns_grouped_history(self):
|
|
conv = ImageConversation.objects.create(team=self.team, created_by=self.user, title="历史")
|
|
mc = ModelConfig.objects.filter(capability=ModelConfig.Capability.IMAGE).first()
|
|
AITask.objects.create(
|
|
team=self.team, created_by=self.user, conversation=conv, task_type=AITask.Type.PRODUCT_IMAGE,
|
|
status=AITask.Status.SUCCEEDED, model_config=mc, idempotency_key="conv-test-1",
|
|
request_payload={"prompt": "猫", "batch_id": "bx", "ratio": "1:1"},
|
|
)
|
|
r = self.client.get(f"/api/ai/image-conversations/{conv.id}/tasks/")
|
|
self.assertEqual(r.status_code, 200, r.content)
|
|
data = r.json()["tasks"]
|
|
self.assertEqual(len(data), 1)
|
|
self.assertEqual(data[0]["prompt"], "猫")
|
|
self.assertEqual(data[0]["batch_id"], "bx")
|
|
|
|
def test_generate_with_refs_persists_and_tasks_endpoint_returns_them(self):
|
|
"""HTTP 全链路:带 reference_image_ids 提交 → 任务 payload 落 ids → tasks 接口把参考图解析回 {name,url}。
|
|
守护「上传的参考图被带进生成 + 切换/刷新后批次头仍能回显参考图」。"""
|
|
CreditAccount.objects.create(team=self.team, balance="100.0000")
|
|
patch("apps.ai.tasks.generate_standalone_image_task.delay").start()
|
|
self.addCleanup(patch.stopall)
|
|
ref = Asset.objects.create(
|
|
team=self.team, created_by=self.user, name="背心参考", asset_type=Asset.Type.IMAGE,
|
|
source=Asset.Source.UPLOAD, category=Asset.Category.UPLOAD,
|
|
)
|
|
AssetFile.objects.create(asset=ref, object_key="r.png", bucket="b", content_type="image/png", preview_url="http://x/ref.png", is_primary=True)
|
|
r = self.client.post(
|
|
"/api/ai/generate-image/",
|
|
{"prompt": "按这件背心生成不同场景", "mode": "image", "count": 1, "reference_image_ids": [str(ref.id)]},
|
|
format="json",
|
|
)
|
|
self.assertEqual(r.status_code, 202, r.content)
|
|
conv_id = r.json()["conversation_id"]
|
|
# 任务 payload 落了参考图 id
|
|
task = AITask.objects.filter(conversation_id=conv_id).first()
|
|
self.assertEqual(task.request_payload.get("reference_image_ids"), [str(ref.id)])
|
|
# tasks 接口把参考图解析成 {id,name,url} 回显 —— id 供前端重跑时原样复用(否则重跑丢参考图)
|
|
data = self.client.get(f"/api/ai/image-conversations/{conv_id}/tasks/").json()["tasks"]
|
|
self.assertEqual(data[0]["reference_images"], [{"id": str(ref.id), "name": "背心参考", "url": "http://x/ref.png"}])
|
|
|
|
def test_rerun_with_batch_id_reuses_batch_and_marks_append(self):
|
|
"""重跑/补图带原批次 batch_id → 新任务沿用同一 batch_id 并打 batch_append 标记(tasks 接口回
|
|
rerun=True)—— 守护「重跑不裂新批次卡 + 失败格计数不膨胀」;非法 batch_id 则铸新批次兜底。"""
|
|
CreditAccount.objects.create(team=self.team, balance="100.0000")
|
|
patch("apps.ai.tasks.generate_standalone_image_task.delay").start()
|
|
self.addCleanup(patch.stopall)
|
|
first = self.client.post(
|
|
"/api/ai/generate-image/", {"prompt": "一只猫", "mode": "image", "count": 2}, format="json"
|
|
)
|
|
self.assertEqual(first.status_code, 202, first.content)
|
|
conv_id = first.json()["conversation_id"]
|
|
batch_id = first.json()["batch_id"]
|
|
self.assertTrue(batch_id)
|
|
# 重跑一张:带原对话 + 原 batch_id → 归回原批次
|
|
second = self.client.post(
|
|
"/api/ai/generate-image/",
|
|
{"prompt": "一只猫", "mode": "image", "count": 1, "conversation_id": conv_id, "batch_id": batch_id},
|
|
format="json",
|
|
)
|
|
self.assertEqual(second.status_code, 202, second.content)
|
|
self.assertEqual(second.json()["batch_id"], batch_id)
|
|
tasks = self.client.get(f"/api/ai/image-conversations/{conv_id}/tasks/").json()["tasks"]
|
|
self.assertEqual(len(tasks), 3)
|
|
self.assertTrue(all(t["batch_id"] == batch_id for t in tasks))
|
|
self.assertEqual([t["rerun"] for t in tasks], [False, False, True])
|
|
# 非法 batch_id(不是 UUID):不沿用,铸新批次,不打 append 标记
|
|
third = self.client.post(
|
|
"/api/ai/generate-image/",
|
|
{"prompt": "一只猫", "mode": "image", "count": 1, "conversation_id": conv_id, "batch_id": "not-a-uuid"},
|
|
format="json",
|
|
)
|
|
self.assertEqual(third.status_code, 202, third.content)
|
|
self.assertNotEqual(third.json()["batch_id"], batch_id)
|
|
newest = self.client.get(f"/api/ai/image-conversations/{conv_id}/tasks/").json()["tasks"][-1]
|
|
self.assertFalse(newest["rerun"])
|
|
|
|
|
|
class StandaloneCategoryTests(TestCase):
|
|
"""图片趴三类归类(期2):模特上身图→model_tryon / 平台套图→platform_kit / 自由创作→free_create;
|
|
生成演员(model 无 product)仍→person(视频角色)。直接跑 worker 函数避开异步 .delay。"""
|
|
|
|
def setUp(self):
|
|
self.user = User.objects.create_user(username="catowner", password="pass")
|
|
self.team = Team.objects.create(name="CT", owner=self.user)
|
|
CreditAccount.objects.create(team=self.team, balance="100.0000")
|
|
self.product = Product.objects.create(team=self.team, created_by=self.user, title="测试商品")
|
|
cover = Asset.objects.create(
|
|
team=self.team, created_by=self.user, name="主图", asset_type=Asset.Type.IMAGE,
|
|
source=Asset.Source.UPLOAD, category=Asset.Category.PRODUCT_IMAGE,
|
|
)
|
|
AssetFile.objects.create(asset=cover, object_key="c.png", bucket="b", content_type="image/png", preview_url="http://x/cover.png", is_primary=True)
|
|
self.product.cover_asset = cover
|
|
self.product.save(update_fields=["cover_asset"])
|
|
|
|
def _patch_provider(self):
|
|
provider = patch("apps.ai.services.get_image_provider").start()
|
|
prov = provider.return_value
|
|
prov.image_edit.return_value = {"data": [{"url": "http://x/out.png"}]}
|
|
prov.image_generation.return_value = {"data": [{"url": "http://x/out.png"}]}
|
|
prov.extract_first_media_url.return_value = "http://x/out.png"
|
|
media = patch("apps.ai.services.VolcanoArkProvider.media_to_bytes").start()
|
|
media.return_value = (BytesIO(b"img"), "image/png")
|
|
store = patch("apps.ai.services.TosStorage").start()
|
|
stored = store.return_value.upload_fileobj.return_value
|
|
stored.object_key, stored.bucket, stored.content_type, stored.size_bytes = "o.png", "b", "image/png", 3
|
|
self.addCleanup(patch.stopall)
|
|
return prov
|
|
|
|
def _run(self, mode, *, product_id=None, model_id=None):
|
|
from apps.ai.services import run_standalone_image_task
|
|
|
|
self._patch_provider()
|
|
tasks = enqueue_standalone_images(
|
|
team=self.team, user=self.user, prompt="x", mode=mode, count=1,
|
|
product_id=product_id, model_id=model_id, model_entity_id="ent-1", ratio="4:5",
|
|
)
|
|
run_standalone_image_task(task_id=str(tasks[0].id))
|
|
return Asset.objects.filter(origin_task=tasks[0]).first()
|
|
|
|
def test_model_tryon_category(self):
|
|
a = self._run("model", product_id=str(self.product.id))
|
|
self.assertEqual(a.category, Asset.Category.MODEL_TRYON)
|
|
self.assertEqual(a.metadata.get("mode"), "model")
|
|
self.assertTrue(a.metadata.get("batch_id")) # 成组
|
|
self.assertEqual(a.metadata.get("model_entity_id"), "ent-1") # 溯源模特库
|
|
self.assertTrue(a.in_library) # R109:图片生成成图自动入资产库
|
|
|
|
def test_platform_kit_category(self):
|
|
a = self._run("cover", product_id=str(self.product.id))
|
|
self.assertEqual(a.category, Asset.Category.PLATFORM_KIT)
|
|
self.assertTrue(a.in_library) # R109
|
|
|
|
def test_free_create_category(self):
|
|
a = self._run("image")
|
|
self.assertEqual(a.category, Asset.Category.FREE_CREATE)
|
|
self.assertTrue(a.in_library) # R109
|
|
|
|
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.MODEL_PORTRAIT)
|
|
self.assertFalse(a.in_library)
|
|
|
|
@override_settings(MODEL_TRYON_PROMPT_V2_ENABLED=True)
|
|
def test_model_tryon_v2_provider_failure_releases_reserved_credit(self):
|
|
from apps.ai.services import run_standalone_image_task
|
|
|
|
self.product.title = "女装裤子"
|
|
self.product.save(update_fields=["title"])
|
|
prov = self._patch_provider()
|
|
prov.image_edit.side_effect = ValueError("生成失败")
|
|
with patch("apps.ai.tasks.generate_standalone_image_task.delay"):
|
|
task = enqueue_standalone_images(
|
|
team=self.team,
|
|
user=self.user,
|
|
prompt="全身站姿,双腿完整露出",
|
|
mode="model",
|
|
count=1,
|
|
product_id=str(self.product.id),
|
|
ratio="3:4",
|
|
)[0]
|
|
|
|
run_standalone_image_task(task_id=str(task.id))
|
|
task.refresh_from_db()
|
|
account = CreditAccount.objects.get(team=self.team)
|
|
self.assertEqual(task.status, AITask.Status.FAILED)
|
|
self.assertEqual(float(task.actual_cost), 0.0)
|
|
self.assertEqual(float(account.balance), 100.0)
|
|
self.assertEqual(float(account.reserved_balance), 0.0)
|
|
self.assertTrue(task.request_payload["tryon_prompt"]["applied"])
|
|
self.assertTrue(
|
|
CreditLedger.objects.filter(task=task, ledger_type=CreditLedger.Type.RELEASE).exists()
|
|
)
|
|
|
|
|
|
class TriviewModelDecouplingTests(TestCase):
|
|
"""项目角色三视图只归项目,不自动创建或更新团队模特。"""
|
|
|
|
def setUp(self):
|
|
from apps.projects.models import Project
|
|
|
|
self.user = User.objects.create_user(username="trio", password="p")
|
|
self.team = Team.objects.create(name="TR", owner=self.user)
|
|
CreditAccount.objects.create(team=self.team, balance="100.0000")
|
|
self.product = Product.objects.create(team=self.team, created_by=self.user, title="P")
|
|
self.project = Project.objects.create(team=self.team, name="项目甲", product=self.product, created_by=self.user)
|
|
self.portrait = 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,
|
|
)
|
|
AssetFile.objects.create(asset=self.portrait, object_key="p.png", bucket="b", content_type="image/png", preview_url="http://x/p.png", is_primary=True)
|
|
|
|
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()
|
|
media.return_value = (BytesIO(b"img"), "image/png")
|
|
store = patch("apps.ai.services.TosStorage").start()
|
|
stored = store.return_value.upload_fileobj.return_value
|
|
stored.object_key, stored.bucket, stored.content_type, stored.size_bytes = "o.png", "b", "image/png", 3
|
|
self.addCleanup(patch.stopall)
|
|
|
|
def _gen_triview(self):
|
|
from apps.ai.services import generate_person_triview, run_triview_task
|
|
|
|
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_stays_on_project_and_does_not_enroll_model(self):
|
|
from apps.assets.models import Model
|
|
|
|
self._patch_provider()
|
|
self._gen_triview()
|
|
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_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()
|
|
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:
|
|
"""模拟 requests 流式响应:支持 with、raise_for_status、可写 encoding、iter_lines。"""
|
|
status_code = 200
|
|
encoding = None
|
|
|
|
def __init__(self, lines):
|
|
self._lines = lines
|
|
|
|
def __enter__(self):
|
|
return self
|
|
|
|
def __exit__(self, *exc):
|
|
return False
|
|
|
|
def raise_for_status(self):
|
|
return None
|
|
|
|
def iter_lines(self, decode_unicode=True): # noqa: ARG002
|
|
yield from self._lines
|
|
|
|
|
|
class ChatStreamReasoningTests(SimpleTestCase):
|
|
"""推理模型(豆包 seed-pro 等)思考期只发 reasoning_content、不发 content。
|
|
provider 必须把它作为独立 `reasoning` 事件转发——否则脚本 agent 思考期零输出 =
|
|
前端「卡在生成分镜」假死。本测试锁住该转发,防回归。"""
|
|
|
|
def test_reasoning_content_forwarded_as_reasoning_event(self):
|
|
def _chunk(delta):
|
|
return "data: " + json.dumps({"choices": [{"delta": delta}]}, ensure_ascii=False)
|
|
|
|
lines = [
|
|
_chunk({"reasoning_content": "先想想"}),
|
|
_chunk({"reasoning_content": "用户要4镜"}),
|
|
_chunk({"content": "正在生成"}),
|
|
_chunk({"content": "脚本…"}),
|
|
"data: [DONE]",
|
|
]
|
|
prov = VolcanoArkProvider(api_key="k", base_url="http://x")
|
|
with patch("apps.ai.providers.volcano.requests.post", return_value=_FakeStreamResp(lines)):
|
|
events = list(prov.chat_completion_stream(model="m", messages=[{"role": "user", "content": "hi"}]))
|
|
|
|
self.assertEqual([e["type"] for e in events], ["reasoning", "reasoning", "delta", "delta", "done"])
|
|
self.assertEqual([e["text"] for e in events if e["type"] == "reasoning"], ["先想想", "用户要4镜"])
|
|
self.assertEqual("".join(e["text"] for e in events if e["type"] == "delta"), "正在生成脚本…")
|
|
|
|
|
|
class WorkbenchAndUnreadTests(TestCase):
|
|
"""R100(工作台记录后端持久化)+ R96(未读生成任务角标)+ R109(删除联动)端到端:
|
|
· /api/ai/tasks/workbench/ 按 mode+product 回放批次流(与任务中心同源),软删的图不回显;
|
|
· /api/ai/tasks/unread/ 只统计 mode∈model/cover/image 的未读任务,按团队隔离;
|
|
· /api/ai/tasks/mark-read/ 按商品清零。"""
|
|
|
|
def setUp(self):
|
|
self.user = User.objects.create_user(username="wbowner", password="pass")
|
|
self.team = Team.objects.create(name="WB", owner=self.user)
|
|
TeamMember.objects.create(team=self.team, user=self.user, role="owner", status="active")
|
|
CreditAccount.objects.create(team=self.team, balance="100.0000")
|
|
self.product = Product.objects.create(team=self.team, created_by=self.user, title="工作台商品")
|
|
cover = Asset.objects.create(
|
|
team=self.team, created_by=self.user, name="主图", asset_type=Asset.Type.IMAGE,
|
|
source=Asset.Source.UPLOAD, category=Asset.Category.PRODUCT_IMAGE,
|
|
)
|
|
AssetFile.objects.create(asset=cover, object_key="c.png", bucket="b", content_type="image/png", preview_url="http://x/cover.png", is_primary=True)
|
|
self.product.cover_asset = cover
|
|
self.product.save(update_fields=["cover_asset"])
|
|
self.client = APIClient()
|
|
self.client.force_authenticate(self.user)
|
|
|
|
def _patch_provider(self):
|
|
provider = patch("apps.ai.services.get_image_provider").start()
|
|
prov = provider.return_value
|
|
prov.image_edit.return_value = {"data": [{"url": "http://x/out.png"}]}
|
|
prov.image_generation.return_value = {"data": [{"url": "http://x/out.png"}]}
|
|
prov.extract_first_media_url.return_value = "http://x/out.png"
|
|
media = patch("apps.ai.services.VolcanoArkProvider.media_to_bytes").start()
|
|
media.return_value = (BytesIO(b"img"), "image/png")
|
|
store = patch("apps.ai.services.TosStorage").start()
|
|
stored = store.return_value.upload_fileobj.return_value
|
|
stored.object_key, stored.bucket, stored.content_type, stored.size_bytes = "o.png", "b", "image/png", 3
|
|
self.addCleanup(patch.stopall)
|
|
return prov
|
|
|
|
def _generate(self, mode="cover", count=2, platform_id=None):
|
|
from apps.ai.services import run_standalone_image_task
|
|
|
|
self._patch_provider()
|
|
tasks = enqueue_standalone_images(
|
|
team=self.team, user=self.user, prompt="平台套图测试", mode=mode, count=count,
|
|
product_id=str(self.product.id), ratio="1:1", platform_id=platform_id,
|
|
)
|
|
for t in tasks:
|
|
run_standalone_image_task(task_id=str(t.id))
|
|
return tasks
|
|
|
|
def test_workbench_returns_batches_and_hides_deleted_assets(self):
|
|
tasks = self._generate(mode="cover", count=2, platform_id="douyin")
|
|
r = self.client.get(f"/api/ai/tasks/workbench/?mode=cover&product={self.product.id}")
|
|
self.assertEqual(r.status_code, 200, r.content)
|
|
rows = r.json()["tasks"]
|
|
self.assertEqual(len(rows), 2)
|
|
# 同一批共享 batch_id;prompt/ratio/platform/product 从 payload 抽出;每任务带成图
|
|
self.assertEqual(len({row["batch_id"] for row in rows}), 1)
|
|
self.assertEqual(rows[0]["prompt"], "平台套图测试")
|
|
self.assertEqual(rows[0]["ratio"], "1:1")
|
|
self.assertEqual(rows[0]["platform_id"], "douyin")
|
|
self.assertEqual(rows[0]["product_id"], str(self.product.id))
|
|
self.assertEqual(len(rows[0]["assets"]), 1)
|
|
# 其他 mode / 其他商品查不到这批
|
|
self.assertEqual(self.client.get("/api/ai/tasks/workbench/?mode=model").json()["tasks"], [])
|
|
# R108/R109:软删其中一张 → 该任务的 assets 随之为空(删除联动),资产不再回显
|
|
asset = Asset.objects.filter(origin_task=tasks[0]).first()
|
|
d = self.client.delete(f"/api/assets/{asset.id}/")
|
|
self.assertEqual(d.status_code, 204, d.content)
|
|
asset.refresh_from_db()
|
|
self.assertTrue(asset.is_deleted) # 回收站软删,不是物理删除
|
|
rows = self.client.get(f"/api/ai/tasks/workbench/?mode=cover&product={self.product.id}").json()["tasks"]
|
|
by_id = {row["id"]: row for row in rows}
|
|
self.assertEqual(by_id[str(tasks[0].id)]["assets"], [])
|
|
self.assertEqual(len(by_id[str(tasks[1].id)]["assets"]), 1)
|
|
|
|
def test_workbench_rejects_unknown_mode(self):
|
|
# mode 白名单:脚本 agent 复用的 auto/theme/revise 不能混进来
|
|
self.assertEqual(self.client.get("/api/ai/tasks/workbench/?mode=auto").status_code, 400)
|
|
self.assertEqual(self.client.get("/api/ai/tasks/workbench/").status_code, 400)
|
|
|
|
def test_unread_counts_and_mark_read(self):
|
|
self._generate(mode="cover", count=2, platform_id="douyin")
|
|
r = self.client.get("/api/ai/tasks/unread/")
|
|
self.assertEqual(r.status_code, 200, r.content)
|
|
body = r.json()
|
|
self.assertEqual(body["total"], 2)
|
|
self.assertEqual(body["by_product"], {str(self.product.id): 2})
|
|
# 按商品标记已读 → 清零
|
|
m = self.client.post("/api/ai/tasks/mark-read/", {"product_id": str(self.product.id)}, format="json")
|
|
self.assertEqual(m.status_code, 200, m.content)
|
|
self.assertEqual(m.json()["updated"], 2)
|
|
body = self.client.get("/api/ai/tasks/unread/").json()
|
|
self.assertEqual(body["total"], 0)
|
|
self.assertEqual(body["by_product"], {})
|
|
|
|
def test_unread_is_team_scoped_and_mode_whitelisted(self):
|
|
self._generate(mode="cover", count=1)
|
|
# 非生成任务(脚本 agent 复用 mode=auto)不计入未读
|
|
AITask.objects.create(
|
|
team=self.team, created_by=self.user, task_type=AITask.Type.SCRIPT_GENERATION,
|
|
model_config=ModelConfig.objects.first(), idempotency_key="script-x",
|
|
request_payload={"mode": "auto"},
|
|
)
|
|
self.assertEqual(self.client.get("/api/ai/tasks/unread/").json()["total"], 1)
|
|
# 另一团队的用户看不到本团队的未读
|
|
other = User.objects.create_user(username="wbother", password="pass")
|
|
other_team = Team.objects.create(name="WB2", owner=other)
|
|
TeamMember.objects.create(team=other_team, user=other, role="owner", status="active")
|
|
c2 = APIClient()
|
|
c2.force_authenticate(other)
|
|
self.assertEqual(c2.get("/api/ai/tasks/unread/").json()["total"], 0)
|