diff --git a/ui/config_tab.py b/ui/config_tab.py index 6f111d4..c313ae3 100644 --- a/ui/config_tab.py +++ b/ui/config_tab.py @@ -163,7 +163,31 @@ def build_ai_config_tab(self): if new_name != config_name: self.loaded_config["llm_configs"][new_name] = self.loaded_config["llm_configs"].pop(config_name) refresh_config_dropdown() - + embedding_config = { + "api_key": self.embedding_api_key_var.get(), + "base_url": self.embedding_url_var.get(), + "model_name": self.embedding_model_name_var.get(), + "retrieval_k": self.safe_get_int(self.embedding_retrieval_k_var, 4), + "interface_format": self.embedding_interface_format_var.get().strip() + + } + other_params = { + "topic": self.topic_text.get("0.0", "end").strip(), + "genre": self.genre_var.get(), + "num_chapters": self.safe_get_int(self.num_chapters_var, 10), + "word_number": self.safe_get_int(self.word_number_var, 3000), + "filepath": self.filepath_var.get(), + "chapter_num": self.chapter_num_var.get(), + "user_guidance": self.user_guide_text.get("0.0", "end").strip(), + "characters_involved": self.characters_involved_var.get(), + "key_items": self.key_items_var.get(), + "scene_location": self.scene_location_var.get(), + "time_constraint": self.time_constraint_var.get() + } + self.loaded_config["embedding_configs"][self.embedding_interface_format_var.get().strip()] = embedding_config + self.loaded_config["other_params"] = other_params + + # 保存到JSON文件 try: save_config(self.loaded_config, self.config_file) @@ -484,7 +508,7 @@ def build_embeddings_config_tab(self): # 1) Embedding API Key create_label_with_help(self, parent=self.embeddings_config_tab, label_text="Embedding API Key:", tooltip_key="embedding_api_key", row=0, column=0, font=("Microsoft YaHei", 12)) - emb_api_key_entry = ctk.CTkEntry(self.embeddings_config_tab, textvariable=self.embedding_api_key_var, font=("Microsoft YaHei", 12)) + emb_api_key_entry = ctk.CTkEntry(self.embeddings_config_tab, textvariable=self.embedding_api_key_var, font=("Microsoft YaHei", 12), show="*") emb_api_key_entry.grid(row=0, column=1, padx=5, pady=5, sticky="nsew") # 2) Embedding 接口格式 diff --git a/ui/generation_handlers.py b/ui/generation_handlers.py index c43fdce..9d86ec6 100644 --- a/ui/generation_handlers.py +++ b/ui/generation_handlers.py @@ -6,6 +6,7 @@ import tkinter as tk from tkinter import messagebox import customtkinter as ctk import traceback +import glob from utils import read_file, save_string_to_txt, clear_file_content from novel_generator import ( Novel_architecture_generate, @@ -14,7 +15,8 @@ from novel_generator import ( finalize_chapter, import_knowledge_file, clear_vector_store, - enrich_chapter_text + enrich_chapter_text, + build_chapter_prompt ) from consistency_checker import check_consistency @@ -157,7 +159,6 @@ def generate_chapter_draft_ui(self): self.safe_log(f"生成第{chap_num}章草稿:准备生成请求提示词...") # 调用新添加的 build_chapter_prompt 函数构造初始提示词 - from novel_generator.chapter import build_chapter_prompt prompt_text = build_chapter_prompt( api_key=api_key, base_url=base_url, @@ -443,6 +444,223 @@ def do_consistency_check(self): finally: self.enable_button_safe(self.btn_check_consistency) threading.Thread(target=task, daemon=True).start() +def generate_batch_ui(self): + def open_batch_dialog(): + dialog = tk.Toplevel() + chapter_file = os.path.join(self.filepath_var.get().strip(), "chapters") + files = glob.glob(os.path.join(chapter_file, "chapter_*.txt")) + if not files: + num = 1 + else: + num = max(int(os.path.basename(f).split('_')[1].split('.')[0]) for f in files) + 1 + dialog.geometry("+500+400") + tk.Label(dialog, text="起始章节").grid(row=0, column=0) + entry_start = tk.Entry(dialog) + entry_start.grid(row=0, column=1) + entry_start.insert(0, str(num)) + tk.Label(dialog, text="结束章节").grid(row=0, column=2) + entry_end = tk.Entry(dialog) + entry_end.grid(row=0, column=3) + tk.Label(dialog, text="期望字数").grid(row=1, column=0) + entry_word = tk.Entry(dialog) + entry_word.grid(row=1, column=1) + entry_word.insert(0, self.word_number_var.get()) + tk.Label(dialog, text="最低字数").grid(row=1, column=2) + entry_min = tk.Entry(dialog) + entry_min.grid(row=1, column=3) + entry_min.insert(0, self.word_number_var.get()) + + auto_enrich_bool = tk.BooleanVar() + auto_enrich_bool_ck = tk.Checkbutton(dialog, text="低于最低字数时自动扩写", variable=auto_enrich_bool) + auto_enrich_bool_ck.grid(row=2, column=0) + + result = {"start": None, "end": None, "word": None, "min": None, "auto_enrich": None, "close": False} + + + def on_confirm(): + nonlocal result + if not entry_start.get() or not entry_end.get() or not entry_word.get() or not entry_min.get(): + messagebox.showwarning("警告", "请填写完整信息。") + return + + result = { + "start": entry_start.get(), + "end": entry_end.get(), + "word": entry_word.get(), + "min": entry_min.get(), + "auto_enrich": auto_enrich_bool.get(), + "close": False + } + dialog.destroy() + + def on_cancel(): + nonlocal result + result["close"] = True + dialog.destroy() + tk.Button(dialog, text="确认", command=on_confirm).grid(row=2, column=1) + dialog.protocol("WM_DELETE_WINDOW", on_cancel) + dialog.transient(self.master) + dialog.grab_set() + dialog.wait_window(dialog) + return result + + def generate_chapter_batch(self ,i ,word, min, auto_enrich): + draft_interface_format = self.loaded_config["llm_configs"][self.prompt_draft_llm_var.get()]["interface_format"] + draft_api_key = self.loaded_config["llm_configs"][self.prompt_draft_llm_var.get()]["api_key"] + draft_base_url = self.loaded_config["llm_configs"][self.prompt_draft_llm_var.get()]["base_url"] + draft_model_name = self.loaded_config["llm_configs"][self.prompt_draft_llm_var.get()]["model_name"] + draft_temperature = self.loaded_config["llm_configs"][self.prompt_draft_llm_var.get()]["temperature"] + draft_max_tokens = self.loaded_config["llm_configs"][self.prompt_draft_llm_var.get()]["max_tokens"] + draft_timeout = self.loaded_config["llm_configs"][self.prompt_draft_llm_var.get()]["timeout"] + user_guidance = self.user_guide_text.get("0.0", "end").strip() + + char_inv = self.characters_involved_var.get().strip() + key_items = self.key_items_var.get().strip() + scene_loc = self.scene_location_var.get().strip() + time_constr = self.time_constraint_var.get().strip() + + embedding_api_key = self.embedding_api_key_var.get().strip() + embedding_url = self.embedding_url_var.get().strip() + embedding_interface_format = self.embedding_interface_format_var.get().strip() + embedding_model_name = self.embedding_model_name_var.get().strip() + embedding_k = self.safe_get_int(self.embedding_retrieval_k_var, 4) + + prompt_text = build_chapter_prompt( + api_key=draft_api_key, + base_url=draft_base_url, + model_name=draft_model_name, + filepath=self.filepath_var.get().strip(), + novel_number=i, + word_number=word, + temperature=draft_temperature, + user_guidance=user_guidance, + characters_involved=char_inv, + key_items=key_items, + scene_location=scene_loc, + time_constraint=time_constr, + embedding_api_key=embedding_api_key, + embedding_url=embedding_url, + embedding_interface_format=embedding_interface_format, + embedding_model_name=embedding_model_name, + embedding_retrieval_k=embedding_k, + interface_format=draft_interface_format, + max_tokens=draft_max_tokens, + timeout=draft_timeout, + ) + final_prompt = prompt_text + role_names = [name.strip() for name in self.char_inv_text.get("0.0", "end").split("\n")] + role_lib_path = os.path.join(self.filepath_var.get().strip(), "角色库") + role_contents = [] + if os.path.exists(role_lib_path): + for root, dirs, files in os.walk(role_lib_path): + for file in files: + if file.endswith(".txt") and os.path.splitext(file)[0] in role_names: + file_path = os.path.join(root, file) + try: + with open(file_path, 'r', encoding='utf-8') as f: + role_contents.append(f.read().strip()) # 直接使用文件内容,不添加重复名字 + except Exception as e: + self.safe_log(f"读取角色文件 {file} 失败: {str(e)}") + if role_contents: + role_content_str = "\n".join(role_contents) + # 更精确的替换逻辑,处理不同情况下的占位符 + placeholder_variations = [ + "核心人物(可能未指定):{characters_involved}", + "核心人物:{characters_involved}", + "核心人物(可能未指定):{characters_involved}", + "核心人物:{characters_involved}" + ] + + for placeholder in placeholder_variations: + if placeholder in final_prompt: + final_prompt = final_prompt.replace( + placeholder, + f"核心人物:\n{role_content_str}" + ) + break + else: # 如果没有找到任何已知占位符变体 + lines = final_prompt.split('\n') + for i, line in enumerate(lines): + if "核心人物" in line and ":" in line: + lines[i] = f"核心人物:\n{role_content_str}" + break + final_prompt = '\n'.join(lines) + draft_text = generate_chapter_draft( + api_key=draft_api_key, + base_url=draft_base_url, + model_name=draft_model_name, + filepath=self.filepath_var.get().strip(), + novel_number=i, + word_number=word, + temperature=draft_temperature, + user_guidance=user_guidance, + characters_involved=char_inv, + key_items=key_items, + scene_location=scene_loc, + time_constraint=time_constr, + embedding_api_key=embedding_api_key, + embedding_url=embedding_url, + embedding_interface_format=embedding_interface_format, + embedding_model_name=embedding_model_name, + embedding_retrieval_k=embedding_k, + interface_format=draft_interface_format, + max_tokens=draft_max_tokens, + timeout=draft_timeout, + custom_prompt_text=final_prompt + ) + + finalize_interface_format = self.loaded_config["llm_configs"][self.final_chapter_llm_var.get()]["interface_format"] + finalize_api_key = self.loaded_config["llm_configs"][self.final_chapter_llm_var.get()]["api_key"] + finalize_base_url = self.loaded_config["llm_configs"][self.final_chapter_llm_var.get()]["base_url"] + finalize_model_name = self.loaded_config["llm_configs"][self.final_chapter_llm_var.get()]["model_name"] + finalize_temperature = self.loaded_config["llm_configs"][self.final_chapter_llm_var.get()]["temperature"] + finalize_max_tokens = self.loaded_config["llm_configs"][self.final_chapter_llm_var.get()]["max_tokens"] + finalize_timeout = self.loaded_config["llm_configs"][self.final_chapter_llm_var.get()]["timeout"] + + chapters_dir = os.path.join(self.filepath_var.get().strip(), "chapters") + os.makedirs(chapters_dir, exist_ok=True) + chapter_path = os.path.join(chapters_dir, f"chapter_{i}.txt") + if len(draft_text) < 0.7 * min and auto_enrich: + self.safe_log(f"第{i}章草稿字数 ({len(draft_text)}) 低于目标字数({min})的70%,正在扩写...") + enriched = enrich_chapter_text( + chapter_text=draft_text, + word_number=word, + api_key=draft_api_key, + base_url=draft_base_url, + model_name=draft_model_name, + temperature=draft_temperature, + interface_format=draft_interface_format, + max_tokens=draft_max_tokens, + timeout=draft_timeout + ) + draft_text = enriched + clear_file_content(chapter_path) + save_string_to_txt(draft_text, chapter_path) + finalize_chapter( + novel_number=i, + word_number=word, + api_key=finalize_api_key, + base_url=finalize_base_url, + model_name=finalize_model_name, + temperature=finalize_temperature, + filepath=self.filepath_var.get().strip(), + embedding_api_key=embedding_api_key, + embedding_url=embedding_url, + embedding_interface_format=embedding_interface_format, + embedding_model_name=embedding_model_name, + interface_format=finalize_interface_format, + max_tokens=finalize_max_tokens, + timeout=finalize_timeout + ) + + + result = open_batch_dialog() + if result["close"]: + return + + for i in range(int(result["start"]), int(result["end"]) + 1): + generate_chapter_batch(self, i, int(result["word"]), int(result["min"]), result["auto_enrich"]) + def import_knowledge_handler(self): selected_file = tk.filedialog.askopenfilename( diff --git a/ui/main_tab.py b/ui/main_tab.py index a5f66d4..55dfdf9 100644 --- a/ui/main_tab.py +++ b/ui/main_tab.py @@ -54,7 +54,8 @@ def build_left_layout(self): # Step 按钮区域 self.step_buttons_frame = ctk.CTkFrame(self.left_frame) self.step_buttons_frame.grid(row=2, column=0, sticky="ew", padx=5, pady=5) - self.step_buttons_frame.columnconfigure((0, 1, 2, 3), weight=1) + self.step_buttons_frame.columnconfigure((0, 1, 2, 3, 4), weight=1) + self.btn_generate_architecture = ctk.CTkButton( self.step_buttons_frame, @@ -88,6 +89,15 @@ def build_left_layout(self): ) self.btn_finalize_chapter.grid(row=0, column=3, padx=5, pady=2, sticky="ew") + self.btn_batch_generate = ctk.CTkButton( + self.step_buttons_frame, + text="批量生成", + command=self.generate_batch_ui, + font=("Microsoft YaHei", 12) + ) + self.btn_batch_generate.grid(row=0, column=4, padx=5, pady=2, sticky="ew") + + # 日志文本框 log_label = ctk.CTkLabel(self.left_frame, text="输出日志 (只读)", font=("Microsoft YaHei", 12)) log_label.grid(row=3, column=0, padx=5, pady=(5, 0), sticky="w") diff --git a/ui/main_window.py b/ui/main_window.py index 11cdf88..630e723 100644 --- a/ui/main_window.py +++ b/ui/main_window.py @@ -26,7 +26,8 @@ from ui.generation_handlers import ( do_consistency_check, import_knowledge_handler, clear_vectorstore_handler, - show_plot_arcs_ui + show_plot_arcs_ui, + generate_batch_ui ) from ui.setting_tab import build_setting_tab, load_novel_architecture, save_novel_architecture from ui.directory_tab import build_directory_tab, load_chapter_blueprint, save_chapter_blueprint @@ -364,6 +365,7 @@ class NovelGeneratorGUI: generate_chapter_draft_ui = generate_chapter_draft_ui finalize_chapter_ui = finalize_chapter_ui do_consistency_check = do_consistency_check + generate_batch_ui = generate_batch_ui import_knowledge_handler = import_knowledge_handler clear_vectorstore_handler = clear_vectorstore_handler show_plot_arcs_ui = show_plot_arcs_ui