Files
cmshoppe/tests/test_prompts.py
T

269 lines
13 KiB
Python
Raw Normal View History

2026-06-27 16:10:50 +08:00
import os
import sys
import unittest
from unittest import mock
2026-06-27 16:10:50 +08:00
sys.path.insert(0, os.path.dirname(__file__))
from _helpers import TempDirMixin
from app import prompts
class PromptTests(TempDirMixin, unittest.TestCase):
def test_title_prompt_load_missing_and_save(self):
with self.make_temp_dir() as temp_dir:
path = os.path.join(temp_dir, "title_prompt.txt")
self.assertEqual("", prompts.load_title_prompt(path))
prompts.save_title_prompt("标题规则", path)
self.assertEqual("标题规则", prompts.load_title_prompt(path))
self.assert_removed(temp_dir)
def test_cover_template_crud_and_validation(self):
with self.make_temp_dir() as temp_dir:
directory = os.path.join(temp_dir, "prompts", "cover")
self.assertEqual([], prompts.list_cover_templates(directory))
prompts.save_cover_template("韩版女装", "封面规则", directory)
prompts.save_cover_template("基础", "基础规则", directory)
self.assertEqual(["基础", "韩版女装"], prompts.list_cover_templates(directory))
self.assertEqual("封面规则", prompts.load_cover_template("韩版女装", directory))
prompts.rename_cover_template("基础", "基础2", directory)
self.assertEqual(["基础2", "韩版女装"], prompts.list_cover_templates(directory))
with self.assertRaises(prompts.PromptError):
prompts.rename_cover_template("基础2", "韩版女装", directory)
with self.assertRaises(prompts.PromptError):
prompts.save_cover_template("../bad", "x", directory)
prompts.delete_cover_template("基础2", directory)
self.assertEqual(["韩版女装"], prompts.list_cover_templates(directory))
self.assert_removed(temp_dir)
2026-07-10 11:43:35 +08:00
def test_generic_template_crud_and_title_cover_directories_are_independent(self):
with self.make_temp_dir() as temp_dir:
title_dir = os.path.join(temp_dir, "prompts", "title")
cover_dir = os.path.join(temp_dir, "prompts", "cover")
prompts.save_template("基础", "通用模板", title_dir)
prompts.save_title_template("标题A", "标题模板", title_dir)
prompts.save_cover_template("封面A", "封面模板", cover_dir)
self.assertEqual(["基础", "标题A"], prompts.list_title_templates(title_dir))
self.assertEqual(["封面A"], prompts.list_cover_templates(cover_dir))
self.assertEqual("通用模板", prompts.load_template("基础", title_dir))
self.assertEqual("标题模板", prompts.load_title_template("标题A", title_dir))
self.assertEqual("封面模板", prompts.load_cover_template("封面A", cover_dir))
prompts.rename_template("基础", "基础2", title_dir)
self.assertEqual(["基础2", "标题A"], prompts.list_templates(title_dir))
with self.assertRaisesRegex(prompts.PromptError, "提示词模板已存在"):
prompts.rename_template("基础2", "标题A", title_dir)
with self.assertRaisesRegex(prompts.PromptError, "提示词模板名非法"):
prompts.save_title_template("../bad", "x", title_dir)
prompts.delete_title_template("标题A", title_dir)
self.assertEqual(["基础2"], prompts.list_title_templates(title_dir))
self.assertEqual(["封面A"], prompts.list_cover_templates(cover_dir))
self.assert_removed(temp_dir)
def test_image_studio_template_crud_isolated_and_preserves_raw_text(self):
with self.make_temp_dir() as temp_dir:
studio_dir = os.path.join(temp_dir, "prompts", "image_studio")
title_dir = os.path.join(temp_dir, "prompts", "title")
cover_dir = os.path.join(temp_dir, "prompts", "cover")
raw_text = "完整提示词第一行\n第二行:不要拆分\n{商品id}"
prompts.save_image_studio_template("工场模板", raw_text, studio_dir)
prompts.save_title_template("标题模板", "标题", title_dir)
prompts.save_cover_template("封面模板", "封面", cover_dir)
self.assertEqual(["工场模板"], prompts.list_image_studio_templates(studio_dir))
self.assertEqual(raw_text, prompts.load_image_studio_template("工场模板", studio_dir))
self.assertEqual(["标题模板"], prompts.list_title_templates(title_dir))
self.assertEqual(["封面模板"], prompts.list_cover_templates(cover_dir))
prompts.rename_image_studio_template("工场模板", "工场模板2", studio_dir)
self.assertEqual(["工场模板2"], prompts.list_image_studio_templates(studio_dir))
with self.assertRaisesRegex(prompts.PromptError, "提示词模板已存在"):
prompts.rename_image_studio_template("工场模板2", "工场模板2", studio_dir)
with self.assertRaisesRegex(prompts.PromptError, "提示词模板名不能为空"):
prompts.save_image_studio_template("", "x", studio_dir)
with self.assertRaisesRegex(prompts.PromptError, "提示词模板名非法"):
prompts.save_image_studio_template("../bad", "x", studio_dir)
prompts.delete_image_studio_template("工场模板2", studio_dir)
self.assertEqual([], prompts.list_image_studio_templates(studio_dir))
with self.assertRaisesRegex(prompts.PromptError, "提示词模板不存在"):
prompts.load_image_studio_template("工场模板2", studio_dir)
self.assert_removed(temp_dir)
2026-07-10 09:00:06 +08:00
def test_ensure_default_prompts_seeds_empty_user_prompt_files(self):
with self.make_temp_dir() as temp_dir:
title_path = os.path.join(temp_dir, "title_prompt.txt")
2026-07-10 11:43:35 +08:00
title_templates_dir = os.path.join(temp_dir, "prompts", "title")
2026-07-10 09:00:06 +08:00
cover_dir = os.path.join(temp_dir, "prompts", "cover")
prompts.ensure_default_prompts(title_path, cover_dir)
self.assertIn("蝦皮台灣站", prompts.load_title_prompt(title_path))
2026-07-10 11:43:35 +08:00
self.assertEqual(["默认"], prompts.list_title_templates(title_templates_dir))
self.assertEqual(
prompts.load_title_prompt(title_path),
prompts.load_title_template("默认", title_templates_dir),
)
2026-07-10 09:00:06 +08:00
self.assertEqual(["papa1"], prompts.list_cover_templates(cover_dir))
self.assertIn(
"商品标题:{新标题}",
prompts.load_cover_template("papa1", cover_dir),
)
self.assert_removed(temp_dir)
def test_ensure_default_prompts_does_not_overwrite_user_prompts(self):
with self.make_temp_dir() as temp_dir:
title_path = os.path.join(temp_dir, "title_prompt.txt")
2026-07-10 11:43:35 +08:00
title_templates_dir = os.path.join(temp_dir, "prompts", "title")
2026-07-10 09:00:06 +08:00
cover_dir = os.path.join(temp_dir, "prompts", "cover")
prompts.save_title_prompt("用户标题提示词", title_path)
2026-07-10 11:43:35 +08:00
prompts.save_title_template("用户标题模板", "用户标题模板内容", title_templates_dir)
2026-07-10 09:00:06 +08:00
prompts.save_cover_template("用户模板", "用户封面提示词", cover_dir)
prompts.ensure_default_prompts(title_path, cover_dir)
self.assertEqual("用户标题提示词", prompts.load_title_prompt(title_path))
2026-07-10 11:43:35 +08:00
self.assertEqual(["用户标题模板"], prompts.list_title_templates(title_templates_dir))
self.assertEqual(
"用户标题模板内容",
prompts.load_title_template("用户标题模板", title_templates_dir),
)
2026-07-10 09:00:06 +08:00
self.assertEqual(["用户模板"], prompts.list_cover_templates(cover_dir))
self.assertEqual(
"用户封面提示词",
prompts.load_cover_template("用户模板", cover_dir),
)
self.assert_removed(temp_dir)
2026-07-10 11:43:35 +08:00
def test_ensure_default_prompts_accepts_explicit_title_templates_dir(self):
with self.make_temp_dir() as temp_dir:
title_path = os.path.join(temp_dir, "custom", "title_prompt.txt")
cover_dir = os.path.join(temp_dir, "custom", "cover")
title_templates_dir = os.path.join(temp_dir, "other", "title_templates")
prompts.ensure_default_prompts(title_path, cover_dir, title_templates_dir)
self.assertEqual(["默认"], prompts.list_title_templates(title_templates_dir))
self.assertFalse(os.path.exists(os.path.join(temp_dir, "custom", "prompts", "title")))
self.assert_removed(temp_dir)
2026-06-27 16:10:50 +08:00
def test_render_prompt_replaces_known_variables(self):
task = {
"old_title": "舊T恤",
"new_title": "新T恤",
"item_id": "51100639510",
"account_name": "主店",
}
rendered = prompts.render_prompt(
"用{旧标题}生成{新标题},商品{商品id},店铺{店铺},未知{不存在}",
task,
)
self.assertEqual(
"用舊T恤生成新T恤,商品51100639510,店铺主店,未知{不存在}",
rendered,
)
def test_product_suite_prompt_seed_save_and_restore_are_isolated(self):
with self.make_temp_dir() as temp_dir:
path = os.path.join(temp_dir, "prompts", "product_suite", "base.txt")
default_text = prompts.load_default_product_suite_prompt()
self.assertEqual(default_text, prompts.ensure_default_product_suite_prompt(path))
self.assertEqual(default_text, prompts.load_product_suite_prompt(path))
custom_text = "自定义说明\n" + default_text
prompts.save_product_suite_prompt(custom_text, path)
self.assertEqual(custom_text, prompts.ensure_default_product_suite_prompt(path))
self.assertEqual(custom_text, prompts.load_product_suite_prompt(path))
self.assertEqual(default_text, prompts.restore_default_product_suite_prompt(path))
self.assertEqual(default_text, prompts.load_product_suite_prompt(path))
self.assert_removed(temp_dir)
def test_product_suite_prompt_invalid_user_file_is_preserved(self):
with self.make_temp_dir() as temp_dir:
path = os.path.join(temp_dir, "prompts", "product_suite", "base.txt")
os.makedirs(os.path.dirname(path), exist_ok=True)
invalid = "坏模板{未知变量}"
with open(path, "w", encoding="utf-8") as handle:
handle.write(invalid)
with self.assertRaisesRegex(prompts.PromptError, "套图提示词模板无效"):
prompts.ensure_default_product_suite_prompt(path)
with open(path, "r", encoding="utf-8") as handle:
self.assertEqual(invalid, handle.read())
self.assert_removed(temp_dir)
def test_product_suite_prompt_unreadable_utf8_is_not_overwritten(self):
with self.make_temp_dir() as temp_dir:
path = os.path.join(temp_dir, "prompts", "product_suite", "base.txt")
os.makedirs(os.path.dirname(path), exist_ok=True)
raw = b"\xff\xfe\x00\x80"
with open(path, "wb") as handle:
handle.write(raw)
with self.assertRaisesRegex(prompts.PromptError, "读取失败"):
prompts.ensure_default_product_suite_prompt(path)
with open(path, "rb") as handle:
self.assertEqual(raw, handle.read())
self.assert_removed(temp_dir)
def test_product_suite_prompt_invalid_packaged_default_does_not_create_file(self):
with self.make_temp_dir() as temp_dir:
path = os.path.join(temp_dir, "prompts", "product_suite", "base.txt")
with mock.patch.object(
prompts,
"load_default_product_suite_prompt",
side_effect=prompts.PromptError(prompts.PRODUCT_SUITE_DEFAULT_ERROR),
):
with self.assertRaisesRegex(prompts.PromptError, "内置套图提示词模板无效"):
prompts.ensure_default_product_suite_prompt(path)
self.assertFalse(os.path.exists(path))
self.assert_removed(temp_dir)
def test_product_suite_prompt_atomic_save_failure_keeps_original(self):
with self.make_temp_dir() as temp_dir:
path = os.path.join(temp_dir, "prompts", "product_suite", "base.txt")
default_text = prompts.ensure_default_product_suite_prompt(path)
custom_text = "自定义说明\n" + default_text
with mock.patch("app.prompts.os.replace", side_effect=OSError("文件占用")):
with self.assertRaisesRegex(prompts.PromptError, "保存失败"):
prompts.save_product_suite_prompt(custom_text, path)
self.assertEqual(default_text, prompts.load_product_suite_prompt(path))
self.assertFalse(
any(name.startswith(".prompt-") for name in os.listdir(os.path.dirname(path)))
)
self.assert_removed(temp_dir)
2026-06-27 16:10:50 +08:00
if __name__ == "__main__":
unittest.main()