diff --git a/src/app/widgets/ai_outfit_panel.py b/src/app/widgets/ai_outfit_panel.py index 1441adc..257dc00 100644 --- a/src/app/widgets/ai_outfit_panel.py +++ b/src/app/widgets/ai_outfit_panel.py @@ -14,6 +14,7 @@ from PySide6.QtWidgets import ( QAbstractItemView, QCheckBox, QComboBox, + QInputDialog, QDoubleSpinBox, QFileDialog, QGridLayout, @@ -39,9 +40,10 @@ from PySide6.QtWidgets import ( ) from services.config_service import ( + DEFAULT_OUTFIT_PROMPT, load_ai_models, - load_outfit_prompt, - save_outfit_prompt, + load_outfit_prompts, + save_outfit_prompts, ) logger = logging.getLogger(__name__) @@ -148,6 +150,9 @@ class AiOutfitPanel(QWidget): self._failures = [] # list of OutfitResult (failed) self._last_resolution = "" # for "value actually changed" check (§10.2) self._last_model = "" + self._prompts = [] # list of {name, text} (§7.2) + self._current_prompt_name = "" + self._saved_text = "" # stored text of the selected template (dirty check) self._build_ui() # -- construction --------------------------------------------------- @@ -187,9 +192,21 @@ class AiOutfitPanel(QWidget): self._output_edit.setPlaceholderText("默认:程序旁的「合并后的图片」") col.addLayout(self._inline_path_row("输出", self._output_edit, self._browse_output)) - # 通用话术(加大,随窗口高度拉伸) + # 通用话术(多套模板;加大,随窗口高度拉伸) prm = QGroupBox("通用话术") pv = QVBoxLayout(prm) + # 模板选择行:下拉 + 新建/另存为/重命名/删除 + self._prompt_combo = QComboBox() + self._compact_combo(self._prompt_combo) + self._prompt_combo.activated.connect(self._on_prompt_template_activated) + pv.addWidget(self._prompt_combo) + trow = QHBoxLayout() + for text, slot in (("新建", self._prompt_new), ("另存为", self._prompt_save_as), + ("重命名", self._prompt_rename), ("删除", self._prompt_delete)): + b = QPushButton(text) + b.clicked.connect(slot) + trow.addWidget(b) + pv.addLayout(trow) self._prompt_edit = QPlainTextEdit() self._prompt_edit.setMinimumHeight(150) self._prompt_edit.textChanged.connect(self._refresh_preview) @@ -387,7 +404,14 @@ class AiOutfitPanel(QWidget): self._set_combo(self._resolution, config.get("outfit_resolution", "1K")) self._set_combo(self._quality, config.get("outfit_quality", "均衡")) - self._prompt_edit.setPlainText(load_outfit_prompt()) + # 话术模板:载入多套 + 选中上次(§7.2) + self._prompts = load_outfit_prompts() + names = [p["name"] for p in self._prompts] + name = config.get("outfit_prompt_name", "") + if name not in names: + name = names[0] + self._rebuild_prompt_combo() + self._apply_prompt(name) self._models = load_ai_models() self._model_combo.clear() @@ -425,6 +449,7 @@ class AiOutfitPanel(QWidget): "outfit_resolution": self._resolution.currentText(), "outfit_quality": self._quality.currentText(), "outfit_retry_failed": self._retry_failed_chk.isChecked(), + "outfit_prompt_name": self._current_prompt_name, }) # -- left actions --------------------------------------------------- @@ -445,9 +470,143 @@ class AiOutfitPanel(QWidget): self._emit_config() def _save_prompt(self): - save_outfit_prompt(self._prompt_edit.toPlainText()) + self._store_current_text() self.statusBar_message("话术已保存") + # -- 话术模板(§7.2)------------------------------------------------ + + def _rebuild_prompt_combo(self): + """Refill the template dropdown from self._prompts (no signal).""" + self._prompt_combo.blockSignals(True) + self._prompt_combo.clear() + self._prompt_combo.addItems([p["name"] for p in self._prompts]) + self._prompt_combo.blockSignals(False) + + def _select_prompt_in_combo(self, name): + idx = self._prompt_combo.findText(name) + if idx >= 0: + self._prompt_combo.blockSignals(True) + self._prompt_combo.setCurrentIndex(idx) + self._prompt_combo.blockSignals(False) + + def _prompt_text(self, name): + return next((p["text"] for p in self._prompts if p["name"] == name), "") + + def _apply_prompt(self, name): + """Load template *name* into the editor (no persistence).""" + self._current_prompt_name = name + self._saved_text = self._prompt_text(name) + self._select_prompt_in_combo(name) + self._prompt_edit.setPlainText(self._saved_text) # fires _refresh_preview + + def _store_current_text(self): + """Save the editor text into the current template + persist to disk.""" + text = self._prompt_edit.toPlainText() + for p in self._prompts: + if p["name"] == self._current_prompt_name: + p["text"] = text + break + save_outfit_prompts(self._prompts) + self._saved_text = text + + def _is_dirty(self): + return self._prompt_edit.toPlainText() != self._saved_text + + def _maybe_save_dirty(self): + """Handle unsaved edits before switching away. Return False = cancel.""" + if not self._is_dirty(): + return True + ans = QMessageBox.question( + self, "未保存", "当前话术「{}」有未保存的修改,是否保存?".format( + self._current_prompt_name), + QMessageBox.Save | QMessageBox.Discard | QMessageBox.Cancel, + QMessageBox.Save) + if ans == QMessageBox.Cancel: + return False + if ans == QMessageBox.Save: + self._store_current_text() + return True + + def _name_exists(self, name): + return any(p["name"] == name for p in self._prompts) + + def _ask_name(self, title, default=""): + """Prompt for a unique non-empty template name; None if cancelled/invalid.""" + name, ok = QInputDialog.getText(self, title, "模板名称:", text=default) + if not ok: + return None + name = name.strip() + if not name: + QMessageBox.warning(self, "名称无效", "模板名称不能为空。") + return None + if self._name_exists(name): + QMessageBox.warning(self, "名称重复", "已存在同名模板:{}".format(name)) + return None + return name + + def _on_prompt_template_activated(self, index): + name = self._prompt_combo.itemText(index) + if name == self._current_prompt_name: + return + if not self._maybe_save_dirty(): + self._select_prompt_in_combo(self._current_prompt_name) # cancel: revert + return + self._apply_prompt(name) + self._emit_config() + + def _prompt_new(self): + if not self._maybe_save_dirty(): + return + name = self._ask_name("新建话术") + if name is None: + return + self._prompts.append({"name": name, "text": DEFAULT_OUTFIT_PROMPT}) + save_outfit_prompts(self._prompts) + self._rebuild_prompt_combo() + self._apply_prompt(name) + self._emit_config() + + def _prompt_save_as(self): + name = self._ask_name("另存为", default=self._current_prompt_name) + if name is None: + return + self._prompts.append({"name": name, "text": self._prompt_edit.toPlainText()}) + save_outfit_prompts(self._prompts) + self._rebuild_prompt_combo() + self._apply_prompt(name) + self._emit_config() + + def _prompt_rename(self): + new = self._ask_name("重命名", default=self._current_prompt_name) + if new is None: + return + for p in self._prompts: + if p["name"] == self._current_prompt_name: + p["name"] = new + break + self._current_prompt_name = new + save_outfit_prompts(self._prompts) + self._rebuild_prompt_combo() + self._select_prompt_in_combo(new) + self._emit_config() + + def _prompt_delete(self): + if len(self._prompts) <= 1: + QMessageBox.information(self, "无法删除", "至少保留一套话术。") + return + ans = QMessageBox.question( + self, "删除话术", "确定删除话术「{}」?".format(self._current_prompt_name), + QMessageBox.Yes | QMessageBox.No, QMessageBox.No) + if ans != QMessageBox.Yes: + return + idx = next((i for i, p in enumerate(self._prompts) + if p["name"] == self._current_prompt_name), 0) + self._prompts.pop(idx) + save_outfit_prompts(self._prompts) + self._rebuild_prompt_combo() + self._apply_prompt(self._prompts[min(idx, len(self._prompts) - 1)]["name"]) + self._emit_config() + # -- inline prompt preview ------------------------------------------ def _fill_sample_combo(self, tasks): @@ -547,7 +706,7 @@ class AiOutfitPanel(QWidget): if answer != QMessageBox.Yes: return - save_outfit_prompt(prompt) + self._store_current_text() # persist editor into the selected template self._emit_config() output_dir = self._output_edit.text().strip() diff --git a/src/services/config_service.py b/src/services/config_service.py index faa2026..dab0d27 100644 --- a/src/services/config_service.py +++ b/src/services/config_service.py @@ -26,13 +26,16 @@ DEFAULT_CONFIG = { "outfit_resolution": "1K", "outfit_quality": "均衡", "outfit_retry_failed": False, + "outfit_prompt_name": "默认", # last-selected 话术模板 name (docs/11 §7.2) } _CONFIG_FILENAME = "app_config.json" _AI_MODELS_FILENAME = "ai_models.json" -_OUTFIT_PROMPT_FILENAME = "outfit_prompt.txt" +_OUTFIT_PROMPT_FILENAME = "outfit_prompt.txt" # legacy single prompt (migrated) +_OUTFIT_PROMPTS_FILENAME = "outfit_prompts.json" # multi named templates (§7.2) +DEFAULT_OUTFIT_PROMPT_NAME = "默认" -# Default outfit prompt (docs/11 §7). Persisted to outfit_prompt.txt on first save. +# Default outfit prompt (docs/11 §7). Seeded into outfit_prompts.json on first run. DEFAULT_OUTFIT_PROMPT = ( "为商品「{title}」生成人物上身实穿效果图:真人模特正面穿着这件衣服," "完整保留款式、版型、颜色与印花图案,自然光、纯色棚拍背景," @@ -139,3 +142,51 @@ def save_outfit_prompt(text): logger.info("Outfit prompt saved to %s", prompt_file) except OSError as exc: logger.error("Failed to save outfit prompt to %s: %s", prompt_file, exc) + + +def _normalize_prompts(data): + """Keep only valid {name, text} entries (non-empty name, string text).""" + if not isinstance(data, list): + return [] + out = [] + for item in data: + if isinstance(item, dict): + name = str(item.get("name", "")).strip() + text = item.get("text", "") + if name and isinstance(text, str): + out.append({"name": name, "text": text}) + return out + + +def load_outfit_prompts(): + """Load named 话术 templates (docs/11 §7.2); always returns >= 1. + + Missing/corrupt → migrate the legacy outfit_prompt.txt into a single + 「默认」template, or seed it from DEFAULT_OUTFIT_PROMPT. + """ + from services.file_service import get_config_path + prompts_file = get_config_path(_OUTFIT_PROMPTS_FILENAME) + if prompts_file.exists(): + try: + with open(str(prompts_file), encoding="utf-8-sig") as f: + prompts = _normalize_prompts(json.load(f)) + if prompts: + return prompts + except (json.JSONDecodeError, ValueError, OSError) as exc: + logger.warning("Outfit prompts unreadable (%s): %s", exc, prompts_file) + + # Seed / migrate (load_outfit_prompt reads the legacy txt or the default). + return [{"name": DEFAULT_OUTFIT_PROMPT_NAME, "text": load_outfit_prompt()}] + + +def save_outfit_prompts(prompts): + """Persist named 话术 templates as JSON (utf-8, no BOM). Does not raise.""" + from services.file_service import get_config_path + prompts_file = get_config_path(_OUTFIT_PROMPTS_FILENAME) + try: + prompts_file.parent.mkdir(parents=True, exist_ok=True) + with open(str(prompts_file), "w", encoding="utf-8") as f: + json.dump(list(prompts), f, ensure_ascii=False, indent=2) + logger.info("Outfit prompts saved to %s", prompts_file) + except OSError as exc: + logger.error("Failed to save outfit prompts to %s: %s", prompts_file, exc) diff --git a/tasks.md b/tasks.md index 811e2d5..b942110 100644 --- a/tasks.md +++ b/tasks.md @@ -1109,8 +1109,8 @@ 需求:把单一话术升级为多套命名话术(下拉切换 / 新建 / 另存为 / 重命名 / 删除 / 保存,记住上次)。决策:**全部自定义**(不分内置,播种一套「默认」);**切换前脏数据弹窗**提醒是否保存;**含重命名**。 -- [ ] `config_service`:`outfit_prompts.json` 读写 `load_outfit_prompts()` / `save_outfit_prompts(list)`(utf-8-sig 读 / 无 BOM 写 / 损坏回退「默认」;始终 ≥1 套);`app_config` 加 `outfit_prompt_name`(当前选中名);首次迁移旧 `outfit_prompt.txt` → 一套「默认」,否则 `DEFAULT_OUTFIT_PROMPT` -- [ ] `ai_outfit_panel.py`:「通用话术」组加 模板下拉 + 新建 / 另存为 / 重命名 / 删除(保留 插入标题 / 保存);切换载入文本 + 刷新预览;记住上次所选并启动恢复(经 `config_changed` 回主窗口集中存) -- [ ] 脏数据保护:编辑框与当前套已存文本不同则在切换/重命名/删除/新建/另存为前弹「保存 / 不保存 / 取消」(取消还原下拉);名字唯一;删后选邻近、不可删到 0 -- [ ] 「开始生成」前把当前编辑存回所选套;生成/预览仍走 `render_prompt`(§7.1 尾巴不变) -- [ ] 单测 `tests/test_config_service.py`:prompts 读写、迁移(有/无旧 txt)、损坏回退、≥1 套;离屏验证切换/脏弹窗/重命名/删除 +- [x] `config_service`:`outfit_prompts.json` 读写 `load_outfit_prompts()` / `save_outfit_prompts(list)`(utf-8-sig 读 / 无 BOM 写 / 损坏回退「默认」;始终 ≥1 套;`_normalize_prompts` 丢弃非法项);`app_config` 加 `outfit_prompt_name`(当前选中名);首次迁移旧 `outfit_prompt.txt` → 一套「默认」,否则 `DEFAULT_OUTFIT_PROMPT` +- [x] `ai_outfit_panel.py`:「通用话术」组加 模板下拉 + 新建 / 另存为 / 重命名 / 删除(保留 插入标题 / 保存);切换载入文本 + 刷新预览;记住上次所选并启动恢复(经 `config_changed` 回主窗口集中存 `outfit_prompt_name`) +- [x] 脏数据保护:编辑框与当前套已存文本不同则在切换/新建/另存为前弹「保存 / 不保存 / 取消」(取消还原下拉);名字唯一(空/重名拒绝);删后选邻近、不可删到 0(删除单独二次确认) +- [x] 「开始生成」前把当前编辑存回所选套(`_store_current_text`);生成/预览仍走 `render_prompt`(§7.1 尾巴不变) +- [x] 单测 `tests/test_config_service.py` +6(seed/迁移/损坏回退/roundtrip 无 BOM/丢弃非法/BOM 容错,共 14);离屏验证 新建·另存为·切换·脏保存持久化·重命名·删除·不可删到 0;全套 12 文件绿 diff --git a/tests/test_config_service.py b/tests/test_config_service.py index 64e4f60..33a9ff5 100644 --- a/tests/test_config_service.py +++ b/tests/test_config_service.py @@ -68,6 +68,49 @@ class TestOutfitConfigHelpers(unittest.TestCase): raw = (self.config_dir / "outfit_prompt.txt").read_bytes() self.assertFalse(raw.startswith(b"\xef\xbb\xbf")) + # -- outfit_prompts.json (multi templates, §7.2) -------------------- + + def test_prompts_seed_default_when_missing(self): + prompts = cs.load_outfit_prompts() + self.assertEqual(len(prompts), 1) + self.assertEqual(prompts[0]["name"], cs.DEFAULT_OUTFIT_PROMPT_NAME) + self.assertEqual(prompts[0]["text"], cs.DEFAULT_OUTFIT_PROMPT) + + def test_prompts_migrate_legacy_txt(self): + (self.config_dir / "outfit_prompt.txt").write_text( + "旧话术 {title}", encoding="utf-8") + prompts = cs.load_outfit_prompts() + self.assertEqual(len(prompts), 1) + self.assertEqual(prompts[0]["name"], "默认") + self.assertEqual(prompts[0]["text"], "旧话术 {title}") + + def test_prompts_roundtrip_and_no_bom(self): + data = [{"name": "A", "text": "甲 {title}"}, {"name": "B", "text": "乙"}] + cs.save_outfit_prompts(data) + self.assertEqual(cs.load_outfit_prompts(), data) + raw = (self.config_dir / "outfit_prompts.json").read_bytes() + self.assertFalse(raw.startswith(b"\xef\xbb\xbf")) + + def test_prompts_corrupt_falls_back(self): + (self.config_dir / "outfit_prompts.json").write_text("{ bad", encoding="utf-8") + prompts = cs.load_outfit_prompts() + self.assertEqual(len(prompts), 1) + self.assertEqual(prompts[0]["name"], cs.DEFAULT_OUTFIT_PROMPT_NAME) + + def test_prompts_drop_invalid_entries(self): + import json + (self.config_dir / "outfit_prompts.json").write_text( + json.dumps([{"name": "", "text": "x"}, {"name": "ok", "text": "y"}, + {"name": "bad", "text": 5}, "nope"]), encoding="utf-8") + prompts = cs.load_outfit_prompts() + self.assertEqual(prompts, [{"name": "ok", "text": "y"}]) + + def test_prompts_tolerate_bom(self): + import json + (self.config_dir / "outfit_prompts.json").write_bytes( + "".encode("utf-8") + json.dumps([{"name": "z", "text": "t"}]).encode("utf-8")) + self.assertEqual(cs.load_outfit_prompts()[0]["name"], "z") + if __name__ == "__main__": unittest.main()