diff --git a/server/routers/base_router.py b/server/routers/base_router.py index 5fc81ae1..b3c746a9 100644 --- a/server/routers/base_router.py +++ b/server/routers/base_router.py @@ -93,7 +93,6 @@ async def update_config_item( async def restart(current_user: User = Depends(get_superadmin_user)): knowledge_base.restart() graph_base.start() - retriever.restart() return {"message": "Restarted!"} @base.get("/log") diff --git a/server/routers/chat_router.py b/server/routers/chat_router.py index ad1e83ca..86626f98 100644 --- a/server/routers/chat_router.py +++ b/server/routers/chat_router.py @@ -10,7 +10,7 @@ from langchain_core.messages import AIMessageChunk, HumanMessage from sqlalchemy.orm import Session from pydantic import BaseModel -from src import executor, config, retriever +from src import executor, config from src.core import HistoryManager from src.agents import agent_manager from src.models import select_model @@ -67,82 +67,6 @@ async def chat_get(current_user: User = Depends(get_required_user)): """聊天服务健康检查(需要登录)""" return "Chat Get!" -@chat.post("/") -async def chat_post( - query: str = Body(...), - meta: dict = Body(None), - history: list[dict] | None = Body(None), - thread_id: str | None = Body(None), - current_user: User = Depends(get_required_user)): - """处理聊天请求的主要端点(需要登录)""" - - model = select_model() - meta["server_model_name"] = model.model_name - history_manager = HistoryManager(history, system_prompt=meta.get("system_prompt")) - logger.debug(f"Received query: {query} with meta: {meta}") - - def make_chunk(content=None, **kwargs): - return json.dumps({ - "response": content, - "meta": meta, - **kwargs - }, ensure_ascii=False).encode('utf-8') + b"\n" - - def need_retrieve(meta): - return meta.get("use_web") or meta.get("use_graph") or meta.get("db_id") - - def generate_response(): - modified_query = query - refs = None - - # 处理知识库检索 - if meta and need_retrieve(meta): - chunk = make_chunk(status="searching") - yield chunk - - try: - modified_query, refs = retriever(modified_query, history_manager.messages, meta) - except Exception as e: - logger.error(f"Retriever error: {e}, {traceback.format_exc()}") - yield make_chunk(message=f"Retriever error: {e}", status="error") - return - - yield make_chunk(status="generating") - - messages = history_manager.get_history_with_msg(modified_query, max_rounds=meta.get('history_round')) - history_manager.add_user(query) # 注意这里使用原始查询 - - content = "" - reasoning_content = "" - try: - for delta in model.predict(messages, stream=True): - if not delta.content and hasattr(delta, 'reasoning_content'): - reasoning_content += delta.reasoning_content or "" - chunk = make_chunk(reasoning_content=reasoning_content, status="reasoning") - yield chunk - continue - - # 文心一言 - if hasattr(delta, 'is_full') and delta.is_full: - content = delta.content - else: - content += delta.content or "" - - chunk = make_chunk(content=delta.content, status="loading") - yield chunk - - logger.debug(f"Final response: {content}") - logger.debug(f"Final reasoning response: {reasoning_content}") - yield make_chunk(status="finished", - history=history_manager.update_ai(content), - refs=refs) - except Exception as e: - logger.error(f"Model error: {e}, {traceback.format_exc()}") - yield make_chunk(message=f"Model error: {e}", status="error") - return - - return StreamingResponse(generate_response(), media_type='application/json') - @chat.post("/call") async def call(query: str = Body(...), meta: dict = Body(None), current_user: User = Depends(get_required_user)): """调用模型进行简单问答(需要登录)""" diff --git a/src/__init__.py b/src/__init__.py index ea1782fc..06abbc9c 100644 --- a/src/__init__.py +++ b/src/__init__.py @@ -13,6 +13,3 @@ knowledge_base = KnowledgeBase() from src.core import GraphDatabase # noqa: E402 graph_base = GraphDatabase() - -from src.core.retriever import Retriever # noqa: E402 -retriever = Retriever() diff --git a/src/config/__init__.py b/src/config/__init__.py index 37cc762d..c9efd1f6 100644 --- a/src/config/__init__.py +++ b/src/config/__init__.py @@ -48,8 +48,6 @@ class Config(SimpleConfig): ### >>> 默认配置 # 功能选项 self.add_item("enable_reranker", default=False, des="是否开启重排序") - self.add_item("enable_knowledge_base", default=False, des="是否开启知识库") - self.add_item("enable_knowledge_graph", default=False, des="是否开启知识图谱") self.add_item("enable_web_search", default=False, des="是否开启网页搜索(注:现阶段会根据 TAVILY_API_KEY 自动开启,无法手动配置,将会在下个版本移除此配置项)") # noqa: E501 # 默认智能体配置 self.add_item("default_agent_id", default="", des="默认智能体ID") @@ -57,7 +55,7 @@ class Config(SimpleConfig): ## 注意这里是模型名,而不是具体的模型路径,默认使用 HuggingFace 的路径 ## 如果需要自定义本地模型路径,则在 src/.env 中配置 MODEL_DIR self.add_item("model_provider", default="siliconflow", des="模型提供商", choices=list(self.model_names.keys())) - self.add_item("model_name", default="Qwen/Qwen2.5-7B-Instruct", des="模型名称") + self.add_item("model_name", default="Qwen/Qwen3-32B", des="模型名称") self.add_item("embed_model", default="siliconflow/BAAI/bge-m3", des="Embedding 模型", choices=list(self.embed_model_names.keys())) self.add_item("reranker", default="siliconflow/BAAI/bge-reranker-v2-m3", des="Re-Ranker 模型", choices=list(self.reranker_names.keys())) # noqa: E501 diff --git a/src/core/graphbase.py b/src/core/graphbase.py index 9f923aed..f45dd446 100644 --- a/src/core/graphbase.py +++ b/src/core/graphbase.py @@ -31,9 +31,6 @@ class GraphDatabase: self.start() def start(self): - if not config.enable_knowledge_graph or not config.enable_knowledge_base: - return - uri = os.environ.get("NEO4J_URI", "bolt://localhost:7687") username = os.environ.get("NEO4J_USERNAME", "neo4j") password = os.environ.get("NEO4J_PASSWORD", "0123456789") @@ -46,7 +43,6 @@ class GraphDatabase: self.save_graph_info(self.kgdb_name) except Exception as e: logger.error(f"Failed to connect to Neo4j: {e}, {uri}, {self.kgdb_name}, {username}, {password}") - self.config.enable_knowledge_graph = False def close(self): """关闭数据库连接""" @@ -54,10 +50,7 @@ class GraphDatabase: def is_running(self): """检查图数据库是否正在运行""" - if not config.enable_knowledge_graph or not config.enable_knowledge_base: - return False - else: - return self.status == "open" + return self.status == "open" def get_sample_nodes(self, kgdb_name='neo4j', num=50): """获取指定数据库的 num 个节点信息""" diff --git a/src/core/retriever.py b/src/core/retriever.py deleted file mode 100644 index a74c772b..00000000 --- a/src/core/retriever.py +++ /dev/null @@ -1,182 +0,0 @@ -import traceback - -from src import config, knowledge_base, graph_base -from src.models.rerank_model import get_reranker -from src.utils.logging_config import logger -from src.models import select_model -from src.core.operators import HyDEOperator - -class Retriever: - - def __init__(self): - self._load_models() - - def _load_models(self): - if config.enable_reranker: - self.reranker = get_reranker() - - if config.enable_web_search: - from src.utils.web_search import WebSearcher - self.web_searcher = WebSearcher() - - def retrieval(self, query, history, meta): - refs = {"query": query, "history": history, "meta": meta} - refs["model_name"] = config.model_name - refs["entities"] = self.reco_entities(query, history, refs) - refs["knowledge_base"] = self.query_knowledgebase(query, history, refs) - refs["graph_base"] = self.query_graph(query, history, refs) - refs["web_search"] = self.query_web(query, history, refs) - - return refs - - def restart(self): - """所有需要重启的模型""" - self._load_models() - - def construct_query(self, query, refs, meta): - logger.debug(f"{refs=}") - if not refs or len(refs) == 0: - return query - - external_parts = [] - - # 解析知识库的结果 - kb_res = refs.get("knowledge_base", {}).get("results", []) - if kb_res: - kb_text = "\n".join(f"{r['id']}: {r['entity']['text']}" for r in kb_res) - external_parts.extend(["知识库信息:", kb_text]) - - # 解析图数据库的结果 - db_res = refs.get("graph_base", {}).get("results", {}) - if db_res.get("nodes") and len(db_res["nodes"]) > 0: - db_text = "\n".join( - [f"{edge['source_name']}和{edge['target_name']}的关系是{edge['type']}" for edge in db_res.get("edges", [])] - ) - external_parts.extend(["图数据库信息:", db_text]) - - # 解析网络搜索的结果 - web_res = refs.get("web_search", {}).get("results", []) - if web_res: - web_text = "\n".join(f"{r['title']}: {r['content']}" for r in web_res) - external_parts.extend(["网络搜索信息:", web_text]) - - # 构造查询 - from src.utils.prompts import knowbase_qa_template - if external_parts and len(external_parts) > 0: - external = "\n\n".join(external_parts) - query = knowbase_qa_template.format(external=external, query=query) - - return query - - def query_classification(self, query): - """判断是否需要查询 - - 对于完全基于用户给定信息的任务,称之为"足够""sufficient",不需要检索; - - 否则,称之为"不足""insufficient",可能需要检索, - """ - raise NotImplementedError - - def query_graph(self, query, history, refs): - results = [] - if refs["meta"].get("use_graph") and config.enable_knowledge_base: - for entity in refs["entities"]: - if entity == "": - continue - result = graph_base.query_node(entity) - if result != []: - results.extend(result) - return {"results": graph_base.format_query_result_to_graph(results)} - - - def query_knowledgebase(self, query, history, refs): - """查询知识库""" - - response = { - "results": [], - "all_results": [], - "rw_query": query, - "message": "", - } - - meta = refs["meta"] - - db_id = meta.get("db_id") - if not db_id or not config.enable_knowledge_base: - response["message"] = "知识库未启用、或未指定知识库、或知识库不存在" - return response - - rw_query = self.rewrite_query(query, history, refs) - - logger.debug(f"{meta=}") - query_result = knowledge_base.query(query_text=rw_query, - db_id=db_id, - distance_threshold=meta.get("distanceThreshold", 0.5), - rerank_threshold=meta.get("rerankThreshold", 0.1), - max_query_count=meta.get("maxQueryCount", 20), - top_k=meta.get("topK", 10)) - - response["results"] = query_result["results"] - response["all_results"] = query_result["all_results"] - response["rw_query"] = rw_query - - return response - - def query_web(self, query, history, refs): - """查询网络""" - - if not (refs["meta"].get("use_web") or not config.enable_web_search): - return {"results": [], "message": "Web search is disabled"} - - try: - search_results = self.web_searcher.search(query, max_results=5) - except Exception as e: - logger.error(f"Web search error: {str(e)}") - return {"results": [], "message": "Web search error"} - - return {"results": search_results} - - def rewrite_query(self, query, history, refs): - """重写查询""" - model_provider = config.model_provider - model_name = config.model_name - model = select_model(model_provider=model_provider, model_name=model_name) - if refs["meta"].get("mode") == "search": # 如果是搜索模式,就使用 meta 的配置,否则就使用全局的配置 - rewrite_query_span = refs["meta"].get("use_rewrite_query", "off") - else: - rewrite_query_span = config.use_rewrite_query - - if rewrite_query_span == "off": - return query - - from src.utils.prompts import rewritten_query_prompt_template2 as rw_template - history_query = [entry["content"] for entry in history if entry["role"] == "user"] if history else "" - rewritten_query_prompt = rw_template.format(history=history_query, query=query) - rewritten_query = model.predict(rewritten_query_prompt).content - - if rewrite_query_span == "hyde": - res = HyDEOperator.call(model_callable=model.predict, query=query, context_str=history_query) - rewritten_query = res.content - - return rewritten_query - - def reco_entities(self, query, history, refs): - """识别句子中的实体""" - query = refs.get("rewritten_query", query) - model_provider = config.model_provider - model_name = config.model_name - model = select_model(model_provider=model_provider, model_name=model_name) - - entities = [] - if refs["meta"].get("use_graph"): - from src.utils.prompts import entity_extraction_prompt_template as entity_template - # from src.utils.prompts import keywords_prompt_template as entity_templat|e - - entity_extraction_prompt = entity_template.format(text=query) - entities = model.predict(entity_extraction_prompt).content.split("<->") - # entities = [entity for entity in entities if all(char.isalnum() or char in "汉字" for char in entity)] - - return entities - - def __call__(self, query, history, meta): - refs = self.retrieval(query, history, meta) - query = self.construct_query(query, refs, meta) - return query, refs diff --git a/src/models/embedding.py b/src/models/embedding.py index 6718c887..e40d9fff 100644 --- a/src/models/embedding.py +++ b/src/models/embedding.py @@ -182,9 +182,6 @@ class OtherEmbedding(BaseEmbeddingModel): } def get_embedding_model(): - if not config.enable_knowledge_base: - return None - provider, model_name = config.embed_model.split('/', 1) support_embed_models = config.embed_model_names.keys() assert config.embed_model in support_embed_models, f"Unsupported embed model: {config.embed_model}, only support {support_embed_models}" diff --git a/web/src/components/ChatComponent.vue b/web/src/components/ChatComponent.vue index 123d2d5c..1f4510b8 100644 --- a/web/src/components/ChatComponent.vue +++ b/web/src/components/ChatComponent.vue @@ -98,14 +98,13 @@