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):
|
||||
|
||||
Reference in New Issue
Block a user