Merge pull request #184 from ahhhhhhhman/main

1.增加在不同环节使用不同的大模型配置的功能 2.增加批量生成章节的功能
This commit is contained in:
xianyun
2025-09-04 11:45:35 +08:00
committed by GitHub
15 changed files with 1214 additions and 163 deletions
+4
View File
@@ -12,3 +12,7 @@ config_test.json
/novel_generator/__pycache__ /novel_generator/__pycache__
/ui/__pycache__ /ui/__pycache__
.idea/ .idea/
/novel
app.log
test.py
/backup
+62
View File
@@ -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"
}
}
+7
View File
@@ -16,6 +16,13 @@ from prompt_definitions import (
plot_architecture_prompt, plot_architecture_prompt,
create_character_state_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 from utils import clear_file_content, save_string_to_txt
def load_partial_architecture_data(filepath: str) -> dict: def load_partial_architecture_data(filepath: str) -> dict:
+8 -2
View File
@@ -10,7 +10,13 @@ from novel_generator.common import invoke_with_cleaning
from llm_adapters import create_llm_adapter from llm_adapters import create_llm_adapter
from prompt_definitions import chapter_blueprint_prompt, chunked_chapter_blueprint_prompt from prompt_definitions import chapter_blueprint_prompt, chunked_chapter_blueprint_prompt
from utils import read_file, clear_file_content, save_string_to_txt 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: def compute_chunk_size(number_of_chapters: int, max_tokens: int) -> int:
""" """
基于“每章约100 tokens”的粗略估算, 基于“每章约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 = (floor(max_tokens/100/10)*10) - 10
并确保 chunk_size 不会小于1或大于实际章节数。 并确保 chunk_size 不会小于1或大于实际章节数。
""" """
tokens_per_chapter = 100.0 tokens_per_chapter = 200.0
ratio = max_tokens / tokens_per_chapter ratio = max_tokens / tokens_per_chapter
ratio_rounded_to_10 = int(ratio // 10) * 10 ratio_rounded_to_10 = int(ratio // 10) * 10
chunk_size = ratio_rounded_to_10 - 10 chunk_size = ratio_rounded_to_10 - 10
+7
View File
@@ -22,6 +22,13 @@ from novel_generator.vectorstore_utils import (
get_relevant_context_from_vector_store, get_relevant_context_from_vector_store,
load_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: def get_last_n_chapters_text(chapters_dir: str, current_chapter_num: int, n: int = 3) -> list:
""" """
+7 -1
View File
@@ -7,7 +7,13 @@ import logging
import re import re
import time import time
import traceback 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): def call_with_retry(func, max_retries=3, sleep_time=2, fallback_return=None, **kwargs):
""" """
通用的重试机制封装。 通用的重试机制封装。
+8 -2
View File
@@ -11,7 +11,13 @@ from prompt_definitions import summary_prompt, update_character_state_prompt
from novel_generator.common import invoke_with_cleaning from novel_generator.common import invoke_with_cleaning
from utils import read_file, clear_file_content, save_string_to_txt from utils import read_file, clear_file_content, save_string_to_txt
from novel_generator.vectorstore_utils import update_vector_store 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( def finalize_chapter(
novel_number: int, novel_number: int,
word_number: int, word_number: int,
@@ -111,7 +117,7 @@ def enrich_chapter_text(
max_tokens=max_tokens, max_tokens=max_tokens,
timeout=timeout timeout=timeout
) )
prompt = f"""以下章节文本较短,请在保持剧情连贯的前提下进行扩写,使其更充实,接近 {word_number} 字左右: prompt = f"""以下章节文本较短,请在保持剧情连贯的前提下进行扩写,使其更充实,接近 {word_number} 字左右,仅给出最终文本,不要解释任何内容。
原内容: 原内容:
{chapter_text} {chapter_text}
""" """
+9 -3
View File
@@ -16,11 +16,17 @@ from langchain.docstore.document import Document
# 禁用特定的Torch警告 # 禁用特定的Torch警告
warnings.filterwarnings('ignore', message='.*Torch was not compiled with flash attention.*') warnings.filterwarnings('ignore', message='.*Torch was not compiled with flash attention.*')
os.environ["TOKENIZERS_PARALLELISM"] = "false" 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: def advanced_split_content(content: str, similarity_threshold: float = 0.7, max_length: int = 500) -> list:
"""使用基本分段策略""" """使用基本分段策略"""
nltk.download('punkt', quiet=True) # nltk.download('punkt', quiet=True)
nltk.download('punkt_tab', quiet=True) # nltk.download('punkt_tab', quiet=True)
sentences = nltk.sent_tokenize(content) sentences = nltk.sent_tokenize(content)
if not sentences: if not sentences:
return [] return []
+9 -3
View File
@@ -13,7 +13,13 @@ import ssl
import requests import requests
import warnings import warnings
from langchain_chroma import Chroma 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警告 # 禁用特定的Torch警告
warnings.filterwarnings('ignore', message='.*Torch was not compiled with flash attention.*') warnings.filterwarnings('ignore', message='.*Torch was not compiled with flash attention.*')
os.environ["TOKENIZERS_PARALLELISM"] = "false" # 禁用tokenizer并行警告 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(): if not chapter_text.strip():
return [] return []
nltk.download('punkt', quiet=True) # nltk.download('punkt', quiet=True)
nltk.download('punkt_tab', quiet=True) # nltk.download('punkt_tab', quiet=True)
sentences = nltk.sent_tokenize(chapter_text) sentences = nltk.sent_tokenize(chapter_text)
if not sentences: if not sentences:
return [] return []
+2 -1
View File
@@ -33,5 +33,6 @@ tooltips = {
"characters_involved": "本章需要重点描写或影响剧情的角色名单。", "characters_involved": "本章需要重点描写或影响剧情的角色名单。",
"key_items": "在本章中出现的重要道具、线索或物品。", "key_items": "在本章中出现的重要道具、线索或物品。",
"scene_location": "本章主要发生的地点或场景描述。", "scene_location": "本章主要发生的地点或场景描述。",
"time_constraint": "本章剧情中涉及的时间压力或时限设置。" "time_constraint": "本章剧情中涉及的时间压力或时限设置。",
"interface_config": "选择你要使用的AI接口配置。"
} }
+485 -93
View File
@@ -1,6 +1,8 @@
# ui/config_tab.py # ui/config_tab.py
# -*- coding: utf-8 -*- # -*- coding: utf-8 -*-
from tkinter import messagebox from tkinter import messagebox
import uuid
import datetime
import customtkinter as ctk 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.ai_config_tab = self.config_tabview.add("LLM Model settings")
self.embeddings_config_tab = self.config_tabview.add("Embedding 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_ai_config_tab(self)
build_embeddings_config_tab(self) build_embeddings_config_tab(self)
build_config_choose_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")
def build_ai_config_tab(self): def build_ai_config_tab(self):
def on_interface_format_changed(new_value): def refresh_config_dropdown():
self.interface_format_var.set(new_value) """刷新配置下拉菜单"""
config_data = load_config(self.config_file) config_names = list(self.loaded_config.get("llm_configs", {}).keys())
if config_data: interface_config_dropdown.configure(values=config_names)
config_data["last_interface_format"] = new_value if config_names and self.interface_config_var.get() not in config_names:
save_config(config_data, self.config_file) self.interface_config_var.set(config_names[0])
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")
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_rowconfigure(i, weight=0)
self.ai_config_tab.grid_columnconfigure(0, 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(1, weight=1)
self.ai_config_tab.grid_columnconfigure(2, weight=0) 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 # 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)) create_label_with_help(self, self.ai_config_tab, "API Key:", "api_key", row_start, 0)
api_key_entry = ctk.CTkEntry(self.ai_config_tab, textvariable=self.api_key_var, font=("Microsoft YaHei", 12),show="*") self.api_key_var = ctk.StringVar(value="")
api_key_entry.grid(row=0, column=1, padx=5, pady=5, columnspan=2, sticky="nsew") 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 # 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)) create_label_with_help(self, self.ai_config_tab, "Base URL:", "base_url", row_start+1, 0)
base_url_entry = ctk.CTkEntry(self.ai_config_tab, textvariable=self.base_url_var, font=("Microsoft YaHei", 12)) self.base_url_var = ctk.StringVar(value="")
base_url_entry.grid(row=1, column=1, padx=5, pady=5, columnspan=2, sticky="nsew") 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) 接口格式 # 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)) create_label_with_help(self, self.ai_config_tab, "接口格式:", "interface_format", row_start+2, 0)
# 在接口选项列表中添加 "Grok" self.interface_format_var = ctk.StringVar(value="OpenAI")
interface_options = ["DeepSeek", "阿里云百炼", "OpenAI", "Azure OpenAI", "Azure AI", "Ollama", "ML Studio", "Gemini", "火山引擎", "硅基流动", "Grok"] 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, command=on_interface_format_changed, font=("Microsoft YaHei", 12)) interface_dropdown = ctk.CTkOptionMenu(
interface_dropdown.grid(row=2, column=1, padx=5, pady=5, columnspan=2, sticky="nsew") 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 # 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)) create_label_with_help(self, self.ai_config_tab, "模型名称:", "model_name", row_start+3, 0)
model_name_entry = ctk.CTkEntry(self.ai_config_tab, textvariable=self.model_name_var, font=("Microsoft YaHei", 12)) self.model_name_var = ctk.StringVar(value="")
model_name_entry.grid(row=3, column=1, padx=5, pady=5, columnspan=2, sticky="nsew") 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 # 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): def update_temp_label(value):
self.temp_value_label.configure(text=f"{float(value):.2f}") 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 = ctk.CTkSlider(
temp_scale.grid(row=4, column=1, padx=5, pady=5, sticky="we") self.ai_config_tab,
self.temp_value_label = ctk.CTkLabel(self.ai_config_tab, text=f"{self.temperature_var.get():.2f}", font=("Microsoft YaHei", 12)) from_=0.0,
self.temp_value_label.grid(row=4, column=2, padx=5, pady=5, sticky="w") 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 # 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): def update_max_tokens_label(value):
self.max_tokens_value_label.configure(text=str(int(float(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 = ctk.CTkSlider(
max_tokens_slider.grid(row=5, column=1, padx=5, pady=5, sticky="we") self.ai_config_tab,
self.max_tokens_value_label = ctk.CTkLabel(self.ai_config_tab, text=str(self.max_tokens_var.get()), font=("Microsoft YaHei", 12)) from_=0,
self.max_tokens_value_label.grid(row=5, column=2, padx=5, pady=5, sticky="w") 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 (sec) # 7) Timeout
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)) 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): def update_timeout_label(value):
integer_val = int(float(value)) self.timeout_value_label.configure(text=str(int(float(value))))
self.timeout_value_label.configure(text=str(integer_val)) timeout_slider = ctk.CTkSlider(
timeout_slider = ctk.CTkSlider(self.ai_config_tab, from_=0, to=3600, number_of_steps=3600, command=update_timeout_label, variable=self.timeout_var) self.ai_config_tab,
timeout_slider.grid(row=6, column=1, padx=5, pady=5, sticky="we") from_=0,
self.timeout_value_label = ctk.CTkLabel(self.ai_config_tab, text=str(self.timeout_var.get()), font=("Microsoft YaHei", 12)) to=3600,
self.timeout_value_label.grid(row=6, column=2, padx=5, pady=5, sticky="w") 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 = ctk.CTkButton(
test_btn.grid(row=7, column=0, columnspan=3, padx=5, pady=5, sticky="ew") 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")
# 初始化当前配置
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 build_embeddings_config_tab(self):
def on_embedding_interface_changed(new_value): def on_embedding_interface_changed(new_value):
@@ -199,11 +508,11 @@ def build_embeddings_config_tab(self):
# 1) Embedding API Key # 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)) 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") emb_api_key_entry.grid(row=0, column=1, padx=5, pady=5, sticky="nsew")
# 2) Embedding 接口格式 # 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"] 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 = 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") 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): def load_config_btn(self):
cfg = load_config(self.config_file) cfg = load_config(self.config_file)
if cfg: if cfg:
@@ -239,6 +623,7 @@ def load_config_btn(self):
llm_configs = cfg.get("llm_configs", {}) llm_configs = cfg.get("llm_configs", {})
if last_llm in llm_configs: if last_llm in llm_configs:
llm_conf = llm_configs[last_llm] 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.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.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")) 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(), "model_name": self.model_name_var.get(),
"temperature": self.temperature_var.get(), "temperature": self.temperature_var.get(),
"max_tokens": self.max_tokens_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 = { embedding_config = {
"api_key": self.embedding_api_key_var.get(), "api_key": self.embedding_api_key_var.get(),
"base_url": self.embedding_url_var.get(), "base_url": self.embedding_url_var.get(),
"model_name": self.embedding_model_name_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 = { other_params = {
"topic": self.topic_text.get("0.0", "end").strip(), "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(), "scene_location": self.scene_location_var.get(),
"time_constraint": self.time_constraint_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) existing_config = load_config(self.config_file)
if not existing_config: if not existing_config:
existing_config = {} existing_config = {}
@@ -307,7 +697,9 @@ def save_config_btn(self):
existing_config["last_embedding_interface_format"] = current_embedding_interface existing_config["last_embedding_interface_format"] = current_embedding_interface
if "llm_configs" not in existing_config: if "llm_configs" not in existing_config:
existing_config["llm_configs"] = {} 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: if "embedding_configs" not in existing_config:
existing_config["embedding_configs"] = {} existing_config["embedding_configs"] = {}
+268 -37
View File
@@ -6,6 +6,7 @@ import tkinter as tk
from tkinter import messagebox from tkinter import messagebox
import customtkinter as ctk import customtkinter as ctk
import traceback import traceback
import glob
from utils import read_file, save_string_to_txt, clear_file_content from utils import read_file, save_string_to_txt, clear_file_content
from novel_generator import ( from novel_generator import (
Novel_architecture_generate, Novel_architecture_generate,
@@ -14,7 +15,8 @@ from novel_generator import (
finalize_chapter, finalize_chapter,
import_knowledge_file, import_knowledge_file,
clear_vector_store, clear_vector_store,
enrich_chapter_text enrich_chapter_text,
build_chapter_prompt
) )
from consistency_checker import check_consistency from consistency_checker import check_consistency
@@ -32,13 +34,17 @@ def generate_novel_architecture_ui(self):
self.disable_button_safe(self.btn_generate_architecture) self.disable_button_safe(self.btn_generate_architecture)
try: 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() interface_format = self.loaded_config["llm_configs"][self.architecture_llm_var.get()]["interface_format"]
model_name = self.model_name_var.get().strip() api_key = self.loaded_config["llm_configs"][self.architecture_llm_var.get()]["api_key"]
temperature = self.temperature_var.get() base_url = self.loaded_config["llm_configs"][self.architecture_llm_var.get()]["base_url"]
max_tokens = self.max_tokens_var.get() model_name = self.loaded_config["llm_configs"][self.architecture_llm_var.get()]["model_name"]
timeout_val = self.safe_get_int(self.timeout_var, 600) 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() topic = self.topic_text.get("0.0", "end").strip()
genre = self.genre_var.get().strip() genre = self.genre_var.get().strip()
@@ -82,14 +88,18 @@ def generate_chapter_blueprint_ui(self):
return return
self.disable_button_safe(self.btn_generate_directory) self.disable_button_safe(self.btn_generate_directory)
try: 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) number_of_chapters = self.safe_get_int(self.num_chapters_var, 10)
temperature = self.temperature_var.get()
max_tokens = self.max_tokens_var.get() interface_format = self.loaded_config["llm_configs"][self.chapter_outline_llm_var.get()]["interface_format"]
timeout_val = self.safe_get_int(self.timeout_var, 600) 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() # 新增获取用户指导 user_guidance = self.user_guide_text.get("0.0", "end").strip() # 新增获取用户指导
self.safe_log("开始生成章节蓝图...") self.safe_log("开始生成章节蓝图...")
@@ -121,13 +131,15 @@ def generate_chapter_draft_ui(self):
def task(): def task():
self.disable_button_safe(self.btn_generate_chapter) self.disable_button_safe(self.btn_generate_chapter)
try: try:
interface_format = self.interface_format_var.get().strip()
api_key = self.api_key_var.get().strip() interface_format = self.loaded_config["llm_configs"][self.prompt_draft_llm_var.get()]["interface_format"]
base_url = self.base_url_var.get().strip() api_key = self.loaded_config["llm_configs"][self.prompt_draft_llm_var.get()]["api_key"]
model_name = self.model_name_var.get().strip() base_url = self.loaded_config["llm_configs"][self.prompt_draft_llm_var.get()]["base_url"]
temperature = self.temperature_var.get() model_name = self.loaded_config["llm_configs"][self.prompt_draft_llm_var.get()]["model_name"]
max_tokens = self.max_tokens_var.get() temperature = self.loaded_config["llm_configs"][self.prompt_draft_llm_var.get()]["temperature"]
timeout_val = self.safe_get_int(self.timeout_var, 600) 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) chap_num = self.safe_get_int(self.chapter_num_var, 1)
word_number = self.safe_get_int(self.word_number_var, 3000) 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}章草稿:准备生成请求提示词...") self.safe_log(f"生成第{chap_num}章草稿:准备生成请求提示词...")
# 调用新添加的 build_chapter_prompt 函数构造初始提示词 # 调用新添加的 build_chapter_prompt 函数构造初始提示词
from novel_generator.chapter import build_chapter_prompt
prompt_text = build_chapter_prompt( prompt_text = build_chapter_prompt(
api_key=api_key, api_key=api_key,
base_url=base_url, base_url=base_url,
@@ -312,13 +323,15 @@ def finalize_chapter_ui(self):
self.disable_button_safe(self.btn_finalize_chapter) self.disable_button_safe(self.btn_finalize_chapter)
try: try:
interface_format = self.interface_format_var.get().strip()
api_key = self.api_key_var.get().strip() interface_format = self.loaded_config["llm_configs"][self.final_chapter_llm_var.get()]["interface_format"]
base_url = self.base_url_var.get().strip() api_key = self.loaded_config["llm_configs"][self.final_chapter_llm_var.get()]["api_key"]
model_name = self.model_name_var.get().strip() base_url = self.loaded_config["llm_configs"][self.final_chapter_llm_var.get()]["base_url"]
temperature = self.temperature_var.get() model_name = self.loaded_config["llm_configs"][self.final_chapter_llm_var.get()]["model_name"]
max_tokens = self.max_tokens_var.get() temperature = self.loaded_config["llm_configs"][self.final_chapter_llm_var.get()]["temperature"]
timeout_val = self.safe_get_int(self.timeout_var, 600) 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_api_key = self.embedding_api_key_var.get().strip()
embedding_url = self.embedding_url_var.get().strip() embedding_url = self.embedding_url_var.get().strip()
@@ -392,13 +405,14 @@ def do_consistency_check(self):
def task(): def task():
self.disable_button_safe(self.btn_check_consistency) self.disable_button_safe(self.btn_check_consistency)
try: try:
api_key = self.api_key_var.get().strip() interface_format = self.loaded_config["llm_configs"][self.consistency_review_llm_var.get()]["interface_format"]
base_url = self.base_url_var.get().strip() api_key = self.loaded_config["llm_configs"][self.consistency_review_llm_var.get()]["api_key"]
model_name = self.model_name_var.get().strip() base_url = self.loaded_config["llm_configs"][self.consistency_review_llm_var.get()]["base_url"]
temperature = self.temperature_var.get() model_name = self.loaded_config["llm_configs"][self.consistency_review_llm_var.get()]["model_name"]
interface_format = self.interface_format_var.get() temperature = self.loaded_config["llm_configs"][self.consistency_review_llm_var.get()]["temperature"]
max_tokens = self.max_tokens_var.get() max_tokens = self.loaded_config["llm_configs"][self.consistency_review_llm_var.get()]["max_tokens"]
timeout = self.timeout_var.get() 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_num = self.safe_get_int(self.chapter_num_var, 1)
chap_file = os.path.join(filepath, "chapters", f"chapter_{chap_num}.txt") chap_file = os.path.join(filepath, "chapters", f"chapter_{chap_num}.txt")
@@ -430,6 +444,223 @@ def do_consistency_check(self):
finally: finally:
self.enable_button_safe(self.btn_check_consistency) self.enable_button_safe(self.btn_check_consistency)
threading.Thread(target=task, daemon=True).start() 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): def import_knowledge_handler(self):
selected_file = tk.filedialog.askopenfilename( selected_file = tk.filedialog.askopenfilename(
+11 -1
View File
@@ -54,7 +54,8 @@ def build_left_layout(self):
# Step 按钮区域 # Step 按钮区域
self.step_buttons_frame = ctk.CTkFrame(self.left_frame) 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.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.btn_generate_architecture = ctk.CTkButton(
self.step_buttons_frame, 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_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 = ctk.CTkLabel(self.left_frame, text="输出日志 (只读)", font=("Microsoft YaHei", 12))
log_label.grid(row=3, column=0, padx=5, pady=(5, 0), sticky="w") log_label.grid(row=3, column=0, padx=5, pady=(5, 0), sticky="w")
+44 -14
View File
@@ -26,13 +26,16 @@ from ui.generation_handlers import (
do_consistency_check, do_consistency_check,
import_knowledge_handler, import_knowledge_handler,
clear_vectorstore_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.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.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.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.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.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: class NovelGeneratorGUI:
""" """
@@ -53,23 +56,27 @@ class NovelGeneratorGUI:
self.loaded_config = load_config(self.config_file) self.loaded_config = load_config(self.config_file)
if self.loaded_config: 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") last_embedding = self.loaded_config.get("last_embedding_interface_format", "OpenAI")
else: else:
last_llm = "OpenAI" last_llm = "OpenAI"
last_embedding = "OpenAI" last_embedding = "OpenAI"
if self.loaded_config and "llm_configs" in self.loaded_config and last_llm in self.loaded_config["llm_configs"]: # 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] # llm_conf = next(iter(self.loaded_config["llm_configs"]))
else: # else:
llm_conf = { # llm_conf = {
"api_key": "", # "api_key": "",
"base_url": "https://api.openai.com/v1", # "base_url": "https://api.openai.com/v1",
"model_name": "gpt-4o-mini", # "model_name": "gpt-4o-mini",
"temperature": 0.7, # "temperature": 0.7,
"max_tokens": 8192, # "max_tokens": 8192,
"timeout": 600 # "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"]: 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] emb_conf = self.loaded_config["embedding_configs"][last_embedding]
@@ -82,13 +89,17 @@ class NovelGeneratorGUI:
} }
# -- LLM通用参数 -- # -- 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.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.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.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.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.max_tokens_var = ctk.IntVar(value=llm_conf.get("max_tokens", 8192))
self.timeout_var = ctk.IntVar(value=llm_conf.get("timeout", 600)) 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相关 -- # -- Embedding相关 --
self.embedding_interface_format_var = ctk.StringVar(value=last_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_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.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: if self.loaded_config and "other_params" in self.loaded_config:
op = self.loaded_config["other_params"] op = self.loaded_config["other_params"]
@@ -111,6 +134,10 @@ class NovelGeneratorGUI:
self.scene_location_var = ctk.StringVar(value=op.get("scene_location", "")) self.scene_location_var = ctk.StringVar(value=op.get("scene_location", ""))
self.time_constraint_var = ctk.StringVar(value=op.get("time_constraint", "")) self.time_constraint_var = ctk.StringVar(value=op.get("time_constraint", ""))
self.user_guidance_default = op.get("user_guidance", "") 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: else:
self.topic_default = "" self.topic_default = ""
self.genre_var = ctk.StringVar(value="玄幻") self.genre_var = ctk.StringVar(value="玄幻")
@@ -138,6 +165,8 @@ class NovelGeneratorGUI:
build_character_tab(self) build_character_tab(self)
build_summary_tab(self) build_summary_tab(self)
build_chapters_tab(self) build_chapters_tab(self)
build_other_settings_tab(self)
# ----------------- 通用辅助函数 ----------------- # ----------------- 通用辅助函数 -----------------
def show_tooltip(self, key: str): def show_tooltip(self, key: str):
@@ -344,6 +373,7 @@ class NovelGeneratorGUI:
generate_chapter_draft_ui = generate_chapter_draft_ui generate_chapter_draft_ui = generate_chapter_draft_ui
finalize_chapter_ui = finalize_chapter_ui finalize_chapter_ui = finalize_chapter_ui
do_consistency_check = do_consistency_check do_consistency_check = do_consistency_check
generate_batch_ui = generate_batch_ui
import_knowledge_handler = import_knowledge_handler import_knowledge_handler = import_knowledge_handler
clear_vectorstore_handler = clear_vectorstore_handler clear_vectorstore_handler = clear_vectorstore_handler
show_plot_arcs_ui = show_plot_arcs_ui show_plot_arcs_ui = show_plot_arcs_ui
+277
View File
@@ -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))