feat: add optional cover generation toggle
This commit is contained in:
@@ -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):
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user