feat: add optional cover generation toggle

This commit is contained in:
chengma
2026-07-02 15:17:11 +08:00
parent f9c6dd3331
commit 21cd02c167
13 changed files with 348 additions and 48 deletions
+55
View File
@@ -247,6 +247,7 @@ class AITests(TempDirMixin, unittest.TestCase):
cfg["image_dir"] = os.path.join(temp_dir, "images")
cfg["ai"]["title_concurrency"] = 2
cfg["ai"]["image_concurrency"] = 2
cfg["ai"]["generate_cover"] = True
batch_id, tasks = self._collected_tasks(temp_dir, cfg)
cover_prompts = []
progress = []
@@ -306,11 +307,65 @@ class AITests(TempDirMixin, unittest.TestCase):
self.assert_removed(temp_dir)
def test_generate_batch_can_skip_cover_generation(self):
with self.make_temp_dir() as temp_dir:
cfg = self._config()
cfg["db_path"] = os.path.join(temp_dir, "cmshopee.db")
cfg["image_dir"] = os.path.join(temp_dir, "images")
batch_id, tasks = self._collected_tasks(temp_dir, cfg)
events = []
progress = []
def fake_title(title_prompt, old_title, **kwargs):
self.assertEqual("标题提示", title_prompt)
return "新" + old_title
with mock.patch("app.ai.gen_title", side_effect=fake_title), \
mock.patch("app.ai.gen_cover") as gen_cover:
summary = ai.generate_batch(
tasks,
{"title": "标题提示", "cover": "封面 {新标题}"},
ai_cfg={
"config": cfg,
"db_path": cfg["db_path"],
"on_event": events.append,
},
on_progress=progress.append,
)
self.assertTrue(summary["ok"])
self.assertFalse(summary["generate_cover"])
self.assertEqual(2, summary["title_done"])
self.assertEqual(0, summary["cover_done"])
self.assertEqual(0, summary["cover_total"])
self.assertEqual(2, summary["generated_done"])
self.assertEqual(0, summary["failed"])
gen_cover.assert_not_called()
updated = db.list_tasks(batch_id=batch_id, path=cfg["db_path"])
self.assertTrue(all(task.stage == "generated" for task in updated))
self.assertTrue(all(task.status == "success" for task in updated))
self.assertEqual({"新旧标题A", "新旧标题B"}, {task.new_title for task in updated})
self.assertTrue(all(task.new_cover_path is None for task in updated))
self.assertEqual(0, progress[-1]["cover_total"])
self.assertEqual(2, progress[-1]["generated_done"])
self.assertTrue(
any(
event.get("phase") == "title"
and event.get("step") == "db_write"
and event.get("result") == "success"
and event.get("detail") == "仅生成标题"
for event in events
)
)
self.assert_removed(temp_dir)
def test_generate_batch_marks_failed_task_without_blocking_others(self):
with self.make_temp_dir() as temp_dir:
cfg = self._config()
cfg["db_path"] = os.path.join(temp_dir, "cmshopee.db")
cfg["image_dir"] = os.path.join(temp_dir, "images")
cfg["ai"]["generate_cover"] = True
batch_id, tasks = self._collected_tasks(temp_dir, cfg, ["好标题", "坏标题"])
def fake_title(title_prompt, old_title, **kwargs):
+1
View File
@@ -19,6 +19,7 @@ class AppConfigTests(TempDirMixin, unittest.TestCase):
self.assertTrue(os.path.exists(config_path))
self.assertEqual("images", appconfig.image_dir(config))
self.assertEqual(240, appconfig.response_timeout(config))
self.assertFalse(appconfig.ai_config(config)["generate_cover"])
updated = appconfig.update_config(
{"ai": {"resolution": "2k"}},
+94 -2
View File
@@ -17,7 +17,7 @@ if gui.QT_IMPORT_ERROR is not None:
raise unittest.SkipTest("PySide6 未安装")
from PySide6.QtGui import QTextCursor
from PySide6.QtWidgets import QApplication, QLineEdit, QPlainTextEdit, QProgressBar, QTableView
from PySide6.QtWidgets import QApplication, QCheckBox, QLineEdit, QPlainTextEdit, QProgressBar, QTableView
from app.gui import (
AccountDialog,
@@ -437,6 +437,7 @@ class GuiTests(TempDirMixin, unittest.TestCase):
cfg["ai"] = appconfig.default_config()["ai"]
cfg["ai"]["default_text_model"] = "Text A"
cfg["ai"]["default_image_model"] = "Image A"
cfg["ai"]["generate_cover"] = True
models_path = cfg["ai_models_path"]
config_path = cfg["config_path"]
appconfig.save_ai_models_config(
@@ -548,6 +549,7 @@ class GuiTests(TempDirMixin, unittest.TestCase):
saved = appconfig.load_config(config_path)
self.assertEqual("Text B", saved["ai"]["default_text_model"])
self.assertEqual("Image B", saved["ai"]["default_image_model"])
self.assertTrue(saved["ai"]["generate_cover"])
self.assertEqual(3, saved["ai"]["title_concurrency"])
self.assertEqual(2, saved["ai"]["image_concurrency"])
self.assertEqual(1, saved["ai"]["retry"])
@@ -688,14 +690,17 @@ class GuiTests(TempDirMixin, unittest.TestCase):
self.assertIsInstance(tab.title_prompt_edit, QPlainTextEdit)
self.assertIsInstance(tab.cover_prompt_edit, QPlainTextEdit)
self.assertIsInstance(tab.task_table, QTableView)
self.assertIsInstance(tab.generate_cover_checkbox, QCheckBox)
self.assertEqual("标题提示词", tab.title_prompt_edit.placeholderText())
self.assertEqual("封面提示词", tab.cover_prompt_edit.placeholderText())
self.assertEqual("保存标题提示词", tab.save_title_button.text())
self.assertEqual("开始生成", tab.generate_button.text())
self.assertEqual("停止", tab.stop_generate_button.text())
self.assertEqual("重置生成结果", tab.reset_generate_button.text())
self.assertEqual("生成封面图片(成本较高)", tab.generate_cover_checkbox.text())
self.assertFalse(tab.generate_cover_checkbox.isChecked())
self.assertFalse(tab.stop_generate_button.isEnabled())
self.assertEqual("进度:标题0/0 · 封面0/0 · 失败0", tab.progress_label.text())
self.assertEqual("进度:标题0/0 · 图片0/0 · 失败0", tab.progress_label.text())
self.assertIsInstance(tab.title_progress_bar, QProgressBar)
self.assertIsInstance(tab.cover_progress_bar, QProgressBar)
self.assertEqual("标题 0/0", tab.title_progress_label.text())
@@ -720,6 +725,23 @@ class GuiTests(TempDirMixin, unittest.TestCase):
self.assert_removed(temp_dir)
def test_generate_tab_persists_generate_cover_toggle(self):
with self.make_temp_dir() as temp_dir:
cfg = self.make_config(temp_dir)
statuses = []
tab = GenerateTab(config=cfg, status_callback=statuses.append)
self.addCleanup(tab.close)
self.assertFalse(appconfig.ai_config(cfg)["generate_cover"])
tab.generate_cover_checkbox.setChecked(True)
saved = appconfig.load_config(cfg["config_path"])
self.assertTrue(saved["ai"]["generate_cover"])
self.assertTrue(cfg["ai"]["generate_cover"])
self.assertIn("会同时生成封面图片", statuses[-1])
self.assert_removed(temp_dir)
def test_generate_tab_manages_prompt_files_and_preview(self):
with self.make_temp_dir() as temp_dir:
cfg = self.make_config(temp_dir)
@@ -893,6 +915,8 @@ class GuiTests(TempDirMixin, unittest.TestCase):
def test_generate_worker_calls_generate_batch_and_emits_signals(self):
with self.make_temp_dir() as temp_dir:
cfg = self.make_config(temp_dir)
cfg["ai"] = appconfig.ai_config(cfg)
cfg["ai"]["generate_cover"] = True
account = accounts.create_account("主店", "alias-a", debug_port=9222, config=cfg)
batch_id = db.create_batch(["input.xlsx"], path=cfg["db_path"])
db.insert_tasks(
@@ -985,6 +1009,8 @@ class GuiTests(TempDirMixin, unittest.TestCase):
def test_generate_worker_writes_run_log_and_diagnostic_log_on_image_failure(self):
with self.make_temp_dir() as temp_dir:
cfg = self.make_config(temp_dir)
cfg["ai"] = appconfig.ai_config(cfg)
cfg["ai"]["generate_cover"] = True
accounts.create_account("主店", "alias-a", debug_port=9222, config=cfg)
batch_id = db.create_batch(["input.xlsx"], path=cfg["db_path"])
db.insert_tasks(
@@ -1058,6 +1084,72 @@ class GuiTests(TempDirMixin, unittest.TestCase):
self.assert_removed(temp_dir)
def test_generate_worker_can_generate_titles_without_covers(self):
with self.make_temp_dir() as temp_dir:
cfg = self.make_config(temp_dir)
accounts.create_account("主店", "alias-a", debug_port=9222, config=cfg)
batch_id = db.create_batch(["input.xlsx"], path=cfg["db_path"])
db.insert_tasks(
batch_id,
[
{
"source_file_abs": os.path.join(temp_dir, "input.xlsx"),
"source_sheet": "商品",
"source_row": 2,
"account_name": "Excel主店",
"alias": "alias-a",
"item_id": "51100639510",
}
],
path=cfg["db_path"],
)
task = db.list_tasks(batch_id=batch_id, path=cfg["db_path"])[0]
db.set_collected(task.id, "旧标题", "old.jpg", path=cfg["db_path"])
tasks = db.list_tasks(batch_id=batch_id, path=cfg["db_path"])
logs = []
progress = []
def fake_title(title_prompt, old_title, **kwargs):
return "新标题"
worker = GenerateWorker(
tasks,
{"title": "标题提示", "cover": "封面提示"},
db_path=cfg["db_path"],
config=cfg,
)
worker.log.connect(logs.append)
worker.progress.connect(progress.append)
with mock.patch("app.ai.gen_title", side_effect=fake_title), \
mock.patch("app.ai.gen_cover") as gen_cover:
summary = worker.execute()
self.assertTrue(summary["ok"])
self.assertFalse(summary["generate_cover"])
self.assertEqual(1, summary["title_done"])
self.assertEqual(0, summary["cover_done"])
self.assertEqual(0, summary["cover_total"])
self.assertEqual(1, summary["generated_done"])
self.assertEqual(0, summary["failed"])
gen_cover.assert_not_called()
updated = db.list_tasks(batch_id=batch_id, path=cfg["db_path"])[0]
self.assertEqual("generated", updated.stage)
self.assertEqual("success", updated.status)
self.assertEqual("新标题", updated.new_title)
self.assertIsNone(updated.new_cover_path)
self.assertEqual(0, progress[-1]["cover_total"])
self.assertEqual(1, progress[-1]["generated_done"])
joined_logs = "\n".join(logs)
self.assertIn("本轮仅生成标题,不生成图片", joined_logs)
self.assertIn("已保存,仅生成标题", joined_logs)
self.assertIn("[完成] AI 生成完成:标题1/1,图片0/0,失败0", joined_logs)
run_log = db.list_run_logs(limit=1, run_type="generate", path=cfg["db_path"])[0]
self.assertEqual(1, run_log.done)
self.assertEqual(1, run_log.success_count)
self.assertTrue(run_log.options["generate_cover"] is False)
self.assert_removed(temp_dir)
def test_generate_tab_loads_latest_generate_run_log(self):
with self.make_temp_dir() as temp_dir:
cfg = self.make_config(temp_dir)