From 5453452452136d754161899cd22f7794243dc2f1 Mon Sep 17 00:00:00 2001 From: littlewwwhite <1095245867@qq.com> Date: Sat, 15 Feb 2025 19:29:32 +0800 Subject: [PATCH] =?UTF-8?q?feat:=20=E6=B7=BB=E5=8A=A0=E4=BA=86=E5=9F=BA?= =?UTF-8?q?=E4=BA=8E=20tavil=E7=9A=84web=20=E6=90=9C=E7=B4=A2?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .gitignore | 5 +- requirements.txt | 3 +- src/config/__init__.py | 2 +- src/config/models.yaml | 1 + src/core/startup.py | 22 ++++- src/models/__init__.py | 59 ++++-------- src/models/chat_model.py | 13 ++- src/routers/chat_router.py | 42 +++++++-- web/src/components/ChatComponent.vue | 132 +++++++++++++++++---------- 9 files changed, 179 insertions(+), 100 deletions(-) diff --git a/.gitignore b/.gitignore index df19050c..439e3613 100644 --- a/.gitignore +++ b/.gitignore @@ -35,4 +35,7 @@ web/package-lock.json saves notebooks graphrag -docker/volumes \ No newline at end of file +docker/volumes + + +.cursorrules diff --git a/requirements.txt b/requirements.txt index ed96727a..c7f87bf4 100644 --- a/requirements.txt +++ b/requirements.txt @@ -24,4 +24,5 @@ opencv-python-headless docx2txt uvicorn[standard] fastapi -python-multipart \ No newline at end of file +python-multipart +tavily-python \ No newline at end of file diff --git a/src/config/__init__.py b/src/config/__init__.py index 5be0e905..2d82e268 100644 --- a/src/config/__init__.py +++ b/src/config/__init__.py @@ -53,7 +53,7 @@ class Config(SimpleConfig): self.add_item("enable_knowledge_base", default=False, des="是否开启知识库") self.add_item("enable_knowledge_graph", default=False, des="是否开启知识图谱") self.add_item("enable_search_engine", default=False, des="是否开启搜索引擎") - + self.add_item("enable_web_search", default=False, des="是否开启网页搜索") # 模型配置 ## 注意这里是模型名,而不是具体的模型路径,默认使用 HuggingFace 的路径 ## 如果需要自定义路径,则在 config/base.yaml 中配置 model_local_paths diff --git a/src/config/models.yaml b/src/config/models.yaml index f42781c2..55dd90f0 100644 --- a/src/config/models.yaml +++ b/src/config/models.yaml @@ -18,6 +18,7 @@ MODEL_NAMES: - DEEPSEEK_API_KEY models: - deepseek-chat + - deepseek-reasoner zhipu: name: 智谱AI (Zhipu) url: https://open.bigmodel.cn/dev/api diff --git a/src/core/startup.py b/src/core/startup.py index 943dd10b..ee4f95dd 100644 --- a/src/core/startup.py +++ b/src/core/startup.py @@ -1,3 +1,4 @@ +import os from src.core import DataBaseManager from src.core.retriever import Retriever from src.models import select_model @@ -9,10 +10,29 @@ logger = setup_logger("Startup") class Startup: def __init__(self): + self.config = Config("config/base.yaml") + self._check_environment() self.start() + def _check_environment(self): + """检查必要的环境变量""" + required_vars = { + "zhipu": ["ZHIPUAI_API_KEY"], + "openai": ["OPENAI_API_KEY"], + "deepseek": ["DEEPSEEK_API_KEY"], + } + + provider = self.config.model_provider + if provider in required_vars: + missing = [var for var in required_vars[provider] if not os.getenv(var)] + if missing: + logger.error(f"Missing required environment variables for {provider}: {missing}") + raise ValueError(f"Missing required environment variables: {missing}") + + if self.config.enable_web_search and not os.getenv("TAVILY_API_KEY"): + logger.warning("TAVILY_API_KEY not set, web search will be disabled") + def start(self): - self.config = Config() self.model = select_model(self.config) self.dbm = DataBaseManager(self.config) self.retriever = Retriever(self.config, self.dbm, self.model) diff --git a/src/models/__init__.py b/src/models/__init__.py index b5efa9d0..4f7a3d39 100644 --- a/src/models/__init__.py +++ b/src/models/__init__.py @@ -1,46 +1,23 @@ from src.utils.logging_config import logger +from src.models.chat_model import OpenModel, DeepSeek, Zhipu, Qianfan, DashScope, SiliconFlow +from src.models.embedding import get_embedding_model def select_model(config): - - model_provider = config.model_provider - model_name = config.model_name - - logger.info(f"Selecting model from {model_provider} with {model_name}") - - if model_provider == "deepseek": - from src.models.chat_model import DeepSeek - return DeepSeek(model_name) - - elif model_provider == "zhipu": - from src.models.chat_model import Zhipu - return Zhipu(model_name) - - elif model_provider == "qianfan": - from src.models.chat_model import Qianfan - return Qianfan(model_name) - - elif model_provider == "dashscope": - from src.models.chat_model import DashScope - return DashScope(model_name) - - elif model_provider == "openai": - from src.models.chat_model import OpenModel - return OpenModel(model_name) - - elif model_provider == "siliconflow": - from src.models.chat_model import SiliconFlow - return SiliconFlow(model_name) - - elif model_provider == "custom": - model_info = next((x for x in config.custom_models if x["custom_id"] == model_name), None) - if model_info is None: - raise ValueError(f"Model {model_name} not found in custom models") - - from src.models.chat_model import CustomModel - return CustomModel(model_info) - - elif model_provider is None: - raise ValueError("Model provider not specified, please modify `model_provider` in `src/config/base.yaml`") + """ + 根据配置选择模型 + """ + if config.model_provider == "deepseek": + return DeepSeek(config.model_name) + elif config.model_provider == "zhipu": + return Zhipu(config.model_name) + elif config.model_provider == "openai": + return OpenModel(config.model_name) + elif config.model_provider == "qianfan": + return Qianfan(config.model_name) + elif config.model_provider == "dashscope": + return DashScope(config.model_name) + elif config.model_provider == "siliconflow": + return SiliconFlow(config.model_name) else: - raise ValueError(f"Model provider {model_provider} not supported") + raise ValueError(f"Unsupported model provider: {config.model_provider}") diff --git a/src/models/chat_model.py b/src/models/chat_model.py index 8f1924e7..bcf331b3 100644 --- a/src/models/chat_model.py +++ b/src/models/chat_model.py @@ -1,6 +1,7 @@ import os from openai import OpenAI from src.utils.logging_config import setup_logger +from zhipuai import ZhipuAI logger = setup_logger(__name__) @@ -50,8 +51,8 @@ class OpenModel(OpenAIBase): class DeepSeek(OpenAIBase): def __init__(self, model_name=None): model_name = model_name or "deepseek-chat" - api_key = os.getenv("DEEPSEEK_API_KEY") - base_url = "https://api.deepseek.com" + api_key = os.getenv("DEEPSEEK_API_KEY", "your-default-api-key") + base_url = os.getenv("DEEPSEEK_API_BASE", "https://api.deepseek.com/v1") super().__init__(api_key=api_key, base_url=base_url, model_name=model_name) @@ -165,6 +166,14 @@ class DashScope: return response.output.choices[0].message +class ChatModel: + def __init__(self, config): + if config.model_provider == "zhipu": + self.client = ZhipuAI(api_key=os.getenv("ZHIPUAI_API_KEY")) + elif config.model_provider == "openai": + self.client = OpenAI(api_key=os.getenv("OPENAI_API_KEY")) + + if __name__ == "__main__": model = SiliconFlow() for a in model.predict("你好", stream=True): diff --git a/src/routers/chat_router.py b/src/routers/chat_router.py index 0f5659ef..ad136dde 100644 --- a/src/routers/chat_router.py +++ b/src/routers/chat_router.py @@ -6,6 +6,7 @@ from concurrent.futures import ThreadPoolExecutor from src.core import HistoryManager from src.core.startup import startup from src.utils.logging_config import setup_logger +from src.utils.web_search import WebSearcher chat = APIRouter(prefix="/chat") logger = setup_logger("server-chat") @@ -13,6 +14,7 @@ logger = setup_logger("server-chat") executor = ThreadPoolExecutor() refs_pool = {} +web_searcher = WebSearcher() @chat.get("/") async def chat_get(): @@ -37,19 +39,45 @@ def chat_post( }, ensure_ascii=False).encode('utf-8') + b"\n" def generate_response(): + modified_query = query + + # 处理网页搜索 + if meta and meta.get("enable_web_search"): + chunk = make_chunk("正在进行网络搜索...", "searching", history=None) + yield chunk - if meta.get("enable_retrieval"): + try: + search_results = web_searcher.search(query) + if search_results: + search_context = web_searcher.format_search_results(search_results) + # 将搜索结果添加到查询中 + modified_query = f"""基于以下网络搜索结果回答问题: + + {search_context} + + 用户问题:{query} + + 请综合以上搜索结果,给出准确、客观的回答。如果搜索结果与问题相关性不大,请直接基于你的知识回答。 + """ + logger.info(f"Web search results added to query") + else: + logger.warning("No web search results found") + except Exception as e: + logger.error(f"Web search error: {str(e)}") + chunk = make_chunk("网络搜索失败,将直接回答问题。", "loading", history=None) + yield chunk + + # 处理知识库检索 + if meta and meta.get("enable_retrieval"): chunk = make_chunk("", "searching", history=None) yield chunk - new_query, refs = startup.retriever(query, history_manager.messages, meta) + modified_query, refs = startup.retriever(modified_query, history_manager.messages, meta) refs_pool[cur_res_id] = refs - else: - new_query = query - messages = history_manager.get_history_with_msg(new_query, max_rounds=meta.get('history_round')) - history_manager.add_user(query) - logger.debug(f"Web history: {history_manager.messages}") + messages = history_manager.get_history_with_msg(modified_query, max_rounds=meta.get('history_round')) + history_manager.add_user(query) # 注意这里使用原始查询 + logger.debug(f"Final query: {modified_query}") content = "" for delta in startup.model.predict(messages, stream=True): diff --git a/web/src/components/ChatComponent.vue b/web/src/components/ChatComponent.vue index 13b37d58..165395f7 100644 --- a/web/src/components/ChatComponent.vue +++ b/web/src/components/ChatComponent.vue @@ -81,6 +81,9 @@