feat: 更新环境模板,添加知识图谱模型名称配置;重构工具工厂,移除冗余代码并优化工具注册逻辑;增强图数据库类,添加数据库连接检查和嵌入模型获取功能;调整设置视图,移除不必要的开关和选择项

This commit is contained in:
Wenjie Zhang 2025-07-02 02:58:13 +08:00
parent 6810f1f84e
commit 9a162d5197
4 changed files with 34 additions and 93 deletions

View File

@ -13,3 +13,6 @@ TOGETHER_API_KEY=tgp_v1_fPjW******irD6zesAn4
# 功能服务 # 功能服务
TAVILY_API_KEY=tvly-3gR4ind9******JMOxw3E2LG # <<< 配置网络搜索 TAVILY_API_KEY=tvly-3gR4ind9******JMOxw3E2LG # <<< 配置网络搜索
# 知识图谱的模型,格式参考 models.yaml
GRAPH_EMBED_MODEL_NAME=""

View File

@ -1,70 +1,14 @@
import json import json
import re import asyncio
from collections.abc import Callable from collections.abc import Callable
from typing import Annotated, Any 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 pydantic import BaseModel, Field
from langchain_core.tools import StructuredTool, tool
from langchain_tavily import TavilySearch
from src import config, graph_base, knowledge_base from src import config, graph_base, knowledge_base
from src.utils import logger
# 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
class KnowledgeRetrieverModel(BaseModel): class KnowledgeRetrieverModel(BaseModel):
@ -75,8 +19,6 @@ class KnowledgeRetrieverModel(BaseModel):
) )
) )
def get_all_tools(): def get_all_tools():
"""获取所有工具""" """获取所有工具"""
tools = _TOOLS_REGISTRY.copy() tools = _TOOLS_REGISTRY.copy()
@ -124,7 +66,7 @@ class BaseToolOutput:
def __init__( def __init__(
self, self,
data: Any, data: Any,
format: str | Callable = None, format: str | Callable | None = None,
data_alias: str = "", data_alias: str = "",
**extras: Any, **extras: Any,
) -> None: ) -> None:

View File

@ -4,8 +4,10 @@ import warnings
import traceback import traceback
from neo4j import GraphDatabase as GD from neo4j import GraphDatabase as GD
from neo4j import Query
from src import config from src import config
from src.models.embedding import get_embedding_model
from src.utils import logger from src.utils import logger
warnings.filterwarnings("ignore", category=UserWarning) warnings.filterwarnings("ignore", category=UserWarning)
@ -19,7 +21,8 @@ class GraphDatabase:
self.files = [] self.files = []
self.status = "closed" self.status = "closed"
self.kgdb_name = "neo4j" 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) self.work_dir = os.path.join(config.save_dir, "knowledge_graph", self.kgdb_name)
os.makedirs(self.work_dir, exist_ok=True) os.makedirs(self.work_dir, exist_ok=True)
@ -45,6 +48,7 @@ class GraphDatabase:
def close(self): def close(self):
"""关闭数据库连接""" """关闭数据库连接"""
assert self.driver is not None, "Database is not connected"
self.driver.close() self.driver.close()
def is_running(self): def is_running(self):
@ -53,6 +57,7 @@ class GraphDatabase:
def get_sample_nodes(self, kgdb_name='neo4j', num=50): def get_sample_nodes(self, kgdb_name='neo4j', num=50):
"""获取指定数据库的 num 个节点信息""" """获取指定数据库的 num 个节点信息"""
assert self.driver is not None, "Database is not connected"
self.use_database(kgdb_name) self.use_database(kgdb_name)
def query(tx, num): def query(tx, num):
result = tx.run("MATCH (n)-[r]->(m) RETURN n, r, m LIMIT $num", num=int(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): def create_graph_database(self, kgdb_name):
"""创建新的数据库,如果已存在则返回已有数据库的名称""" """创建新的数据库,如果已存在则返回已有数据库的名称"""
assert self.driver is not None, "Database is not connected"
with self.driver.session() as session: with self.driver.session() as session:
existing_databases = session.run("SHOW DATABASES") existing_databases = session.run("SHOW DATABASES")
existing_db_names = [db['name'] for db in existing_databases] existing_db_names = [db['name'] for db in existing_databases]
@ -71,7 +77,7 @@ class GraphDatabase:
print(f"已存在数据库: {existing_db_names[0]}") print(f"已存在数据库: {existing_db_names[0]}")
return 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}' 创建成功.") print(f"数据库 '{kgdb_name}' 创建成功.")
return kgdb_name # 返回创建的数据库名称 return kgdb_name # 返回创建的数据库名称
@ -83,6 +89,7 @@ class GraphDatabase:
def txt_add_entity(self, triples, kgdb_name='neo4j'): def txt_add_entity(self, triples, kgdb_name='neo4j'):
"""添加实体三元组""" """添加实体三元组"""
assert self.driver is not None, "Database is not connected"
self.use_database(kgdb_name) self.use_database(kgdb_name)
def create(tx, triples): def create(tx, triples):
for triple in triples: for triple in triples:
@ -101,6 +108,7 @@ class GraphDatabase:
async def txt_add_vector_entity(self, triples, kgdb_name='neo4j'): 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) self.use_database(kgdb_name)
def _index_exists(tx, index_name): def _index_exists(tx, index_name):
"""检查索引是否存在""" """检查索引是否存在"""
@ -211,6 +219,7 @@ class GraphDatabase:
self.save_graph_info() self.save_graph_info()
async def jsonl_file_add_entity(self, file_path, kgdb_name='neo4j'): 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" self.status = "processing"
kgdb_name = kgdb_name or 'neo4j' kgdb_name = kgdb_name or 'neo4j'
self.use_database(kgdb_name) # 切换到指定数据库 self.use_database(kgdb_name) # 切换到指定数据库
@ -233,6 +242,7 @@ class GraphDatabase:
def delete_entity(self, entity_name=None, kgdb_name="neo4j"): def delete_entity(self, entity_name=None, kgdb_name="neo4j"):
"""删除数据库中的指定实体三元组, 参数entity_name为空则删除全部实体""" """删除数据库中的指定实体三元组, 参数entity_name为空则删除全部实体"""
assert self.driver is not None, "Database is not connected"
self.use_database(kgdb_name) self.use_database(kgdb_name)
with self.driver.session() as session: with self.driver.session() as session:
if entity_name: 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): 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 停止检索 # TODO 添加判断节点数量为 0 停止检索
# 判断是否启动 # 判断是否启动
if not self.is_running(): 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): 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: if not entity_name:
logger.warning("实体名称为空") logger.warning("实体名称为空")
return [] return []
@ -343,6 +355,7 @@ class GraphDatabase:
def query_all_nodes_and_relationships(self, kgdb_name='neo4j', hops = 2): def query_all_nodes_and_relationships(self, kgdb_name='neo4j', hops = 2):
"""查询图数据库中所有三元组信息 NEVER USE""" """查询图数据库中所有三元组信息 NEVER USE"""
assert self.driver is not None, "Database is not connected"
self.use_database(kgdb_name) self.use_database(kgdb_name)
def query(tx, hops): def query(tx, hops):
result = tx.run(f""" result = tx.run(f"""
@ -358,6 +371,7 @@ class GraphDatabase:
def query_by_relationship_type(self, relationship_type, kgdb_name='neo4j', hops = 2): def query_by_relationship_type(self, relationship_type, kgdb_name='neo4j', hops = 2):
"""查询指定关系三元组信息 NEVER USE""" """查询指定关系三元组信息 NEVER USE"""
assert self.driver is not None, "Database is not connected"
self.use_database(kgdb_name) self.use_database(kgdb_name)
def query(tx, relationship_type, hops): def query(tx, relationship_type, hops):
result = tx.run(f""" result = tx.run(f"""
@ -373,6 +387,7 @@ class GraphDatabase:
def query_entity_like(self, keyword, kgdb_name='neo4j', hops = 2): def query_entity_like(self, keyword, kgdb_name='neo4j', hops = 2):
"""模糊查询 NEVER USE""" """模糊查询 NEVER USE"""
assert self.driver is not None, "Database is not connected"
self.use_database(kgdb_name) self.use_database(kgdb_name)
def query(tx, keyword, hops): def query(tx, keyword, hops):
result = tx.run(f""" result = tx.run(f"""
@ -390,6 +405,7 @@ class GraphDatabase:
def query_node_info(self, node_name, kgdb_name='neo4j', hops = 2): def query_node_info(self, node_name, kgdb_name='neo4j', hops = 2):
"""查询指定节点的详细信息返回信息 NEVER USE""" """查询指定节点的详细信息返回信息 NEVER USE"""
assert self.driver is not None, "Database is not connected"
self.use_database(kgdb_name) # 切换到指定数据库 self.use_database(kgdb_name) # 切换到指定数据库
def query(tx, node_name, hops): def query(tx, node_name, hops):
result = tx.run(f""" result = tx.run(f"""
@ -405,23 +421,19 @@ class GraphDatabase:
return session.execute_read(query, node_name, hops) return session.execute_read(query, node_name, hops)
async def aget_embedding(self, text): async def aget_embedding(self, text):
from src import knowledge_base
if isinstance(text, list): 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 return outputs
else: else:
outputs = await knowledge_base.embed_model.aencode(text) outputs = await self.embed_model.aencode(text)
return outputs return outputs
def get_embedding(self, text): def get_embedding(self, text):
from src import knowledge_base
if isinstance(text, list): 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 return outputs
else: else:
outputs = knowledge_base.embed_model.encode([text])[0] outputs = self.embed_model.encode([text])[0]
return outputs return outputs
def set_embedding(self, tx, entity_name, embedding): def set_embedding(self, tx, entity_name, embedding):
@ -431,6 +443,7 @@ class GraphDatabase:
""", name=entity_name, embedding=embedding) """, name=entity_name, embedding=embedding)
def get_graph_info(self, graph_name="neo4j"): def get_graph_info(self, graph_name="neo4j"):
assert self.driver is not None, "Database is not connected"
self.use_database(graph_name) self.use_database(graph_name)
def query(tx): def query(tx):
entity_count = tx.run("MATCH (n) RETURN count(n) AS count").single()["count"] entity_count = tx.run("MATCH (n) RETURN count(n) AS count").single()["count"]
@ -493,6 +506,7 @@ class GraphDatabase:
Returns: Returns:
list: 没有嵌入向量的节点列表 list: 没有嵌入向量的节点列表
""" """
assert self.driver is not None, "Database is not connected"
self.use_database(kgdb_name) self.use_database(kgdb_name)
def query(tx): def query(tx):
@ -543,6 +557,7 @@ class GraphDatabase:
Returns: Returns:
int: 成功添加嵌入向量的节点数量 int: 成功添加嵌入向量的节点数量
""" """
assert self.driver is not None, "Database is not connected"
self.use_database(kgdb_name) self.use_database(kgdb_name)
# 如果node_names为None则获取所有没有嵌入向量的节点 # 如果node_names为None则获取所有没有嵌入向量的节点
@ -575,6 +590,7 @@ class GraphDatabase:
source_id = source.element_id source_id = source.element_id
target_id = target.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 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 target_name = node_dict[target_id]["name"] if target_name is None else target_name

View File

@ -43,7 +43,6 @@
<a-select style="width: 300px" <a-select style="width: 300px"
:value="configStore.config?.reranker" :value="configStore.config?.reranker"
@change="handleChange('reranker', $event)" @change="handleChange('reranker', $event)"
:disabled="!configStore.config.enable_reranker"
> >
<a-select-option <a-select-option
v-for="(name, idx) in items?.reranker.choices" :key="idx" v-for="(name, idx) in items?.reranker.choices" :key="idx"
@ -51,25 +50,6 @@
</a-select-option> </a-select-option>
</a-select> </a-select>
</div> </div>
<div class="card">
<span class="label">{{ items?.enable_reranker.des }}</span>
<a-switch
:checked="configStore.config.enable_reranker"
@change="handleChange('enable_reranker', !configStore.config.enable_reranker)"
/>
</div>
<div class="card card-select">
<span class="label">{{ items?.use_rewrite_query.des }}</span>
<a-select style="width: 200px"
:value="configStore.config?.use_rewrite_query"
@change="handleChange('use_rewrite_query', $event)"
>
<a-select-option
v-for="(name, idx) in items?.use_rewrite_query.choices" :key="idx"
:value="name">{{ name }}
</a-select-option>
</a-select>
</div>
</div> </div>
</div> </div>
<div class="setting" v-if="(state.windowWidth <= 520 || state.section === 'model') && userStore.isSuperAdmin"> <div class="setting" v-if="(state.windowWidth <= 520 || state.section === 'model') && userStore.isSuperAdmin">