diff --git a/config_manager.py b/config_manager.py index c3b45f3..a65be54 100644 --- a/config_manager.py +++ b/config_manager.py @@ -2,6 +2,10 @@ # -*- coding: utf-8 -*- import json import os +import threading +from llm_adapters import create_llm_adapter +from embedding_adapters import create_embedding_adapter + def load_config(config_file: str) -> dict: """从指定的 config_file 加载配置,若不存在则返回空字典。""" @@ -21,3 +25,56 @@ def save_config(config_data: dict, config_file: str) -> bool: return True except: return False + +def test_llm_config(interface_format, api_key, base_url, model_name, temperature, max_tokens, timeout, log_func, handle_exception_func): + """测试当前的LLM配置是否可用""" + def task(): + try: + log_func("开始测试LLM配置...") + llm_adapter = create_llm_adapter( + interface_format=interface_format, + base_url=base_url, + model_name=model_name, + api_key=api_key, + temperature=temperature, + max_tokens=max_tokens, + timeout=timeout + ) + + test_prompt = "Please reply 'OK'" + response = llm_adapter.invoke(test_prompt) + if response: + log_func("✅ LLM配置测试成功!") + log_func(f"测试回复: {response}") + else: + log_func("❌ LLM配置测试失败:未获取到响应") + except Exception as e: + log_func(f"❌ LLM配置测试出错: {str(e)}") + handle_exception_func("测试LLM配置时出错") + + threading.Thread(target=task, daemon=True).start() + +def test_embedding_config(api_key, base_url, interface_format, model_name, log_func, handle_exception_func): + """测试当前的Embedding配置是否可用""" + def task(): + try: + log_func("开始测试Embedding配置...") + embedding_adapter = create_embedding_adapter( + interface_format=interface_format, + api_key=api_key, + base_url=base_url, + model_name=model_name + ) + + test_text = "测试文本" + embeddings = embedding_adapter.embed_query(test_text) + if embeddings and len(embeddings) > 0: + log_func("✅ Embedding配置测试成功!") + log_func(f"生成的向量维度: {len(embeddings)}") + else: + log_func("❌ Embedding配置测试失败:未获取到向量") + except Exception as e: + log_func(f"❌ Embedding配置测试出错: {str(e)}") + handle_exception_func("测试Embedding配置时出错") + + threading.Thread(target=task, daemon=True).start() \ No newline at end of file diff --git a/ui.py b/ui.py index c27964f..7538b97 100644 --- a/ui.py +++ b/ui.py @@ -9,7 +9,7 @@ from tkinter import filedialog, messagebox import tkinter as tk import traceback -from config_manager import load_config, save_config +from config_manager import load_config, save_config, test_llm_config, test_embedding_config from utils import read_file, save_string_to_txt, clear_file_content from novel_generator import ( @@ -525,6 +525,15 @@ class NovelGeneratorGUI: ) self.timeout_value_label.grid(row=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=7, column=0, columnspan=3, padx=5, pady=5, sticky="ew") + # --------------- Embedding 模型配置 --------------- def build_embeddings_config_tab(self): def on_embedding_interface_changed(new_value): @@ -615,6 +624,15 @@ class NovelGeneratorGUI: emb_retrieval_k_entry = ctk.CTkEntry(self.embeddings_config_tab, textvariable=self.embedding_retrieval_k_var, font=("Microsoft YaHei", 12)) emb_retrieval_k_entry.grid(row=4, column=1, padx=5, pady=5, sticky="nsew") + # 添加测试按钮 + 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_novel_params_area(self, start_row=1): """ @@ -1658,6 +1676,48 @@ class NovelGeneratorGUI: else: messagebox.showinfo("提示", "已经是最后一章了。") + def test_llm_config(self): + """ + 测试当前的LLM配置是否可用 + """ + 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 = self.timeout_var.get() + + test_llm_config( + interface_format=interface_format, + api_key=api_key, + base_url=base_url, + model_name=model_name, + temperature=temperature, + max_tokens=max_tokens, + timeout=timeout, + log_func=self.safe_log, + handle_exception_func=self.handle_exception + ) + + def test_embedding_config(self): + """ + 测试当前的Embedding配置是否可用 + """ + api_key = self.embedding_api_key_var.get().strip() + base_url = self.embedding_url_var.get().strip() + interface_format = self.embedding_interface_format_var.get().strip() + model_name = self.embedding_model_name_var.get().strip() + + test_embedding_config( + api_key=api_key, + base_url=base_url, + interface_format=interface_format, + model_name=model_name, + log_func=self.safe_log, + handle_exception_func=self.handle_exception + ) + # ----------------- 程序入口 ----------------- if __name__ == "__main__":