Merge pull request #184 from ahhhhhhhman/main
1.增加在不同环节使用不同的大模型配置的功能 2.增加批量生成章节的功能
This commit is contained in:
@@ -12,3 +12,7 @@ config_test.json
|
||||
/novel_generator/__pycache__
|
||||
/ui/__pycache__
|
||||
.idea/
|
||||
/novel
|
||||
app.log
|
||||
test.py
|
||||
/backup
|
||||
|
||||
@@ -0,0 +1,62 @@
|
||||
{
|
||||
"last_interface_format": "OpenAI",
|
||||
"last_embedding_interface_format": "OpenAI",
|
||||
"llm_configs": {
|
||||
"DeepSeek V3": {
|
||||
"api_key": "",
|
||||
"base_url": "https://api.deepseek.com/v1",
|
||||
"model_name": "deepseek-chat",
|
||||
"temperature": 0.7,
|
||||
"max_tokens": 8192,
|
||||
"timeout": 600,
|
||||
"interface_format": "OpenAI"
|
||||
},
|
||||
"GPT 5": {
|
||||
"api_key": "",
|
||||
"base_url": "https://api.openai.com/v1",
|
||||
"model_name": "gpt-5",
|
||||
"temperature": 0.7,
|
||||
"max_tokens": 32768,
|
||||
"timeout": 600,
|
||||
"interface_format": "OpenAI"
|
||||
},
|
||||
"Gemini 2.5 Pro": {
|
||||
"api_key": "",
|
||||
"base_url": "https://generativelanguage.googleapis.com/v1beta/openai",
|
||||
"model_name": "gemini-2.5-pro",
|
||||
"temperature": 0.7,
|
||||
"max_tokens": 32768,
|
||||
"timeout": 600,
|
||||
"interface_format": "OpenAI"
|
||||
}
|
||||
},
|
||||
"embedding_configs": {
|
||||
"OpenAI": {
|
||||
"api_key": "",
|
||||
"base_url": "https://api.openai.com/v1",
|
||||
"model_name": "text-embedding-ada-002",
|
||||
"retrieval_k": 4,
|
||||
"interface_format": "OpenAI"
|
||||
}
|
||||
},
|
||||
"other_params": {
|
||||
"topic": "",
|
||||
"genre": "",
|
||||
"num_chapters": 0,
|
||||
"word_number": 0,
|
||||
"filepath": "",
|
||||
"chapter_num": "120",
|
||||
"user_guidance": "",
|
||||
"characters_involved": "",
|
||||
"key_items": "",
|
||||
"scene_location": "",
|
||||
"time_constraint": ""
|
||||
},
|
||||
"choose_configs": {
|
||||
"prompt_draft_llm": "DeepSeek V3",
|
||||
"chapter_outline_llm": "DeepSeek V3",
|
||||
"architecture_llm": "Gemini 2.5 Pro",
|
||||
"final_chapter_llm": "GPT 5",
|
||||
"consistency_review_llm": "DeepSeek V3"
|
||||
}
|
||||
}
|
||||
@@ -16,6 +16,13 @@ from prompt_definitions import (
|
||||
plot_architecture_prompt,
|
||||
create_character_state_prompt
|
||||
)
|
||||
logging.basicConfig(
|
||||
filename='app.log', # 日志文件名
|
||||
filemode='a', # 追加模式('w' 会覆盖)
|
||||
level=logging.INFO, # 记录 INFO 及以上级别的日志
|
||||
format='%(asctime)s - %(name)s - %(levelname)s - %(message)s',
|
||||
datefmt='%Y-%m-%d %H:%M:%S'
|
||||
)
|
||||
from utils import clear_file_content, save_string_to_txt
|
||||
|
||||
def load_partial_architecture_data(filepath: str) -> dict:
|
||||
|
||||
@@ -10,7 +10,13 @@ from novel_generator.common import invoke_with_cleaning
|
||||
from llm_adapters import create_llm_adapter
|
||||
from prompt_definitions import chapter_blueprint_prompt, chunked_chapter_blueprint_prompt
|
||||
from utils import read_file, clear_file_content, save_string_to_txt
|
||||
|
||||
logging.basicConfig(
|
||||
filename='app.log', # 日志文件名
|
||||
filemode='a', # 追加模式('w' 会覆盖)
|
||||
level=logging.INFO, # 记录 INFO 及以上级别的日志
|
||||
format='%(asctime)s - %(name)s - %(levelname)s - %(message)s',
|
||||
datefmt='%Y-%m-%d %H:%M:%S'
|
||||
)
|
||||
def compute_chunk_size(number_of_chapters: int, max_tokens: int) -> int:
|
||||
"""
|
||||
基于“每章约100 tokens”的粗略估算,
|
||||
@@ -18,7 +24,7 @@ def compute_chunk_size(number_of_chapters: int, max_tokens: int) -> int:
|
||||
chunk_size = (floor(max_tokens/100/10)*10) - 10
|
||||
并确保 chunk_size 不会小于1或大于实际章节数。
|
||||
"""
|
||||
tokens_per_chapter = 100.0
|
||||
tokens_per_chapter = 200.0
|
||||
ratio = max_tokens / tokens_per_chapter
|
||||
ratio_rounded_to_10 = int(ratio // 10) * 10
|
||||
chunk_size = ratio_rounded_to_10 - 10
|
||||
|
||||
@@ -22,6 +22,13 @@ from novel_generator.vectorstore_utils import (
|
||||
get_relevant_context_from_vector_store,
|
||||
load_vector_store # 添加导入
|
||||
)
|
||||
logging.basicConfig(
|
||||
filename='app.log', # 日志文件名
|
||||
filemode='a', # 追加模式('w' 会覆盖)
|
||||
level=logging.INFO, # 记录 INFO 及以上级别的日志
|
||||
format='%(asctime)s - %(name)s - %(levelname)s - %(message)s',
|
||||
datefmt='%Y-%m-%d %H:%M:%S'
|
||||
)
|
||||
|
||||
def get_last_n_chapters_text(chapters_dir: str, current_chapter_num: int, n: int = 3) -> list:
|
||||
"""
|
||||
|
||||
@@ -7,7 +7,13 @@ import logging
|
||||
import re
|
||||
import time
|
||||
import traceback
|
||||
|
||||
logging.basicConfig(
|
||||
filename='app.log', # 日志文件名
|
||||
filemode='a', # 追加模式('w' 会覆盖)
|
||||
level=logging.INFO, # 记录 INFO 及以上级别的日志
|
||||
format='%(asctime)s - %(name)s - %(levelname)s - %(message)s',
|
||||
datefmt='%Y-%m-%d %H:%M:%S'
|
||||
)
|
||||
def call_with_retry(func, max_retries=3, sleep_time=2, fallback_return=None, **kwargs):
|
||||
"""
|
||||
通用的重试机制封装。
|
||||
|
||||
@@ -11,7 +11,13 @@ from prompt_definitions import summary_prompt, update_character_state_prompt
|
||||
from novel_generator.common import invoke_with_cleaning
|
||||
from utils import read_file, clear_file_content, save_string_to_txt
|
||||
from novel_generator.vectorstore_utils import update_vector_store
|
||||
|
||||
logging.basicConfig(
|
||||
filename='app.log', # 日志文件名
|
||||
filemode='a', # 追加模式('w' 会覆盖)
|
||||
level=logging.INFO, # 记录 INFO 及以上级别的日志
|
||||
format='%(asctime)s - %(name)s - %(levelname)s - %(message)s',
|
||||
datefmt='%Y-%m-%d %H:%M:%S'
|
||||
)
|
||||
def finalize_chapter(
|
||||
novel_number: int,
|
||||
word_number: int,
|
||||
@@ -111,7 +117,7 @@ def enrich_chapter_text(
|
||||
max_tokens=max_tokens,
|
||||
timeout=timeout
|
||||
)
|
||||
prompt = f"""以下章节文本较短,请在保持剧情连贯的前提下进行扩写,使其更充实,接近 {word_number} 字左右:
|
||||
prompt = f"""以下章节文本较短,请在保持剧情连贯的前提下进行扩写,使其更充实,接近 {word_number} 字左右,仅给出最终文本,不要解释任何内容。:
|
||||
原内容:
|
||||
{chapter_text}
|
||||
"""
|
||||
|
||||
@@ -16,11 +16,17 @@ from langchain.docstore.document import Document
|
||||
# 禁用特定的Torch警告
|
||||
warnings.filterwarnings('ignore', message='.*Torch was not compiled with flash attention.*')
|
||||
os.environ["TOKENIZERS_PARALLELISM"] = "false"
|
||||
|
||||
logging.basicConfig(
|
||||
filename='app.log', # 日志文件名
|
||||
filemode='a', # 追加模式('w' 会覆盖)
|
||||
level=logging.INFO, # 记录 INFO 及以上级别的日志
|
||||
format='%(asctime)s - %(name)s - %(levelname)s - %(message)s',
|
||||
datefmt='%Y-%m-%d %H:%M:%S'
|
||||
)
|
||||
def advanced_split_content(content: str, similarity_threshold: float = 0.7, max_length: int = 500) -> list:
|
||||
"""使用基本分段策略"""
|
||||
nltk.download('punkt', quiet=True)
|
||||
nltk.download('punkt_tab', quiet=True)
|
||||
# nltk.download('punkt', quiet=True)
|
||||
# nltk.download('punkt_tab', quiet=True)
|
||||
sentences = nltk.sent_tokenize(content)
|
||||
if not sentences:
|
||||
return []
|
||||
|
||||
@@ -13,7 +13,13 @@ import ssl
|
||||
import requests
|
||||
import warnings
|
||||
from langchain_chroma import Chroma
|
||||
|
||||
logging.basicConfig(
|
||||
filename='app.log', # 日志文件名
|
||||
filemode='a', # 追加模式('w' 会覆盖)
|
||||
level=logging.INFO, # 记录 INFO 及以上级别的日志
|
||||
format='%(asctime)s - %(name)s - %(levelname)s - %(message)s',
|
||||
datefmt='%Y-%m-%d %H:%M:%S'
|
||||
)
|
||||
# 禁用特定的Torch警告
|
||||
warnings.filterwarnings('ignore', message='.*Torch was not compiled with flash attention.*')
|
||||
os.environ["TOKENIZERS_PARALLELISM"] = "false" # 禁用tokenizer并行警告
|
||||
@@ -146,8 +152,8 @@ def split_text_for_vectorstore(chapter_text: str, max_length: int = 500, similar
|
||||
if not chapter_text.strip():
|
||||
return []
|
||||
|
||||
nltk.download('punkt', quiet=True)
|
||||
nltk.download('punkt_tab', quiet=True)
|
||||
# nltk.download('punkt', quiet=True)
|
||||
# nltk.download('punkt_tab', quiet=True)
|
||||
sentences = nltk.sent_tokenize(chapter_text)
|
||||
if not sentences:
|
||||
return []
|
||||
|
||||
+2
-1
@@ -33,5 +33,6 @@ tooltips = {
|
||||
"characters_involved": "本章需要重点描写或影响剧情的角色名单。",
|
||||
"key_items": "在本章中出现的重要道具、线索或物品。",
|
||||
"scene_location": "本章主要发生的地点或场景描述。",
|
||||
"time_constraint": "本章剧情中涉及的时间压力或时限设置。"
|
||||
"time_constraint": "本章剧情中涉及的时间压力或时限设置。",
|
||||
"interface_config": "选择你要使用的AI接口配置。"
|
||||
}
|
||||
|
||||
+491
-99
@@ -1,6 +1,8 @@
|
||||
# ui/config_tab.py
|
||||
# -*- coding: utf-8 -*-
|
||||
from tkinter import messagebox
|
||||
import uuid
|
||||
import datetime
|
||||
|
||||
import customtkinter as ctk
|
||||
|
||||
@@ -41,123 +43,430 @@ def build_config_tabview(self):
|
||||
|
||||
self.ai_config_tab = self.config_tabview.add("LLM Model settings")
|
||||
self.embeddings_config_tab = self.config_tabview.add("Embedding settings")
|
||||
self.config_choose = self.config_tabview.add("Config choose")
|
||||
|
||||
|
||||
build_ai_config_tab(self)
|
||||
build_embeddings_config_tab(self)
|
||||
|
||||
# 底部的"保存配置"和"加载配置"按钮
|
||||
self.btn_frame_config = ctk.CTkFrame(self.config_frame)
|
||||
self.btn_frame_config.grid(row=1, column=0, padx=5, pady=5, sticky="ew")
|
||||
self.btn_frame_config.columnconfigure(0, weight=1)
|
||||
self.btn_frame_config.columnconfigure(1, weight=1)
|
||||
|
||||
save_config_btn = ctk.CTkButton(self.btn_frame_config, text="保存当前选择接口配置到文件", command=self.save_config_btn, font=("Microsoft YaHei", 12))
|
||||
save_config_btn.grid(row=0, column=0, padx=5, pady=5, sticky="ew")
|
||||
|
||||
load_config_btn = ctk.CTkButton(self.btn_frame_config, text="加载当前选择接口配置到程序", command=self.load_config_btn, font=("Microsoft YaHei", 12))
|
||||
load_config_btn.grid(row=0, column=1, padx=5, pady=5, sticky="ew")
|
||||
build_config_choose_tab(self)
|
||||
|
||||
def build_ai_config_tab(self):
|
||||
def on_interface_format_changed(new_value):
|
||||
self.interface_format_var.set(new_value)
|
||||
config_data = load_config(self.config_file)
|
||||
if config_data:
|
||||
config_data["last_interface_format"] = new_value
|
||||
save_config(config_data, self.config_file)
|
||||
if self.loaded_config and "llm_configs" in self.loaded_config and new_value in self.loaded_config["llm_configs"]:
|
||||
llm_conf = self.loaded_config["llm_configs"][new_value]
|
||||
self.api_key_var.set(llm_conf.get("api_key", ""))
|
||||
self.base_url_var.set(llm_conf.get("base_url", self.base_url_var.get()))
|
||||
self.model_name_var.set(llm_conf.get("model_name", ""))
|
||||
self.temperature_var.set(llm_conf.get("temperature", 0.7))
|
||||
self.max_tokens_var.set(llm_conf.get("max_tokens", 8192))
|
||||
self.timeout_var.set(llm_conf.get("timeout", 600))
|
||||
else:
|
||||
if new_value == "Ollama":
|
||||
self.base_url_var.set("http://localhost:11434/v1")
|
||||
elif new_value == "ML Studio":
|
||||
self.base_url_var.set("http://localhost:1234/v1")
|
||||
elif new_value == "OpenAI":
|
||||
self.base_url_var.set("https://api.openai.com/v1")
|
||||
self.model_name_var.set("gpt-4o-mini")
|
||||
elif new_value == "Azure OpenAI":
|
||||
self.base_url_var.set("https://[az].openai.azure.com/openai/deployments/[model]/chat/completions?api-version=2024-08-01-preview")
|
||||
elif new_value == "DeepSeek":
|
||||
self.base_url_var.set("https://api.deepseek.com/v1")
|
||||
self.model_name_var.set("deepseek-chat")
|
||||
elif new_value == "Gemini":
|
||||
self.base_url_var.set("")
|
||||
elif new_value == "Azure AI":
|
||||
self.base_url_var.set("https://<your-endpoint>.services.ai.azure.com/models/chat/completions?api-version=2024-05-01-preview")
|
||||
elif new_value == "阿里云百炼":
|
||||
self.base_url_var.set("https://dashscope.aliyuncs.com/compatible-mode/v1")
|
||||
self.model_name_var.set("qwen-plus")
|
||||
elif new_value == "硅基流动":
|
||||
self.base_url_var.set("https://api.siliconflow.cn/v1")
|
||||
self.model_name_var.set("deepseek-ai/DeepSeek-V3")
|
||||
elif new_value == "Grok":
|
||||
self.base_url_var.set("https://api.x.ai/v1")
|
||||
self.model_name_var.set("grok-3")
|
||||
def refresh_config_dropdown():
|
||||
"""刷新配置下拉菜单"""
|
||||
config_names = list(self.loaded_config.get("llm_configs", {}).keys())
|
||||
interface_config_dropdown.configure(values=config_names)
|
||||
if config_names and self.interface_config_var.get() not in config_names:
|
||||
self.interface_config_var.set(config_names[0])
|
||||
|
||||
for i in range(7):
|
||||
def on_config_selected(new_value):
|
||||
"""当选择不同配置时的回调"""
|
||||
if new_value in self.loaded_config.get("llm_configs", {}):
|
||||
config = self.loaded_config["llm_configs"][new_value]
|
||||
# 更新所有UI变量
|
||||
self.api_key_var.set(config.get("api_key", ""))
|
||||
self.base_url_var.set(config.get("base_url", ""))
|
||||
self.model_name_var.set(config.get("model_name", ""))
|
||||
self.temperature_var.set(float(config.get("temperature", 0.7)))
|
||||
self.max_tokens_var.set(int(config.get("max_tokens", 8192)))
|
||||
self.timeout_var.set(int(config.get("timeout", 600)))
|
||||
self.interface_format_var.set(config.get("interface_format", "OpenAI"))
|
||||
|
||||
# 更新显示标签
|
||||
self.temp_value_label.configure(text=f"{float(config.get('temperature', 0.7)):.2f}")
|
||||
self.max_tokens_value_label.configure(text=str(int(config.get('max_tokens', 8192))))
|
||||
self.timeout_value_label.configure(text=str(int(config.get('timeout', 600))))
|
||||
|
||||
def add_new_config():
|
||||
"""添加新配置 - 弹出对话框让用户输入名称"""
|
||||
dialog = ctk.CTkInputDialog(
|
||||
text="请输入新配置名称:",
|
||||
title="新增配置"
|
||||
)
|
||||
new_name = dialog.get_input()
|
||||
|
||||
if not new_name:
|
||||
return
|
||||
|
||||
new_name = new_name.strip()
|
||||
|
||||
if new_name in self.loaded_config.get("llm_configs", {}):
|
||||
messagebox.showerror("错误", f"配置名称 '{new_name}' 已存在!")
|
||||
return
|
||||
|
||||
if "llm_configs" not in self.loaded_config:
|
||||
self.loaded_config["llm_configs"] = {}
|
||||
|
||||
self.loaded_config["llm_configs"][new_name] = {
|
||||
"id": str(uuid.uuid4()),
|
||||
"api_key": "",
|
||||
"base_url": "",
|
||||
"model_name": "",
|
||||
"temperature": 0.7,
|
||||
"max_tokens": 8192,
|
||||
"timeout": 600,
|
||||
"interface_format": "OpenAI",
|
||||
"created_at": datetime.datetime.now().isoformat()
|
||||
}
|
||||
|
||||
refresh_config_dropdown()
|
||||
self.interface_config_var.set(new_name)
|
||||
messagebox.showinfo("提示", f"已成功创建新配置: {new_name}")
|
||||
|
||||
def delete_current_config():
|
||||
"""删除当前选中的配置并保存到JSON文件"""
|
||||
selected_config = self.interface_config_var.get()
|
||||
if selected_config in self.loaded_config.get("llm_configs", {}):
|
||||
if len(self.loaded_config["llm_configs"]) <= 1:
|
||||
messagebox.showerror("错误", "至少需要保留一个配置!")
|
||||
return
|
||||
|
||||
confirm = messagebox.askyesno(
|
||||
"确认删除",
|
||||
f"确定要删除配置 '{selected_config}' 吗?\n此操作不可撤销!"
|
||||
)
|
||||
if not confirm:
|
||||
return
|
||||
|
||||
del self.loaded_config["llm_configs"][selected_config]
|
||||
refresh_config_dropdown()
|
||||
|
||||
# 保存到JSON文件
|
||||
try:
|
||||
save_config(self.loaded_config, self.config_file)
|
||||
messagebox.showinfo("提示", f"已删除配置: {selected_config},并已更新配置文件")
|
||||
except Exception as e:
|
||||
messagebox.showerror("错误", f"保存配置文件失败: {str(e)}")
|
||||
else:
|
||||
messagebox.showerror("错误", "未找到选中的配置!")
|
||||
|
||||
def save_current_config():
|
||||
"""保存当前配置的修改到JSON文件"""
|
||||
config_name = self.interface_config_var.get()
|
||||
if config_name not in self.loaded_config.get("llm_configs", {}):
|
||||
messagebox.showerror("错误", "配置不存在!")
|
||||
return
|
||||
|
||||
config = self.loaded_config["llm_configs"][config_name]
|
||||
config.update({
|
||||
"api_key": self.api_key_var.get(),
|
||||
"base_url": self.base_url_var.get(),
|
||||
"model_name": self.model_name_var.get(),
|
||||
"temperature": float(self.temperature_var.get()),
|
||||
"max_tokens": int(self.max_tokens_var.get()),
|
||||
"timeout": int(self.timeout_var.get()),
|
||||
"interface_format": self.interface_format_var.get(),
|
||||
"updated_at": datetime.datetime.now().isoformat()
|
||||
})
|
||||
|
||||
# 如果修改了配置名称
|
||||
new_name = self.interface_config_var.get()
|
||||
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)
|
||||
messagebox.showinfo("提示", f"配置 {new_name} 已保存并持久化到文件")
|
||||
except Exception as e:
|
||||
messagebox.showerror("错误", f"保存配置文件失败: {str(e)}")
|
||||
|
||||
def rename_current_config():
|
||||
"""重命名当前配置"""
|
||||
old_name = self.interface_config_var.get()
|
||||
if old_name not in self.loaded_config.get("llm_configs", {}):
|
||||
messagebox.showerror("错误", "当前配置不存在!")
|
||||
return
|
||||
|
||||
dialog = ctk.CTkInputDialog(
|
||||
text=f"请输入新的配置名称 (原名称: {old_name}):",
|
||||
title="重命名配置"
|
||||
)
|
||||
new_name = dialog.get_input()
|
||||
|
||||
if not new_name:
|
||||
return
|
||||
|
||||
new_name = new_name.strip()
|
||||
|
||||
if new_name == old_name:
|
||||
return
|
||||
|
||||
if new_name in self.loaded_config.get("llm_configs", {}):
|
||||
messagebox.showerror("错误", f"配置名称 '{new_name}' 已存在!")
|
||||
return
|
||||
|
||||
# 更新配置名称
|
||||
self.loaded_config["llm_configs"][new_name] = self.loaded_config["llm_configs"].pop(old_name)
|
||||
self.interface_config_var.set(new_name)
|
||||
refresh_config_dropdown()
|
||||
|
||||
messagebox.showinfo("提示", f"配置已从 '{old_name}' 重命名为 '{new_name}'")
|
||||
|
||||
# 初始化UI布局
|
||||
for i in range(10):
|
||||
self.ai_config_tab.grid_rowconfigure(i, weight=0)
|
||||
self.ai_config_tab.grid_columnconfigure(0, weight=0)
|
||||
self.ai_config_tab.grid_columnconfigure(1, weight=1)
|
||||
self.ai_config_tab.grid_columnconfigure(2, weight=0)
|
||||
|
||||
# 配置选择控件
|
||||
create_label_with_help(self, self.ai_config_tab, "当前配置", "interface_config", 0, 0)
|
||||
config_names = list(self.loaded_config.get("llm_configs", {}).keys())
|
||||
if not config_names:
|
||||
self.loaded_config["llm_configs"] = {
|
||||
"默认配置": {
|
||||
"id": str(uuid.uuid4()),
|
||||
"api_key": "",
|
||||
"base_url": "https://api.openai.com/v1",
|
||||
"model_name": "gpt-4",
|
||||
"temperature": 0.7,
|
||||
"max_tokens": 8192,
|
||||
"timeout": 600,
|
||||
"interface_format": "OpenAI",
|
||||
"created_at": datetime.datetime.now().isoformat()
|
||||
}
|
||||
}
|
||||
config_names = ["默认配置"]
|
||||
|
||||
self.interface_config_var = ctk.StringVar(value=config_names[0])
|
||||
|
||||
interface_config_dropdown = ctk.CTkOptionMenu(
|
||||
self.ai_config_tab,
|
||||
values=config_names,
|
||||
variable=self.interface_config_var,
|
||||
command=on_config_selected,
|
||||
font=("Microsoft YaHei", 12)
|
||||
)
|
||||
interface_config_dropdown.grid(row=0, column=1, columnspan=2, padx=5, pady=5, sticky="nsew")
|
||||
|
||||
# 配置管理按钮组
|
||||
btn_frame = ctk.CTkFrame(self.ai_config_tab)
|
||||
btn_frame.grid(row=1, column=0, columnspan=3, padx=5, pady=5, sticky="ew")
|
||||
btn_frame.columnconfigure(0, weight=1)
|
||||
btn_frame.columnconfigure(1, weight=1)
|
||||
btn_frame.columnconfigure(2, weight=1)
|
||||
btn_frame.columnconfigure(3, weight=1)
|
||||
|
||||
add_btn = ctk.CTkButton(
|
||||
btn_frame,
|
||||
text="➕ 新增",
|
||||
command=add_new_config,
|
||||
font=("Microsoft YaHei", 12),
|
||||
fg_color="#2E8B57",
|
||||
width=80
|
||||
)
|
||||
add_btn.grid(row=0, column=0, padx=2, pady=2, sticky="ew")
|
||||
|
||||
rename_btn = ctk.CTkButton(
|
||||
btn_frame,
|
||||
text="✏️ 重命名",
|
||||
command=rename_current_config,
|
||||
font=("Microsoft YaHei", 12),
|
||||
fg_color="#DAA520",
|
||||
width=80
|
||||
)
|
||||
rename_btn.grid(row=0, column=1, padx=2, pady=2, sticky="ew")
|
||||
|
||||
del_btn = ctk.CTkButton(
|
||||
btn_frame,
|
||||
text="🗑️ 删除",
|
||||
command=delete_current_config,
|
||||
font=("Microsoft YaHei", 12),
|
||||
fg_color="#8B0000",
|
||||
width=80
|
||||
)
|
||||
del_btn.grid(row=0, column=2, padx=2, pady=2, sticky="ew")
|
||||
|
||||
save_btn = ctk.CTkButton(
|
||||
btn_frame,
|
||||
text="💾 保存",
|
||||
command=save_current_config,
|
||||
font=("Microsoft YaHei", 12),
|
||||
fg_color="#1E90FF",
|
||||
width=80
|
||||
)
|
||||
save_btn.grid(row=0, column=3, padx=2, pady=2, sticky="ew")
|
||||
|
||||
# 配置参数控件
|
||||
row_start = 2
|
||||
# 1) API Key
|
||||
create_label_with_help(self, parent=self.ai_config_tab, label_text="LLM API Key:", tooltip_key="api_key", row=0, column=0, font=("Microsoft YaHei", 12))
|
||||
api_key_entry = ctk.CTkEntry(self.ai_config_tab, textvariable=self.api_key_var, font=("Microsoft YaHei", 12),show="*")
|
||||
api_key_entry.grid(row=0, column=1, padx=5, pady=5, columnspan=2, sticky="nsew")
|
||||
|
||||
create_label_with_help(self, self.ai_config_tab, "API Key:", "api_key", row_start, 0)
|
||||
self.api_key_var = ctk.StringVar(value="")
|
||||
api_key_entry = ctk.CTkEntry(
|
||||
self.ai_config_tab,
|
||||
textvariable=self.api_key_var,
|
||||
font=("Microsoft YaHei", 12),
|
||||
show="*"
|
||||
)
|
||||
api_key_entry.grid(row=row_start, column=1, columnspan=2, padx=5, pady=5, sticky="nsew")
|
||||
|
||||
# 2) Base URL
|
||||
create_label_with_help(self, parent=self.ai_config_tab, label_text="LLM Base URL:", tooltip_key="base_url", row=1, column=0, font=("Microsoft YaHei", 12))
|
||||
base_url_entry = ctk.CTkEntry(self.ai_config_tab, textvariable=self.base_url_var, font=("Microsoft YaHei", 12))
|
||||
base_url_entry.grid(row=1, column=1, padx=5, pady=5, columnspan=2, sticky="nsew")
|
||||
|
||||
create_label_with_help(self, self.ai_config_tab, "Base URL:", "base_url", row_start+1, 0)
|
||||
self.base_url_var = ctk.StringVar(value="")
|
||||
base_url_entry = ctk.CTkEntry(
|
||||
self.ai_config_tab,
|
||||
textvariable=self.base_url_var,
|
||||
font=("Microsoft YaHei", 12)
|
||||
)
|
||||
base_url_entry.grid(row=row_start+1, column=1, columnspan=2, padx=5, pady=5, sticky="nsew")
|
||||
|
||||
# 3) 接口格式
|
||||
create_label_with_help(self, parent=self.ai_config_tab, label_text="LLM 接口格式:", tooltip_key="interface_format", row=2, column=0, font=("Microsoft YaHei", 12))
|
||||
# 在接口选项列表中添加 "Grok"
|
||||
interface_options = ["DeepSeek", "阿里云百炼", "OpenAI", "Azure OpenAI", "Azure AI", "Ollama", "ML Studio", "Gemini", "火山引擎", "硅基流动", "Grok"]
|
||||
interface_dropdown = ctk.CTkOptionMenu(self.ai_config_tab, values=interface_options, variable=self.interface_format_var, command=on_interface_format_changed, font=("Microsoft YaHei", 12))
|
||||
interface_dropdown.grid(row=2, column=1, padx=5, pady=5, columnspan=2, sticky="nsew")
|
||||
|
||||
create_label_with_help(self, self.ai_config_tab, "接口格式:", "interface_format", row_start+2, 0)
|
||||
self.interface_format_var = ctk.StringVar(value="OpenAI")
|
||||
interface_options = ["OpenAI", "Azure OpenAI", "Ollama", "DeepSeek", "Gemini", "ML Studio"]
|
||||
interface_dropdown = ctk.CTkOptionMenu(
|
||||
self.ai_config_tab,
|
||||
values=interface_options,
|
||||
variable=self.interface_format_var,
|
||||
font=("Microsoft YaHei", 12)
|
||||
)
|
||||
interface_dropdown.grid(row=row_start+2, column=1, columnspan=2, padx=5, pady=5, sticky="nsew")
|
||||
|
||||
# 4) Model Name
|
||||
create_label_with_help(self, parent=self.ai_config_tab, label_text="Model Name:", tooltip_key="model_name", row=3, column=0, font=("Microsoft YaHei", 12))
|
||||
model_name_entry = ctk.CTkEntry(self.ai_config_tab, textvariable=self.model_name_var, font=("Microsoft YaHei", 12))
|
||||
model_name_entry.grid(row=3, column=1, padx=5, pady=5, columnspan=2, sticky="nsew")
|
||||
|
||||
create_label_with_help(self, self.ai_config_tab, "模型名称:", "model_name", row_start+3, 0)
|
||||
self.model_name_var = ctk.StringVar(value="")
|
||||
model_name_entry = ctk.CTkEntry(
|
||||
self.ai_config_tab,
|
||||
textvariable=self.model_name_var,
|
||||
font=("Microsoft YaHei", 12)
|
||||
)
|
||||
model_name_entry.grid(row=row_start+3, column=1, columnspan=2, padx=5, pady=5, sticky="nsew")
|
||||
|
||||
# 5) Temperature
|
||||
create_label_with_help(self, parent=self.ai_config_tab, label_text="Temperature:", tooltip_key="temperature", row=4, column=0, font=("Microsoft YaHei", 12))
|
||||
create_label_with_help(self, self.ai_config_tab, "Temperature:", "temperature", row_start+4, 0)
|
||||
self.temperature_var = ctk.DoubleVar(value=0.7)
|
||||
def update_temp_label(value):
|
||||
self.temp_value_label.configure(text=f"{float(value):.2f}")
|
||||
temp_scale = ctk.CTkSlider(self.ai_config_tab, from_=0.0, to=2.0, number_of_steps=200, command=update_temp_label, variable=self.temperature_var)
|
||||
temp_scale.grid(row=4, column=1, padx=5, pady=5, sticky="we")
|
||||
self.temp_value_label = ctk.CTkLabel(self.ai_config_tab, text=f"{self.temperature_var.get():.2f}", font=("Microsoft YaHei", 12))
|
||||
self.temp_value_label.grid(row=4, column=2, padx=5, pady=5, sticky="w")
|
||||
|
||||
temp_scale = ctk.CTkSlider(
|
||||
self.ai_config_tab,
|
||||
from_=0.0,
|
||||
to=2.0,
|
||||
number_of_steps=200,
|
||||
command=update_temp_label,
|
||||
variable=self.temperature_var
|
||||
)
|
||||
temp_scale.grid(row=row_start+4, column=1, padx=5, pady=5, sticky="we")
|
||||
self.temp_value_label = ctk.CTkLabel(
|
||||
self.ai_config_tab,
|
||||
text=f"{self.temperature_var.get():.2f}",
|
||||
font=("Microsoft YaHei", 12)
|
||||
)
|
||||
self.temp_value_label.grid(row=row_start+4, column=2, padx=5, pady=5, sticky="w")
|
||||
|
||||
# 6) Max Tokens
|
||||
create_label_with_help(self, parent=self.ai_config_tab, label_text="Max Tokens:", tooltip_key="max_tokens", row=5, column=0, font=("Microsoft YaHei", 12))
|
||||
create_label_with_help(self, self.ai_config_tab, "Max Tokens:", "max_tokens", row_start+5, 0)
|
||||
self.max_tokens_var = ctk.IntVar(value=8192)
|
||||
def update_max_tokens_label(value):
|
||||
self.max_tokens_value_label.configure(text=str(int(float(value))))
|
||||
max_tokens_slider = ctk.CTkSlider(self.ai_config_tab, from_=0, to=102400, number_of_steps=100, command=update_max_tokens_label, variable=self.max_tokens_var)
|
||||
max_tokens_slider.grid(row=5, column=1, padx=5, pady=5, sticky="we")
|
||||
self.max_tokens_value_label = ctk.CTkLabel(self.ai_config_tab, text=str(self.max_tokens_var.get()), font=("Microsoft YaHei", 12))
|
||||
self.max_tokens_value_label.grid(row=5, column=2, padx=5, pady=5, sticky="w")
|
||||
|
||||
# 7) Timeout (sec)
|
||||
create_label_with_help(self, parent=self.ai_config_tab, label_text="Timeout (sec):", tooltip_key="timeout", row=6, column=0, font=("Microsoft YaHei", 12))
|
||||
max_tokens_slider = ctk.CTkSlider(
|
||||
self.ai_config_tab,
|
||||
from_=0,
|
||||
to=102400,
|
||||
number_of_steps=100,
|
||||
command=update_max_tokens_label,
|
||||
variable=self.max_tokens_var
|
||||
)
|
||||
max_tokens_slider.grid(row=row_start+5, column=1, padx=5, pady=5, sticky="we")
|
||||
self.max_tokens_value_label = ctk.CTkLabel(
|
||||
self.ai_config_tab,
|
||||
text=str(self.max_tokens_var.get()),
|
||||
font=("Microsoft YaHei", 12)
|
||||
)
|
||||
self.max_tokens_value_label.grid(row=row_start+5, column=2, padx=5, pady=5, sticky="w")
|
||||
|
||||
# 7) Timeout
|
||||
create_label_with_help(self, self.ai_config_tab, "Timeout (sec):", "timeout", row_start+6, 0)
|
||||
self.timeout_var = ctk.IntVar(value=600)
|
||||
def update_timeout_label(value):
|
||||
integer_val = int(float(value))
|
||||
self.timeout_value_label.configure(text=str(integer_val))
|
||||
timeout_slider = ctk.CTkSlider(self.ai_config_tab, from_=0, to=3600, number_of_steps=3600, command=update_timeout_label, variable=self.timeout_var)
|
||||
timeout_slider.grid(row=6, column=1, padx=5, pady=5, sticky="we")
|
||||
self.timeout_value_label = ctk.CTkLabel(self.ai_config_tab, text=str(self.timeout_var.get()), font=("Microsoft YaHei", 12))
|
||||
self.timeout_value_label.grid(row=6, column=2, padx=5, pady=5, sticky="w")
|
||||
self.timeout_value_label.configure(text=str(int(float(value))))
|
||||
timeout_slider = ctk.CTkSlider(
|
||||
self.ai_config_tab,
|
||||
from_=0,
|
||||
to=3600,
|
||||
number_of_steps=3600,
|
||||
command=update_timeout_label,
|
||||
variable=self.timeout_var
|
||||
)
|
||||
timeout_slider.grid(row=row_start+6, column=1, padx=5, pady=5, sticky="we")
|
||||
self.timeout_value_label = ctk.CTkLabel(
|
||||
self.ai_config_tab,
|
||||
text=str(self.timeout_var.get()),
|
||||
font=("Microsoft YaHei", 12)
|
||||
)
|
||||
self.timeout_value_label.grid(row=row_start+6, column=2, padx=5, pady=5, sticky="w")
|
||||
|
||||
# 测试按钮
|
||||
test_btn = ctk.CTkButton(
|
||||
self.ai_config_tab,
|
||||
text="测试配置",
|
||||
command=self.test_llm_config,
|
||||
font=("Microsoft YaHei", 12)
|
||||
)
|
||||
test_btn.grid(row=row_start+7, column=0, columnspan=3, padx=5, pady=5, sticky="ew")
|
||||
|
||||
# 添加测试按钮
|
||||
test_btn = ctk.CTkButton(self.ai_config_tab, text="测试配置", command=self.test_llm_config, font=("Microsoft YaHei", 12))
|
||||
test_btn.grid(row=7, column=0, columnspan=3, padx=5, pady=5, sticky="ew")
|
||||
# 初始化当前配置
|
||||
on_config_selected(config_names[0])
|
||||
|
||||
|
||||
# 初始化UI布局
|
||||
for i in range(10): # 增加一行给按钮组
|
||||
self.ai_config_tab.grid_rowconfigure(i, weight=0)
|
||||
self.ai_config_tab.grid_columnconfigure(0, weight=0)
|
||||
self.ai_config_tab.grid_columnconfigure(1, weight=1)
|
||||
self.ai_config_tab.grid_columnconfigure(2, weight=0)
|
||||
|
||||
# 配置选择控件
|
||||
create_label_with_help(self, self.ai_config_tab, "当前配置", "interface_config", 0, 0)
|
||||
config_names = list(self.loaded_config.get("llm_configs", {}).keys())
|
||||
if not config_names: # 如果没有配置,创建一个默认配置
|
||||
self.loaded_config["llm_configs"] = {
|
||||
"默认配置": {
|
||||
"id": str(uuid.uuid4()),
|
||||
"api_key": "",
|
||||
"base_url": "https://api.openai.com/v1",
|
||||
"model_name": "gpt-4",
|
||||
"temperature": 0.7,
|
||||
"max_tokens": 8192,
|
||||
"timeout": 600,
|
||||
"interface_format": "OpenAI",
|
||||
"created_at": datetime.datetime.now().isoformat()
|
||||
}
|
||||
}
|
||||
config_names = ["默认配置"]
|
||||
|
||||
interface_config_dropdown = ctk.CTkOptionMenu(
|
||||
self.ai_config_tab,
|
||||
values=config_names,
|
||||
variable=self.interface_config_var,
|
||||
command=on_config_selected,
|
||||
font=("Microsoft YaHei", 12)
|
||||
)
|
||||
interface_config_dropdown.grid(row=0, column=1, columnspan=2, padx=5, pady=5, sticky="nsew")
|
||||
|
||||
def build_embeddings_config_tab(self):
|
||||
def on_embedding_interface_changed(new_value):
|
||||
@@ -199,11 +508,11 @@ 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 接口格式
|
||||
create_label_with_help(self, parent=self.embeddings_config_tab, label_text="Embedding 接口格式:", tooltip_key="embedding_interface_format", row=1, column=0, font=("Microsoft YaHei", 12))
|
||||
create_label_with_help(self, parent=self.embeddings_config_tab, label_text="Embedding 接口格式:", tooltip_key="embedding_intexrface_format", row=1, column=0, font=("Microsoft YaHei", 12))
|
||||
|
||||
emb_interface_options = ["DeepSeek", "OpenAI", "Azure OpenAI", "Gemini", "Ollama", "ML Studio","SiliconFlow"]
|
||||
|
||||
@@ -229,6 +538,81 @@ def build_embeddings_config_tab(self):
|
||||
test_btn = ctk.CTkButton(self.embeddings_config_tab, text="测试配置", command=self.test_embedding_config, font=("Microsoft YaHei", 12))
|
||||
test_btn.grid(row=5, column=0, columnspan=2, padx=5, pady=5, sticky="ew")
|
||||
|
||||
def build_config_choose_tab(self):
|
||||
|
||||
|
||||
self.config_choose.grid_rowconfigure(0, weight=0)
|
||||
self.config_choose.grid_columnconfigure(0, weight=0)
|
||||
self.config_choose.grid_columnconfigure(1, weight=1)
|
||||
config_choose_options = list(self.loaded_config.get("llm_configs", {}).keys())
|
||||
create_label_with_help(self, parent=self.config_choose, label_text="生成架构所用大模型", tooltip_key="architecture_llm_config", row=0, column=0, font=("Microsoft YaHei", 12))
|
||||
architecture_dropdown = ctk.CTkOptionMenu(self.config_choose, values=config_choose_options, variable=self.architecture_llm_var, font=("Microsoft YaHei", 12))
|
||||
architecture_dropdown.grid(row=0, column=1, padx=5, pady=5, sticky="nsew")
|
||||
|
||||
create_label_with_help(self, parent=self.config_choose, label_text="生成大目录所用大模型", tooltip_key="chapter_outline_llm_config", row=1, column=0, font=("Microsoft YaHei", 12))
|
||||
chapter_outline_dropdown = ctk.CTkOptionMenu(self.config_choose, values=config_choose_options, variable=self.chapter_outline_llm_var, font=("Microsoft YaHei", 12))
|
||||
chapter_outline_dropdown.grid(row=1, column=1, padx=5, pady=5, sticky="nsew")
|
||||
|
||||
create_label_with_help(self, parent=self.config_choose, label_text="生成草稿所用大模型", tooltip_key="prompt_draft_llm_config", row=2, column=0, font=("Microsoft YaHei", 12))
|
||||
prompt_draft_dropdown = ctk.CTkOptionMenu(self.config_choose, values=config_choose_options, variable=self.prompt_draft_llm_var, font=("Microsoft YaHei", 12))
|
||||
prompt_draft_dropdown.grid(row=2, column=1, padx=5, pady=5, sticky="nsew")
|
||||
|
||||
create_label_with_help(self, parent=self.config_choose, label_text="定稿章节所用大模型", tooltip_key="final_chapter_llm_config", row=3, column=0, font=("Microsoft YaHei", 12))
|
||||
final_chapter_dropdown = ctk.CTkOptionMenu(self.config_choose, values=config_choose_options, variable=self.final_chapter_llm_var, font=("Microsoft YaHei", 12))
|
||||
final_chapter_dropdown.grid(row=3, column=1, padx=5, pady=5, sticky="nsew")
|
||||
|
||||
create_label_with_help(self, parent=self.config_choose, label_text="一致性审校所用大模型", tooltip_key="consistency_review_llm_config", row=4, column=0, font=("Microsoft YaHei", 12))
|
||||
consistency_review_dropdown = ctk.CTkOptionMenu(self.config_choose, values=config_choose_options, variable=self.consistency_review_llm_var, font=("Microsoft YaHei", 12))
|
||||
consistency_review_dropdown.grid(row=4, column=1, padx=5, pady=5, sticky="nsew")
|
||||
|
||||
def save_config_choose():
|
||||
config_data = load_config(self.config_file)["choose_configs"]
|
||||
if not config_data:
|
||||
config_data = {}
|
||||
config_data["architecture_llm"] = self.architecture_llm_var.get()
|
||||
config_data["chapter_outline_llm"] = self.chapter_outline_llm_var.get()
|
||||
config_data["prompt_draft_llm"] = self.prompt_draft_llm_var.get()
|
||||
config_data["final_chapter_llm"] = self.final_chapter_llm_var.get()
|
||||
config_data["consistency_review_llm"] = self.consistency_review_llm_var.get()
|
||||
|
||||
|
||||
config_data_full = load_config(self.config_file)
|
||||
config_data_full["choose_configs"] = config_data
|
||||
save_config(config_data_full, self.config_file)
|
||||
messagebox.showinfo("提示", "配置已保存。")
|
||||
|
||||
def refresh_config_dropdowns():
|
||||
"""刷新所有配置下拉菜单"""
|
||||
config_names = list(self.loaded_config.get("llm_configs", {}).keys())
|
||||
for dropdown in [architecture_dropdown, chapter_outline_dropdown, prompt_draft_dropdown, final_chapter_dropdown, consistency_review_dropdown]:
|
||||
dropdown.configure(values=config_names)
|
||||
if config_names and dropdown.cget("variable").get() not in config_names:
|
||||
dropdown.cget("variable").set(config_names[0])
|
||||
|
||||
save_btn = ctk.CTkButton(
|
||||
self.config_choose,
|
||||
text="保存配置",
|
||||
command=save_config_choose,
|
||||
font=("Microsoft YaHei", 12)
|
||||
)
|
||||
save_btn.grid(row=10, column=0,padx=2, pady=2, sticky="ew")
|
||||
|
||||
refresh_btn = ctk.CTkButton(
|
||||
self.config_choose,
|
||||
text="刷新配置",
|
||||
command=refresh_config_dropdowns,
|
||||
font=("Microsoft YaHei", 12)
|
||||
)
|
||||
refresh_btn.grid(row=10, column=1, padx=2, pady=2, sticky="ew")
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
def load_config_btn(self):
|
||||
cfg = load_config(self.config_file)
|
||||
if cfg:
|
||||
@@ -239,6 +623,7 @@ def load_config_btn(self):
|
||||
llm_configs = cfg.get("llm_configs", {})
|
||||
if last_llm in llm_configs:
|
||||
llm_conf = llm_configs[last_llm]
|
||||
self.interface_format_var.set(llm_conf.get("interface_format", "OpenAI"))
|
||||
self.api_key_var.set(llm_conf.get("api_key", ""))
|
||||
self.base_url_var.set(llm_conf.get("base_url", "https://api.openai.com/v1"))
|
||||
self.model_name_var.set(llm_conf.get("model_name", "gpt-4o-mini"))
|
||||
@@ -279,13 +664,16 @@ def save_config_btn(self):
|
||||
"model_name": self.model_name_var.get(),
|
||||
"temperature": self.temperature_var.get(),
|
||||
"max_tokens": self.max_tokens_var.get(),
|
||||
"timeout": self.safe_get_int(self.timeout_var, 600)
|
||||
"timeout": self.safe_get_int(self.timeout_var, 600),
|
||||
"interface_format": current_llm_interface
|
||||
}
|
||||
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)
|
||||
"retrieval_k": self.safe_get_int(self.embedding_retrieval_k_var, 4),
|
||||
"interface_format": current_embedding_interface
|
||||
|
||||
}
|
||||
other_params = {
|
||||
"topic": self.topic_text.get("0.0", "end").strip(),
|
||||
@@ -300,6 +688,8 @@ def save_config_btn(self):
|
||||
"scene_location": self.scene_location_var.get(),
|
||||
"time_constraint": self.time_constraint_var.get()
|
||||
}
|
||||
llm_config_name = self.base_url_var.get().split("/")[2] + " " + self.model_name_var.get()
|
||||
|
||||
existing_config = load_config(self.config_file)
|
||||
if not existing_config:
|
||||
existing_config = {}
|
||||
@@ -307,7 +697,9 @@ def save_config_btn(self):
|
||||
existing_config["last_embedding_interface_format"] = current_embedding_interface
|
||||
if "llm_configs" not in existing_config:
|
||||
existing_config["llm_configs"] = {}
|
||||
existing_config["llm_configs"][current_llm_interface] = llm_config
|
||||
llm_config["config_name"] = llm_config_name
|
||||
|
||||
existing_config["llm_configs"][llm_config_name] = llm_config
|
||||
|
||||
if "embedding_configs" not in existing_config:
|
||||
existing_config["embedding_configs"] = {}
|
||||
|
||||
+268
-37
@@ -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
|
||||
|
||||
@@ -32,13 +34,17 @@ def generate_novel_architecture_ui(self):
|
||||
|
||||
self.disable_button_safe(self.btn_generate_architecture)
|
||||
try:
|
||||
interface_format = self.interface_format_var.get().strip()
|
||||
api_key = self.api_key_var.get().strip()
|
||||
base_url = self.base_url_var.get().strip()
|
||||
model_name = self.model_name_var.get().strip()
|
||||
temperature = self.temperature_var.get()
|
||||
max_tokens = self.max_tokens_var.get()
|
||||
timeout_val = self.safe_get_int(self.timeout_var, 600)
|
||||
|
||||
|
||||
interface_format = self.loaded_config["llm_configs"][self.architecture_llm_var.get()]["interface_format"]
|
||||
api_key = self.loaded_config["llm_configs"][self.architecture_llm_var.get()]["api_key"]
|
||||
base_url = self.loaded_config["llm_configs"][self.architecture_llm_var.get()]["base_url"]
|
||||
model_name = self.loaded_config["llm_configs"][self.architecture_llm_var.get()]["model_name"]
|
||||
temperature = self.loaded_config["llm_configs"][self.architecture_llm_var.get()]["temperature"]
|
||||
max_tokens = self.loaded_config["llm_configs"][self.architecture_llm_var.get()]["max_tokens"]
|
||||
timeout_val = self.loaded_config["llm_configs"][self.architecture_llm_var.get()]["timeout"]
|
||||
|
||||
|
||||
|
||||
topic = self.topic_text.get("0.0", "end").strip()
|
||||
genre = self.genre_var.get().strip()
|
||||
@@ -82,14 +88,18 @@ def generate_chapter_blueprint_ui(self):
|
||||
return
|
||||
self.disable_button_safe(self.btn_generate_directory)
|
||||
try:
|
||||
interface_format = self.interface_format_var.get().strip()
|
||||
api_key = self.api_key_var.get().strip()
|
||||
base_url = self.base_url_var.get().strip()
|
||||
model_name = self.model_name_var.get().strip()
|
||||
|
||||
number_of_chapters = self.safe_get_int(self.num_chapters_var, 10)
|
||||
temperature = self.temperature_var.get()
|
||||
max_tokens = self.max_tokens_var.get()
|
||||
timeout_val = self.safe_get_int(self.timeout_var, 600)
|
||||
|
||||
interface_format = self.loaded_config["llm_configs"][self.chapter_outline_llm_var.get()]["interface_format"]
|
||||
api_key = self.loaded_config["llm_configs"][self.chapter_outline_llm_var.get()]["api_key"]
|
||||
base_url = self.loaded_config["llm_configs"][self.chapter_outline_llm_var.get()]["base_url"]
|
||||
model_name = self.loaded_config["llm_configs"][self.chapter_outline_llm_var.get()]["model_name"]
|
||||
temperature = self.loaded_config["llm_configs"][self.chapter_outline_llm_var.get()]["temperature"]
|
||||
max_tokens = self.loaded_config["llm_configs"][self.chapter_outline_llm_var.get()]["max_tokens"]
|
||||
timeout_val = self.loaded_config["llm_configs"][self.chapter_outline_llm_var.get()]["timeout"]
|
||||
|
||||
|
||||
user_guidance = self.user_guide_text.get("0.0", "end").strip() # 新增获取用户指导
|
||||
|
||||
self.safe_log("开始生成章节蓝图...")
|
||||
@@ -121,13 +131,15 @@ def generate_chapter_draft_ui(self):
|
||||
def task():
|
||||
self.disable_button_safe(self.btn_generate_chapter)
|
||||
try:
|
||||
interface_format = self.interface_format_var.get().strip()
|
||||
api_key = self.api_key_var.get().strip()
|
||||
base_url = self.base_url_var.get().strip()
|
||||
model_name = self.model_name_var.get().strip()
|
||||
temperature = self.temperature_var.get()
|
||||
max_tokens = self.max_tokens_var.get()
|
||||
timeout_val = self.safe_get_int(self.timeout_var, 600)
|
||||
|
||||
interface_format = self.loaded_config["llm_configs"][self.prompt_draft_llm_var.get()]["interface_format"]
|
||||
api_key = self.loaded_config["llm_configs"][self.prompt_draft_llm_var.get()]["api_key"]
|
||||
base_url = self.loaded_config["llm_configs"][self.prompt_draft_llm_var.get()]["base_url"]
|
||||
model_name = self.loaded_config["llm_configs"][self.prompt_draft_llm_var.get()]["model_name"]
|
||||
temperature = self.loaded_config["llm_configs"][self.prompt_draft_llm_var.get()]["temperature"]
|
||||
max_tokens = self.loaded_config["llm_configs"][self.prompt_draft_llm_var.get()]["max_tokens"]
|
||||
timeout_val = self.loaded_config["llm_configs"][self.prompt_draft_llm_var.get()]["timeout"]
|
||||
|
||||
|
||||
chap_num = self.safe_get_int(self.chapter_num_var, 1)
|
||||
word_number = self.safe_get_int(self.word_number_var, 3000)
|
||||
@@ -147,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,
|
||||
@@ -312,13 +323,15 @@ def finalize_chapter_ui(self):
|
||||
|
||||
self.disable_button_safe(self.btn_finalize_chapter)
|
||||
try:
|
||||
interface_format = self.interface_format_var.get().strip()
|
||||
api_key = self.api_key_var.get().strip()
|
||||
base_url = self.base_url_var.get().strip()
|
||||
model_name = self.model_name_var.get().strip()
|
||||
temperature = self.temperature_var.get()
|
||||
max_tokens = self.max_tokens_var.get()
|
||||
timeout_val = self.safe_get_int(self.timeout_var, 600)
|
||||
|
||||
interface_format = self.loaded_config["llm_configs"][self.final_chapter_llm_var.get()]["interface_format"]
|
||||
api_key = self.loaded_config["llm_configs"][self.final_chapter_llm_var.get()]["api_key"]
|
||||
base_url = self.loaded_config["llm_configs"][self.final_chapter_llm_var.get()]["base_url"]
|
||||
model_name = self.loaded_config["llm_configs"][self.final_chapter_llm_var.get()]["model_name"]
|
||||
temperature = self.loaded_config["llm_configs"][self.final_chapter_llm_var.get()]["temperature"]
|
||||
max_tokens = self.loaded_config["llm_configs"][self.final_chapter_llm_var.get()]["max_tokens"]
|
||||
timeout_val = self.loaded_config["llm_configs"][self.final_chapter_llm_var.get()]["timeout"]
|
||||
|
||||
|
||||
embedding_api_key = self.embedding_api_key_var.get().strip()
|
||||
embedding_url = self.embedding_url_var.get().strip()
|
||||
@@ -392,13 +405,14 @@ def do_consistency_check(self):
|
||||
def task():
|
||||
self.disable_button_safe(self.btn_check_consistency)
|
||||
try:
|
||||
api_key = self.api_key_var.get().strip()
|
||||
base_url = self.base_url_var.get().strip()
|
||||
model_name = self.model_name_var.get().strip()
|
||||
temperature = self.temperature_var.get()
|
||||
interface_format = self.interface_format_var.get()
|
||||
max_tokens = self.max_tokens_var.get()
|
||||
timeout = self.timeout_var.get()
|
||||
interface_format = self.loaded_config["llm_configs"][self.consistency_review_llm_var.get()]["interface_format"]
|
||||
api_key = self.loaded_config["llm_configs"][self.consistency_review_llm_var.get()]["api_key"]
|
||||
base_url = self.loaded_config["llm_configs"][self.consistency_review_llm_var.get()]["base_url"]
|
||||
model_name = self.loaded_config["llm_configs"][self.consistency_review_llm_var.get()]["model_name"]
|
||||
temperature = self.loaded_config["llm_configs"][self.consistency_review_llm_var.get()]["temperature"]
|
||||
max_tokens = self.loaded_config["llm_configs"][self.consistency_review_llm_var.get()]["max_tokens"]
|
||||
timeout = self.loaded_config["llm_configs"][self.consistency_review_llm_var.get()]["timeout"]
|
||||
|
||||
|
||||
chap_num = self.safe_get_int(self.chapter_num_var, 1)
|
||||
chap_file = os.path.join(filepath, "chapters", f"chapter_{chap_num}.txt")
|
||||
@@ -430,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(
|
||||
|
||||
+11
-1
@@ -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")
|
||||
|
||||
+44
-14
@@ -26,13 +26,16 @@ 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
|
||||
from ui.character_tab import build_character_tab, load_character_state, save_character_state
|
||||
from ui.summary_tab import build_summary_tab, load_global_summary, save_global_summary
|
||||
from ui.chapters_tab import build_chapters_tab, refresh_chapters_list, on_chapter_selected, load_chapter_content, save_current_chapter, prev_chapter, next_chapter
|
||||
from ui.other_settings import build_other_settings_tab
|
||||
|
||||
|
||||
class NovelGeneratorGUI:
|
||||
"""
|
||||
@@ -53,23 +56,27 @@ class NovelGeneratorGUI:
|
||||
self.loaded_config = load_config(self.config_file)
|
||||
|
||||
if self.loaded_config:
|
||||
last_llm = self.loaded_config.get("last_interface_format", "OpenAI")
|
||||
last_llm = next(iter(self.loaded_config["llm_configs"].values())).get("interface_format", "OpenAI")
|
||||
|
||||
last_embedding = self.loaded_config.get("last_embedding_interface_format", "OpenAI")
|
||||
else:
|
||||
last_llm = "OpenAI"
|
||||
last_embedding = "OpenAI"
|
||||
|
||||
if self.loaded_config and "llm_configs" in self.loaded_config and last_llm in self.loaded_config["llm_configs"]:
|
||||
llm_conf = self.loaded_config["llm_configs"][last_llm]
|
||||
else:
|
||||
llm_conf = {
|
||||
"api_key": "",
|
||||
"base_url": "https://api.openai.com/v1",
|
||||
"model_name": "gpt-4o-mini",
|
||||
"temperature": 0.7,
|
||||
"max_tokens": 8192,
|
||||
"timeout": 600
|
||||
}
|
||||
# if self.loaded_config and "llm_configs" in self.loaded_config and last_llm in self.loaded_config["llm_configs"]:
|
||||
# llm_conf = next(iter(self.loaded_config["llm_configs"]))
|
||||
# else:
|
||||
# llm_conf = {
|
||||
# "api_key": "",
|
||||
# "base_url": "https://api.openai.com/v1",
|
||||
# "model_name": "gpt-4o-mini",
|
||||
# "temperature": 0.7,
|
||||
# "max_tokens": 8192,
|
||||
# "timeout": 600
|
||||
# }
|
||||
llm_conf = next(iter(self.loaded_config["llm_configs"].values()))
|
||||
choose_configs = self.loaded_config.get("choose_configs", {})
|
||||
|
||||
|
||||
if self.loaded_config and "embedding_configs" in self.loaded_config and last_embedding in self.loaded_config["embedding_configs"]:
|
||||
emb_conf = self.loaded_config["embedding_configs"][last_embedding]
|
||||
@@ -82,13 +89,17 @@ class NovelGeneratorGUI:
|
||||
}
|
||||
|
||||
# -- LLM通用参数 --
|
||||
# self.llm_conf_name = next(iter(self.loaded_config["llm_configs"]))
|
||||
self.api_key_var = ctk.StringVar(value=llm_conf.get("api_key", ""))
|
||||
self.base_url_var = ctk.StringVar(value=llm_conf.get("base_url", "https://api.openai.com/v1"))
|
||||
self.interface_format_var = ctk.StringVar(value=last_llm)
|
||||
self.interface_format_var = ctk.StringVar(value=llm_conf.get("interface_format", "OpenAI"))
|
||||
self.model_name_var = ctk.StringVar(value=llm_conf.get("model_name", "gpt-4o-mini"))
|
||||
self.temperature_var = ctk.DoubleVar(value=llm_conf.get("temperature", 0.7))
|
||||
self.max_tokens_var = ctk.IntVar(value=llm_conf.get("max_tokens", 8192))
|
||||
self.timeout_var = ctk.IntVar(value=llm_conf.get("timeout", 600))
|
||||
self.interface_config_var = ctk.StringVar(value=next(iter(self.loaded_config["llm_configs"])))
|
||||
|
||||
|
||||
|
||||
# -- Embedding相关 --
|
||||
self.embedding_interface_format_var = ctk.StringVar(value=last_embedding)
|
||||
@@ -97,6 +108,18 @@ class NovelGeneratorGUI:
|
||||
self.embedding_model_name_var = ctk.StringVar(value=emb_conf.get("model_name", "text-embedding-ada-002"))
|
||||
self.embedding_retrieval_k_var = ctk.StringVar(value=str(emb_conf.get("retrieval_k", 4)))
|
||||
|
||||
|
||||
# -- 生成配置相关 --
|
||||
self.architecture_llm_var = ctk.StringVar(value=choose_configs.get("architecture_llm", "DeepSeek"))
|
||||
self.chapter_outline_llm_var = ctk.StringVar(value=choose_configs.get("chapter_outline_llm", "DeepSeek"))
|
||||
self.final_chapter_llm_var = ctk.StringVar(value=choose_configs.get("final_chapter_llm", "DeepSeek"))
|
||||
self.consistency_review_llm_var = ctk.StringVar(value=choose_configs.get("consistency_review_llm", "DeepSeek"))
|
||||
self.prompt_draft_llm_var = ctk.StringVar(value=choose_configs.get("prompt_draft_llm", "DeepSeek"))
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
# -- 小说参数相关 --
|
||||
if self.loaded_config and "other_params" in self.loaded_config:
|
||||
op = self.loaded_config["other_params"]
|
||||
@@ -111,6 +134,10 @@ class NovelGeneratorGUI:
|
||||
self.scene_location_var = ctk.StringVar(value=op.get("scene_location", ""))
|
||||
self.time_constraint_var = ctk.StringVar(value=op.get("time_constraint", ""))
|
||||
self.user_guidance_default = op.get("user_guidance", "")
|
||||
self.webdav_url_var = ctk.StringVar(value=op.get("webdav_url", ""))
|
||||
self.webdav_username_var = ctk.StringVar(value=op.get("webdav_username", ""))
|
||||
self.webdav_password_var = ctk.StringVar(value=op.get("webdav_password", ""))
|
||||
|
||||
else:
|
||||
self.topic_default = ""
|
||||
self.genre_var = ctk.StringVar(value="玄幻")
|
||||
@@ -138,6 +165,8 @@ class NovelGeneratorGUI:
|
||||
build_character_tab(self)
|
||||
build_summary_tab(self)
|
||||
build_chapters_tab(self)
|
||||
build_other_settings_tab(self)
|
||||
|
||||
|
||||
# ----------------- 通用辅助函数 -----------------
|
||||
def show_tooltip(self, key: str):
|
||||
@@ -344,6 +373,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
|
||||
|
||||
@@ -0,0 +1,277 @@
|
||||
# ui/other_settings.py
|
||||
import customtkinter as ctk
|
||||
from ui.config_tab import create_label_with_help
|
||||
from tkinter import messagebox
|
||||
from config_manager import load_config, save_config
|
||||
import requests
|
||||
from requests.auth import HTTPBasicAuth
|
||||
import os
|
||||
from xml.etree import ElementTree as ET
|
||||
import shutil
|
||||
import time
|
||||
def build_other_settings_tab(self):
|
||||
self.other_settings_tab = self.tabview.add("Other Settings")
|
||||
self.other_settings_tab.rowconfigure(0, weight=1)
|
||||
self.other_settings_tab.columnconfigure(0, weight=1)
|
||||
if "webdav_config" not in self.loaded_config:
|
||||
self.loaded_config["webdav_config"] = {
|
||||
"webdav_url": "",
|
||||
"webdav_username": "",
|
||||
"webdav_password": ""
|
||||
}
|
||||
|
||||
self.webdav_url_var.set(self.loaded_config["webdav_config"].get("webdav_url", ""))
|
||||
self.webdav_username_var.set(self.loaded_config["webdav_config"].get("webdav_username", ""))
|
||||
self.webdav_password_var.set(self.loaded_config["webdav_config"].get("webdav_password", ""))
|
||||
|
||||
|
||||
def save_webdav_settings():
|
||||
self.loaded_config["webdav_config"]["webdav_url"] = self.webdav_url_var.get().strip()
|
||||
self.loaded_config["webdav_config"]["webdav_username"] = self.webdav_username_var.get().strip()
|
||||
self.loaded_config["webdav_config"]["webdav_password"] = self.webdav_password_var.get().strip()
|
||||
save_config(self.loaded_config, self.config_file)
|
||||
|
||||
|
||||
def test_webdav_connection(test = True):
|
||||
try:
|
||||
client = WebDAVClient(self.webdav_url_var.get().strip(),self.webdav_username_var.get().strip(),self.webdav_password_var.get().strip())
|
||||
client.list_directory()
|
||||
if not test:
|
||||
save_webdav_settings()
|
||||
return True
|
||||
messagebox.showinfo("成功", "WebDAV 连接成功!")
|
||||
save_webdav_settings()
|
||||
return True
|
||||
|
||||
except Exception as e:
|
||||
print(e)
|
||||
|
||||
messagebox.showerror("错误", f"发生未知错误: {e}")
|
||||
return False
|
||||
|
||||
def backup_to_webdav():
|
||||
try:
|
||||
target_dir = "AI_Novel_Generator"
|
||||
client = WebDAVClient(self.webdav_url_var.get().strip(),self.webdav_username_var.get().strip(),self.webdav_password_var.get().strip())
|
||||
if not client.ensure_directory_exists(target_dir):
|
||||
client.create_directory(target_dir)
|
||||
client.upload_file(self.config_file, f"{target_dir}/config.json")
|
||||
messagebox.showinfo("成功", "配置备份成功!")
|
||||
except Exception as e:
|
||||
print(e)
|
||||
messagebox.showerror("错误", f"发生未知错误: {e}")
|
||||
return False
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
def restore_from_webdav():
|
||||
try:
|
||||
target_dir = "AI_Novel_Generator"
|
||||
client = WebDAVClient(self.webdav_url_var.get().strip(),self.webdav_username_var.get().strip(),self.webdav_password_var.get().strip())
|
||||
client.download_file(f"{target_dir}/config.json", self.config_file)
|
||||
self.loaded_config = load_config(self.config_file)
|
||||
messagebox.showinfo("成功", "配置恢复成功!")
|
||||
|
||||
except Exception as e:
|
||||
print(e)
|
||||
messagebox.showerror("错误", f"发生未知错误: {e}")
|
||||
return False
|
||||
|
||||
|
||||
|
||||
|
||||
dav_frame = ctk.CTkFrame(self.other_settings_tab)
|
||||
dav_frame.pack(padx=20, pady=20, fill="x")
|
||||
|
||||
dav_title = ctk.CTkLabel(dav_frame, text="webdav设置", font=("Microsoft YaHei", 16, "bold"))
|
||||
dav_title.pack(anchor="w", padx=5, pady=(0, 5))
|
||||
dav_warp_frame = ctk.CTkFrame(dav_frame, corner_radius=10, border_width=2, border_color="gray")
|
||||
dav_warp_frame.pack(fill="x", padx=5)
|
||||
dav_warp_frame.columnconfigure(1, weight=1)
|
||||
|
||||
|
||||
|
||||
create_label_with_help(self, parent=dav_warp_frame, label_text="Webdav URL", tooltip_key="webdav_url",row=0, column=0, font=("Microsoft YaHei", 12), sticky="w")
|
||||
dav_url_entry = ctk.CTkEntry(dav_warp_frame, textvariable=self.webdav_url_var, font=("Microsoft YaHei", 12))
|
||||
dav_url_entry.grid(row=0, column=1, padx=5, pady=5, sticky="w")
|
||||
|
||||
create_label_with_help(self, parent=dav_warp_frame, label_text="Webdav用户名", tooltip_key="webdav_username",row=1, column=0, font=("Microsoft YaHei", 12), sticky="w")
|
||||
dav_username_entry = ctk.CTkEntry(dav_warp_frame, textvariable=self.webdav_username_var, font=("Microsoft YaHei", 12))
|
||||
dav_username_entry.grid(row=1, column=1, padx=5, pady=5, sticky="w")
|
||||
|
||||
create_label_with_help(self, parent=dav_warp_frame, label_text="Webdav密码", tooltip_key="webdav_password",row=2, column=0, font=("Microsoft YaHei", 12), sticky="w")
|
||||
dav_password_entry = ctk.CTkEntry(dav_warp_frame, textvariable=self.webdav_password_var, font=("Microsoft YaHei", 12), show="*")
|
||||
dav_password_entry.grid(row=2, column=1, padx=5, pady=5, sticky="w")
|
||||
|
||||
button_frame = ctk.CTkFrame(dav_warp_frame)
|
||||
button_frame.grid(row=3, column=0, columnspan=2, padx=5, pady=10, sticky="w")
|
||||
|
||||
# 测试连接按钮
|
||||
test_btn = ctk.CTkButton(button_frame, text="测试连接", font=("Microsoft YaHei", 12),
|
||||
command=test_webdav_connection)
|
||||
test_btn.pack(side="left", padx=5)
|
||||
|
||||
# 保存设置按钮
|
||||
save_btn = ctk.CTkButton(button_frame, text="备份", font=("Microsoft YaHei", 12),
|
||||
command=backup_to_webdav)
|
||||
save_btn.pack(side="left", padx=5)
|
||||
|
||||
# 重置按钮
|
||||
reset_btn = ctk.CTkButton(button_frame, text="恢复", font=("Microsoft YaHei", 12),
|
||||
command=restore_from_webdav)
|
||||
reset_btn.pack(side="left", padx=5)
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
class WebDAVClient:
|
||||
def __init__(self, base_url, username, password):
|
||||
"""初始化WebDAV客户端"""
|
||||
self.base_url = base_url.rstrip('/') + '/'
|
||||
self.auth = HTTPBasicAuth(username, password)
|
||||
self.headers = {
|
||||
'User-Agent': 'Python WebDAV Client',
|
||||
'Accept': '*/*'
|
||||
}
|
||||
# WebDAV命名空间
|
||||
self.ns = {'d': 'DAV:'}
|
||||
|
||||
def _get_url(self, path):
|
||||
"""获取完整的资源URL"""
|
||||
return self.base_url + path.lstrip('/')
|
||||
|
||||
def directory_exists(self, path):
|
||||
"""
|
||||
检查目录是否存在
|
||||
:param path: 目录路径
|
||||
:return: 布尔值,表示目录是否存在
|
||||
"""
|
||||
url = self._get_url(path)
|
||||
headers = self.headers.copy()
|
||||
headers['Depth'] = '0' # 只检查当前资源
|
||||
|
||||
try:
|
||||
# 发送PROPFIND请求检查资源是否存在
|
||||
response = requests.request('PROPFIND', url, headers=headers, auth=self.auth)
|
||||
|
||||
# 207 Multi-Status表示成功,说明资源存在
|
||||
if response.status_code == 207:
|
||||
# 解析XML响应,确认是目录
|
||||
root = ET.fromstring(response.content)
|
||||
# 查找资源类型属性
|
||||
res_type = root.find('.//d:resourcetype', namespaces=self.ns)
|
||||
# 如果包含collection元素,则是目录
|
||||
if res_type is not None and res_type.find('d:collection', namespaces=self.ns) is not None:
|
||||
return True
|
||||
return False
|
||||
except requests.exceptions.RequestException as e:
|
||||
print(f"检查目录存在性时出错: {e}")
|
||||
return False
|
||||
|
||||
def create_directory(self, path):
|
||||
"""
|
||||
创建远程目录
|
||||
:param path: 要创建的目录路径
|
||||
:return: 是否创建成功
|
||||
"""
|
||||
url = self._get_url(path)
|
||||
|
||||
try:
|
||||
response = requests.request('MKCOL', url, auth=self.auth, headers=self.headers)
|
||||
response.raise_for_status()
|
||||
|
||||
print(f"目录创建成功: {path}")
|
||||
return True
|
||||
except requests.exceptions.RequestException as e:
|
||||
print(f"目录创建失败: {e}")
|
||||
return False
|
||||
|
||||
def ensure_directory_exists(self, path):
|
||||
"""
|
||||
确保目录存在,如果不存在则创建
|
||||
:param path: 目录路径
|
||||
:return: 布尔值,表示最终目录是否存在
|
||||
"""
|
||||
# 移除末尾的斜杠(如果有)
|
||||
path = path.rstrip('/')
|
||||
|
||||
# 如果目录已经存在,直接返回True
|
||||
if self.directory_exists(path):
|
||||
print(f"目录已存在: {path}")
|
||||
return True
|
||||
|
||||
# 递归创建父目录
|
||||
parent_dir = os.path.dirname(path)
|
||||
if parent_dir and not self.directory_exists(parent_dir):
|
||||
# 如果父目录不存在,则先创建父目录
|
||||
if not self.ensure_directory_exists(parent_dir):
|
||||
print(f"创建父目录失败: {parent_dir}")
|
||||
return False
|
||||
|
||||
# 创建当前目录
|
||||
return self.create_directory(path)
|
||||
def upload_file(self, local_path, remote_path):
|
||||
"""
|
||||
上传文件到WebDAV服务器
|
||||
:param local_path: 本地文件路径
|
||||
:param remote_path: 远程文件路径
|
||||
:return: 是否上传成功
|
||||
"""
|
||||
if not os.path.isfile(local_path):
|
||||
print(f"本地文件不存在: {local_path}")
|
||||
return False
|
||||
|
||||
url = self._get_url(remote_path)
|
||||
|
||||
try:
|
||||
with open(local_path, 'rb') as f:
|
||||
response = requests.put(url, data=f, auth=self.auth, headers=self.headers)
|
||||
response.raise_for_status()
|
||||
|
||||
print(f"文件上传成功: {local_path} -> {remote_path}")
|
||||
return True
|
||||
except requests.exceptions.RequestException as e:
|
||||
print(f"文件上传失败: {e}")
|
||||
return False
|
||||
def download_file(self, remote_path, local_path):
|
||||
"""
|
||||
从WebDAV服务器下载文件
|
||||
:param remote_path: 远程文件路径
|
||||
:param local_path: 本地保存路径
|
||||
:return: 是否下载成功
|
||||
"""
|
||||
url = self._get_url(remote_path)
|
||||
local_path = os.path.join(os.path.dirname(os.path.dirname(os.path.abspath(__file__))), local_path)
|
||||
self.backup(local_path)
|
||||
try:
|
||||
response = requests.get(url, auth=self.auth, headers=self.headers, stream=True)
|
||||
response.raise_for_status()
|
||||
|
||||
# 创建本地目录(如果需要)
|
||||
os.makedirs(os.path.dirname(local_path), exist_ok=True)
|
||||
|
||||
with open(local_path, 'wb') as f:
|
||||
for chunk in response.iter_content(chunk_size=8192):
|
||||
f.write(chunk)
|
||||
|
||||
print(f"文件下载成功: {remote_path} -> {local_path}")
|
||||
return True
|
||||
except requests.exceptions.RequestException as e:
|
||||
print(f"文件下载失败: {e}")
|
||||
return False
|
||||
def backup(self, local_path):
|
||||
name_parts = os.path.basename(local_path).rsplit('.', 1) # 只分割最后一个点
|
||||
base_name = name_parts[0]
|
||||
extension = name_parts[1]
|
||||
timestamp = time.strftime("%Y%m%d%H%M%S")
|
||||
if not os.path.exists(os.path.join(os.path.dirname(local_path), "backup")):
|
||||
os.makedirs(os.path.join(os.path.dirname(local_path), "backup"))
|
||||
backup_file_name = f"{base_name}_{timestamp}_bak.{extension}"
|
||||
shutil.copy2(os.path.basename(local_path), os.path.join(os.path.dirname(local_path), "backup", backup_file_name))
|
||||
Reference in New Issue
Block a user