From 9a162d5197ff7acbf731b09505ee2a332aeb665f Mon Sep 17 00:00:00 2001 From: Wenjie Zhang Date: Wed, 2 Jul 2025 02:58:13 +0800 Subject: [PATCH] =?UTF-8?q?feat:=20=E6=9B=B4=E6=96=B0=E7=8E=AF=E5=A2=83?= =?UTF-8?q?=E6=A8=A1=E6=9D=BF=EF=BC=8C=E6=B7=BB=E5=8A=A0=E7=9F=A5=E8=AF=86?= =?UTF-8?q?=E5=9B=BE=E8=B0=B1=E6=A8=A1=E5=9E=8B=E5=90=8D=E7=A7=B0=E9=85=8D?= =?UTF-8?q?=E7=BD=AE=EF=BC=9B=E9=87=8D=E6=9E=84=E5=B7=A5=E5=85=B7=E5=B7=A5?= =?UTF-8?q?=E5=8E=82=EF=BC=8C=E7=A7=BB=E9=99=A4=E5=86=97=E4=BD=99=E4=BB=A3?= =?UTF-8?q?=E7=A0=81=E5=B9=B6=E4=BC=98=E5=8C=96=E5=B7=A5=E5=85=B7=E6=B3=A8?= =?UTF-8?q?=E5=86=8C=E9=80=BB=E8=BE=91=EF=BC=9B=E5=A2=9E=E5=BC=BA=E5=9B=BE?= =?UTF-8?q?=E6=95=B0=E6=8D=AE=E5=BA=93=E7=B1=BB=EF=BC=8C=E6=B7=BB=E5=8A=A0?= =?UTF-8?q?=E6=95=B0=E6=8D=AE=E5=BA=93=E8=BF=9E=E6=8E=A5=E6=A3=80=E6=9F=A5?= =?UTF-8?q?=E5=92=8C=E5=B5=8C=E5=85=A5=E6=A8=A1=E5=9E=8B=E8=8E=B7=E5=8F=96?= =?UTF-8?q?=E5=8A=9F=E8=83=BD=EF=BC=9B=E8=B0=83=E6=95=B4=E8=AE=BE=E7=BD=AE?= =?UTF-8?q?=E8=A7=86=E5=9B=BE=EF=BC=8C=E7=A7=BB=E9=99=A4=E4=B8=8D=E5=BF=85?= =?UTF-8?q?=E8=A6=81=E7=9A=84=E5=BC=80=E5=85=B3=E5=92=8C=E9=80=89=E6=8B=A9?= =?UTF-8?q?=E9=A1=B9?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- src/.env.template | 3 ++ src/agents/tools_factory.py | 68 +++-------------------------------- src/core/graphbase.py | 36 +++++++++++++------ web/src/views/SettingView.vue | 20 ----------- 4 files changed, 34 insertions(+), 93 deletions(-) diff --git a/src/.env.template b/src/.env.template index b685c22f..1a712494 100644 --- a/src/.env.template +++ b/src/.env.template @@ -13,3 +13,6 @@ TOGETHER_API_KEY=tgp_v1_fPjW******irD6zesAn4 # 功能服务 TAVILY_API_KEY=tvly-3gR4ind9******JMOxw3E2LG # <<< 配置网络搜索 + +# 知识图谱的模型,格式参考 models.yaml +GRAPH_EMBED_MODEL_NAME="" \ No newline at end of file diff --git a/src/agents/tools_factory.py b/src/agents/tools_factory.py index feb75d03..3f3a8ee4 100644 --- a/src/agents/tools_factory.py +++ b/src/agents/tools_factory.py @@ -1,70 +1,14 @@ import json -import re +import asyncio from collections.abc import Callable from typing import Annotated, Any -import asyncio -import logging -from langchain_tavily import TavilySearch -from langchain_core.tools import BaseTool, StructuredTool, tool from pydantic import BaseModel, Field +from langchain_core.tools import StructuredTool, tool +from langchain_tavily import TavilySearch from src import config, graph_base, knowledge_base - - -# refs https://github.com/chatchat-space/LangGraph-Chatchat chatchat-server/chatchat/server/agent/tools_factory/tools_registry.py -def regist_tool( - *args: Any, - title: str = "", - description: str = "", - return_direct: bool = False, - args_schema: type[BaseModel] | None = None, - infer_schema: bool = True, -) -> Callable | BaseTool: - """ - wrapper of langchain tool decorator - add tool to registry automatically - """ - - def _parse_tool(t: BaseTool): - nonlocal description, title - - _TOOLS_REGISTRY[t.name] = t - - # change default description - if not description: - if t.func is not None: - description = t.func.__doc__ - elif t.coroutine is not None: - description = t.coroutine.__doc__ - t.description = " ".join(re.split(r"\n+\s*", description)) - # set a default title for human - if not title: - title = "".join([x.capitalize() for x in t.name.split("_")]) - setattr(t, "_title", title) - - def wrapper(def_func: Callable) -> BaseTool: - partial_ = tool( - *args, - return_direct=return_direct, - args_schema=args_schema, - infer_schema=infer_schema, - ) - t = partial_(def_func) - _parse_tool(t) - return t - - if len(args) == 0: - return wrapper - else: - t = tool( - *args, - return_direct=return_direct, - args_schema=args_schema, - infer_schema=infer_schema, - ) - _parse_tool(t) - return t +from src.utils import logger class KnowledgeRetrieverModel(BaseModel): @@ -75,8 +19,6 @@ class KnowledgeRetrieverModel(BaseModel): ) ) - - def get_all_tools(): """获取所有工具""" tools = _TOOLS_REGISTRY.copy() @@ -124,7 +66,7 @@ class BaseToolOutput: def __init__( self, data: Any, - format: str | Callable = None, + format: str | Callable | None = None, data_alias: str = "", **extras: Any, ) -> None: diff --git a/src/core/graphbase.py b/src/core/graphbase.py index 21a0a1ae..3a5fcd5a 100644 --- a/src/core/graphbase.py +++ b/src/core/graphbase.py @@ -4,8 +4,10 @@ import warnings import traceback from neo4j import GraphDatabase as GD +from neo4j import Query from src import config +from src.models.embedding import get_embedding_model from src.utils import logger warnings.filterwarnings("ignore", category=UserWarning) @@ -19,7 +21,8 @@ class GraphDatabase: self.files = [] self.status = "closed" self.kgdb_name = "neo4j" - self.embed_model_name = None + self.embed_model_name = os.getenv("GRAPH_EMBED_MODEL_NAME") or "siliconflow/BAAI/bge-m3" + self.embed_model = get_embedding_model(self.embed_model_name) self.work_dir = os.path.join(config.save_dir, "knowledge_graph", self.kgdb_name) os.makedirs(self.work_dir, exist_ok=True) @@ -45,6 +48,7 @@ class GraphDatabase: def close(self): """关闭数据库连接""" + assert self.driver is not None, "Database is not connected" self.driver.close() def is_running(self): @@ -53,6 +57,7 @@ class GraphDatabase: def get_sample_nodes(self, kgdb_name='neo4j', num=50): """获取指定数据库的 num 个节点信息""" + assert self.driver is not None, "Database is not connected" self.use_database(kgdb_name) def query(tx, num): result = tx.run("MATCH (n)-[r]->(m) RETURN n, r, m LIMIT $num", num=int(num)) @@ -63,6 +68,7 @@ class GraphDatabase: def create_graph_database(self, kgdb_name): """创建新的数据库,如果已存在则返回已有数据库的名称""" + assert self.driver is not None, "Database is not connected" with self.driver.session() as session: existing_databases = session.run("SHOW DATABASES") existing_db_names = [db['name'] for db in existing_databases] @@ -71,7 +77,7 @@ class GraphDatabase: print(f"已存在数据库: {existing_db_names[0]}") return existing_db_names[0] # 返回所有已有数据库名称 - session.run(f"CREATE DATABASE {kgdb_name}") + session.run(f"CREATE DATABASE {kgdb_name}") # type: ignore print(f"数据库 '{kgdb_name}' 创建成功.") return kgdb_name # 返回创建的数据库名称 @@ -83,6 +89,7 @@ class GraphDatabase: def txt_add_entity(self, triples, kgdb_name='neo4j'): """添加实体三元组""" + assert self.driver is not None, "Database is not connected" self.use_database(kgdb_name) def create(tx, triples): for triple in triples: @@ -101,6 +108,7 @@ class GraphDatabase: async def txt_add_vector_entity(self, triples, kgdb_name='neo4j'): """添加实体三元组""" + assert self.driver is not None, "Database is not connected" self.use_database(kgdb_name) def _index_exists(tx, index_name): """检查索引是否存在""" @@ -211,6 +219,7 @@ class GraphDatabase: self.save_graph_info() async def jsonl_file_add_entity(self, file_path, kgdb_name='neo4j'): + assert self.driver is not None, "Database is not connected" self.status = "processing" kgdb_name = kgdb_name or 'neo4j' self.use_database(kgdb_name) # 切换到指定数据库 @@ -233,6 +242,7 @@ class GraphDatabase: def delete_entity(self, entity_name=None, kgdb_name="neo4j"): """删除数据库中的指定实体三元组, 参数entity_name为空则删除全部实体""" + assert self.driver is not None, "Database is not connected" self.use_database(kgdb_name) with self.driver.session() as session: if entity_name: @@ -256,6 +266,7 @@ class GraphDatabase: def query_node(self, entity_name, threshold=0.9, kgdb_name='neo4j', hops=2, max_entities=5, **kwargs): """知识图谱查询节点的入口:""" + assert self.driver is not None, "Database is not connected" # TODO 添加判断节点数量为 0 停止检索 # 判断是否启动 if not self.is_running(): @@ -306,6 +317,7 @@ class GraphDatabase: def query_specific_entity(self, entity_name, kgdb_name='neo4j', hops=2, limit=100): """查询指定实体三元组信息(无向关系)""" + assert self.driver is not None, "Database is not connected" if not entity_name: logger.warning("实体名称为空") return [] @@ -343,6 +355,7 @@ class GraphDatabase: def query_all_nodes_and_relationships(self, kgdb_name='neo4j', hops = 2): """查询图数据库中所有三元组信息 NEVER USE""" + assert self.driver is not None, "Database is not connected" self.use_database(kgdb_name) def query(tx, hops): result = tx.run(f""" @@ -358,6 +371,7 @@ class GraphDatabase: def query_by_relationship_type(self, relationship_type, kgdb_name='neo4j', hops = 2): """查询指定关系三元组信息 NEVER USE""" + assert self.driver is not None, "Database is not connected" self.use_database(kgdb_name) def query(tx, relationship_type, hops): result = tx.run(f""" @@ -373,6 +387,7 @@ class GraphDatabase: def query_entity_like(self, keyword, kgdb_name='neo4j', hops = 2): """模糊查询 NEVER USE""" + assert self.driver is not None, "Database is not connected" self.use_database(kgdb_name) def query(tx, keyword, hops): result = tx.run(f""" @@ -390,6 +405,7 @@ class GraphDatabase: def query_node_info(self, node_name, kgdb_name='neo4j', hops = 2): """查询指定节点的详细信息返回信息 NEVER USE""" + assert self.driver is not None, "Database is not connected" self.use_database(kgdb_name) # 切换到指定数据库 def query(tx, node_name, hops): result = tx.run(f""" @@ -405,23 +421,19 @@ class GraphDatabase: return session.execute_read(query, node_name, hops) async def aget_embedding(self, text): - from src import knowledge_base - if isinstance(text, list): - outputs = await knowledge_base.embed_model.abatch_encode(text, batch_size=40) + outputs = await self.embed_model.abatch_encode(text, batch_size=40) return outputs else: - outputs = await knowledge_base.embed_model.aencode(text) + outputs = await self.embed_model.aencode(text) return outputs def get_embedding(self, text): - from src import knowledge_base - if isinstance(text, list): - outputs = knowledge_base.embed_model.batch_encode(text, batch_size=40) + outputs = self.embed_model.batch_encode(text, batch_size=40) return outputs else: - outputs = knowledge_base.embed_model.encode([text])[0] + outputs = self.embed_model.encode([text])[0] return outputs def set_embedding(self, tx, entity_name, embedding): @@ -431,6 +443,7 @@ class GraphDatabase: """, name=entity_name, embedding=embedding) def get_graph_info(self, graph_name="neo4j"): + assert self.driver is not None, "Database is not connected" self.use_database(graph_name) def query(tx): entity_count = tx.run("MATCH (n) RETURN count(n) AS count").single()["count"] @@ -493,6 +506,7 @@ class GraphDatabase: Returns: list: 没有嵌入向量的节点列表 """ + assert self.driver is not None, "Database is not connected" self.use_database(kgdb_name) def query(tx): @@ -543,6 +557,7 @@ class GraphDatabase: Returns: int: 成功添加嵌入向量的节点数量 """ + assert self.driver is not None, "Database is not connected" self.use_database(kgdb_name) # 如果node_names为None,则获取所有没有嵌入向量的节点 @@ -575,6 +590,7 @@ class GraphDatabase: source_id = source.element_id target_id = target.element_id + assert node_dict is not None, "node_dict is required" source_name = node_dict[source_id]["name"] if source_name is None else source_name target_name = node_dict[target_id]["name"] if target_name is None else target_name diff --git a/web/src/views/SettingView.vue b/web/src/views/SettingView.vue index e6fae9b0..be99ca85 100644 --- a/web/src/views/SettingView.vue +++ b/web/src/views/SettingView.vue @@ -43,7 +43,6 @@ -
- {{ items?.enable_reranker.des }} - -
-
- {{ items?.use_rewrite_query.des }} - - {{ name }} - - -