refactor(graph): 重构知识图谱查询和格式化逻辑
- 修改查询方法以支持不同返回格式(graph/triples) - 移除不再使用的格式化方法和废弃查询方法 - 优化节点和关系的查询性能及结果处理 - 更新相关API端点以直接返回原始查询结果
This commit is contained in:
parent
007e6aeb1f
commit
612bd9f7b0
@ -189,16 +189,15 @@ async def get_neo4j_nodes(
|
||||
raise HTTPException(status_code=400, detail="图数据库未启动")
|
||||
|
||||
result = graph_base.get_sample_nodes(kgdb_name, num)
|
||||
formatted_result = graph_base.format_general_results(result)
|
||||
|
||||
return {
|
||||
"success": True,
|
||||
"result": formatted_result,
|
||||
"result": result,
|
||||
"message": "success"
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"获取图节点数据失败: {e}")
|
||||
logger.error(f"获取图节点数据失败: {e}\n{traceback.format_exc()}")
|
||||
raise HTTPException(status_code=500, detail=f"获取图节点数据失败: {str(e)}")
|
||||
|
||||
@graph.get("/neo4j/node")
|
||||
@ -214,11 +213,10 @@ async def get_neo4j_node(
|
||||
raise HTTPException(status_code=400, detail="图数据库未启动")
|
||||
|
||||
result = graph_base.query_node(entity_name=entity_name)
|
||||
formatted_result = graph_base.format_query_result_to_graph(result)
|
||||
|
||||
return {
|
||||
"success": True,
|
||||
"result": formatted_result,
|
||||
"result": result,
|
||||
"message": "success"
|
||||
}
|
||||
|
||||
|
||||
@ -18,7 +18,7 @@ from src.agents.chatbot.configuration import ChatbotConfiguration
|
||||
from src.agents.tools_factory import get_runnable_tools
|
||||
|
||||
class ChatbotAgent(BaseAgent):
|
||||
name = "对话机器人(Chatbot)"
|
||||
name = "问答助手"
|
||||
description = "基础的对话机器人,可以回答问题,默认不使用任何工具,可在配置中启用需要的工具。"
|
||||
config_schema = ChatbotConfiguration
|
||||
|
||||
|
||||
@ -139,8 +139,8 @@ def calculator(a: float, b: float, operation: str) -> float:
|
||||
|
||||
@tool
|
||||
def query_knowledge_graph(query: Annotated[str, "The keyword to query knowledge graph."]):
|
||||
"""Use this to query knowledge graph."""
|
||||
return graph_base.query_node(query, hops=2)
|
||||
"""Use this to query knowledge graph, which include some food domain knowledge."""
|
||||
return graph_base.query_node(query, hops=2, return_format='triples')
|
||||
|
||||
# 更新工具注册表
|
||||
_TOOLS_REGISTRY.update({
|
||||
|
||||
@ -61,11 +61,26 @@ class GraphDatabase:
|
||||
self.use_database(kgdb_name)
|
||||
def query(tx, num):
|
||||
"""Note: 这里只查询带有 Entity 标签的节点"""
|
||||
result = tx.run("MATCH (n:Entity)-[r]->(m:Entity) RETURN n, r, m LIMIT $num", num=int(num))
|
||||
return result.values()
|
||||
query_str = """
|
||||
MATCH (n:Entity)-[r]->(m:Entity)
|
||||
RETURN
|
||||
{id: elementId(n), name: n.name} AS h,
|
||||
{type: r.type, source_id: elementId(n), target_id: elementId(m)} AS r,
|
||||
{id: elementId(m), name: m.name} AS t
|
||||
LIMIT $num
|
||||
"""
|
||||
results = tx.run(query_str, num=int(num))
|
||||
formatted_results = {'nodes': [], 'edges': []}
|
||||
|
||||
for item in results:
|
||||
formatted_results['nodes'].extend([item['h'], item['t']])
|
||||
formatted_results['edges'].append(item['r'])
|
||||
|
||||
return formatted_results
|
||||
|
||||
with self.driver.session() as session:
|
||||
return session.execute_read(query, num)
|
||||
results = session.execute_read(query, num)
|
||||
return results
|
||||
|
||||
def create_graph_database(self, kgdb_name):
|
||||
"""创建新的数据库,如果已存在则返回已有数据库的名称"""
|
||||
@ -266,7 +281,7 @@ class GraphDatabase:
|
||||
"""
|
||||
tx.run(query)
|
||||
|
||||
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, return_format='graph', **kwargs):
|
||||
"""知识图谱查询节点的入口:"""
|
||||
assert self.driver is not None, "Database is not connected"
|
||||
assert self.is_running(), "图数据库未启动"
|
||||
@ -274,26 +289,32 @@ class GraphDatabase:
|
||||
self.use_database(kgdb_name)
|
||||
|
||||
# 使用向量索引进行查询
|
||||
results_sim = self._query_with_vector_sim(entity_name, kgdb_name, hops, threshold)
|
||||
results_fuzzy = self._query_with_fuzzy_match(entity_name, kgdb_name, hops)
|
||||
results_sim = self._query_with_vector_sim(entity_name, kgdb_name, threshold)
|
||||
results_fuzzy = self._query_with_fuzzy_match(entity_name, kgdb_name)
|
||||
results = results_sim + results_fuzzy
|
||||
qualified_entities = [result[0] for result in results][:max_entities]
|
||||
|
||||
logger.debug(f"Graph Query Entities: {entity_name}, {qualified_entities=}")
|
||||
|
||||
# 对每个合格的实体进行查询
|
||||
all_query_results = []
|
||||
all_query_results = {'nodes': [], 'edges': [], 'triples': []}
|
||||
for entity in qualified_entities:
|
||||
query_result = self._query_specific_entity(entity_name=entity, hops=hops, kgdb_name=kgdb_name)
|
||||
all_query_results.extend(query_result)
|
||||
query_result = self._query_specific_entity(entity_name=entity, kgdb_name=kgdb_name, hops=hops)
|
||||
if return_format == 'graph':
|
||||
all_query_results['nodes'].extend(query_result['nodes'])
|
||||
all_query_results['edges'].extend(query_result['edges'])
|
||||
elif return_format == 'triples':
|
||||
all_query_results['triples'].extend(query_result['triples'])
|
||||
else:
|
||||
raise ValueError(f"Invalid return_format: {return_format}")
|
||||
|
||||
return all_query_results
|
||||
|
||||
def _query_with_fuzzy_match(self, keyword, kgdb_name='neo4j', hops = 2):
|
||||
def _query_with_fuzzy_match(self, keyword, kgdb_name='neo4j'):
|
||||
"""模糊查询"""
|
||||
assert self.driver is not None, "Database is not connected"
|
||||
self.use_database(kgdb_name)
|
||||
def query_fuzzy_match(tx, keyword, hops):
|
||||
def query_fuzzy_match(tx, keyword):
|
||||
result = tx.run("""
|
||||
MATCH (n:Entity)
|
||||
WHERE n.name CONTAINS $keyword
|
||||
@ -304,9 +325,9 @@ class GraphDatabase:
|
||||
return values
|
||||
|
||||
with self.driver.session() as session:
|
||||
return session.execute_read(query_fuzzy_match, keyword, hops)
|
||||
return session.execute_read(query_fuzzy_match, keyword)
|
||||
|
||||
def _query_with_vector_sim(self, keyword, kgdb_name='neo4j', hops = 2, threshold=0.9):
|
||||
def _query_with_vector_sim(self, keyword, kgdb_name='neo4j', threshold=0.9):
|
||||
"""向量查询"""
|
||||
assert self.driver is not None, "Database is not connected"
|
||||
self.use_database(kgdb_name)
|
||||
@ -334,7 +355,6 @@ class GraphDatabase:
|
||||
|
||||
with self.driver.session() as session:
|
||||
results = session.execute_read(query_by_vector, keyword, threshold=threshold)
|
||||
results = clean_triples_embedding(results)
|
||||
return results
|
||||
|
||||
|
||||
@ -349,21 +369,35 @@ class GraphDatabase:
|
||||
|
||||
def query(tx, entity_name, hops, limit):
|
||||
try:
|
||||
query_str = f"""
|
||||
MATCH (n {{name: $entity_name}})-[r*1..{hops}]-(m)
|
||||
RETURN n AS n, r, m AS m
|
||||
query_str = """
|
||||
MATCH (n {name: $entity_name})-[r1]-(m1)
|
||||
RETURN
|
||||
{id: elementId(n), name: n.name} AS h,
|
||||
{type: r1.type, source_id: elementId(n), target_id: elementId(m1)} AS r,
|
||||
{id: elementId(m1), name: m1.name} AS t
|
||||
UNION
|
||||
MATCH (n {name: $entity_name})-[r1]-(m1)-[r2]-(m2)
|
||||
RETURN
|
||||
{id: elementId(m1), name: m1.name} AS h,
|
||||
{type: r2.type, source_id: elementId(m1), target_id: elementId(m2)} AS r,
|
||||
{id: elementId(m2), name: m2.name} AS t
|
||||
LIMIT $limit
|
||||
"""
|
||||
result = tx.run(query_str, entity_name=entity_name, limit=limit)
|
||||
results = tx.run(query_str, entity_name=entity_name, limit=limit)
|
||||
|
||||
if not result:
|
||||
if not results:
|
||||
logger.info(f"未找到实体 {entity_name} 的相关信息")
|
||||
return []
|
||||
return {}
|
||||
|
||||
values = result.values()
|
||||
# 安全地处理embedding属性
|
||||
values = clean_triples_embedding(values)
|
||||
return values
|
||||
formatted_results = {'nodes': [], 'edges': [], 'triples': []}
|
||||
|
||||
for item in results:
|
||||
formatted_results['nodes'].extend([item['h'], item['t']])
|
||||
formatted_results['edges'].append(item['r'])
|
||||
formatted_results['triples'].append((item['h']['name'], item['r']['type'], item['t']['name']))
|
||||
|
||||
logger.debug(f"Query Results: {results}")
|
||||
return formatted_results
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"查询实体 {entity_name} 失败: {str(e)}")
|
||||
@ -372,62 +406,11 @@ class GraphDatabase:
|
||||
try:
|
||||
with self.driver.session() as session:
|
||||
return session.execute_read(query, entity_name, hops, limit)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"数据库会话异常: {str(e)}")
|
||||
return []
|
||||
|
||||
def query_all_nodes_and_relationships(self, kgdb_name='neo4j', hops = 2):
|
||||
"""查询图数据库中所有三元组信息 NEVER USE"""
|
||||
raise Exception("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"""
|
||||
MATCH (n)-[r*1..{hops}]->(m)
|
||||
RETURN n AS n, r, m AS m
|
||||
""")
|
||||
values = result.values()
|
||||
values = clean_triples_embedding(values)
|
||||
return values
|
||||
|
||||
with self.driver.session() as session:
|
||||
return session.execute_read(query, hops)
|
||||
|
||||
def query_by_relationship_type(self, relationship_type, kgdb_name='neo4j', hops = 2):
|
||||
"""查询指定关系三元组信息 NEVER USE"""
|
||||
raise Exception("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"""
|
||||
MATCH (n)-[r:`{relationship_type}`*1..{hops}]->(m)
|
||||
RETURN n AS n, r, m AS m
|
||||
""")
|
||||
values = result.values()
|
||||
values = clean_triples_embedding(values)
|
||||
return values
|
||||
|
||||
with self.driver.session() as session:
|
||||
return session.execute_read(query, relationship_type, hops)
|
||||
|
||||
def query_node_info(self, node_name, kgdb_name='neo4j', hops = 2):
|
||||
"""查询指定节点的详细信息返回信息 NEVER USE"""
|
||||
raise Exception("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"""
|
||||
MATCH (n {{name: $node_name}})
|
||||
OPTIONAL MATCH (n)-[r*1..{hops}]->(m)
|
||||
RETURN n AS n, r, m AS m
|
||||
""", node_name=node_name)
|
||||
values = result.values()
|
||||
values = clean_triples_embedding(values)
|
||||
return values
|
||||
|
||||
with self.driver.session() as session:
|
||||
return session.execute_read(query, node_name, hops)
|
||||
|
||||
async def aget_embedding(self, text):
|
||||
if isinstance(text, list):
|
||||
outputs = await self.embed_model.abatch_encode(text, batch_size=40)
|
||||
@ -587,64 +570,15 @@ class GraphDatabase:
|
||||
|
||||
return count
|
||||
|
||||
|
||||
def _extract_relationship_info(self, relationship, source_name=None, target_name=None, node_dict=None):
|
||||
"""
|
||||
提取关系信息并返回格式化的节点和边信息
|
||||
"""
|
||||
rel_id = relationship.element_id
|
||||
nodes = relationship.nodes
|
||||
if len(nodes) != 2:
|
||||
return None, None
|
||||
|
||||
source, target = nodes
|
||||
source_id = source.element_id
|
||||
target_id = target.element_id
|
||||
|
||||
# 如果没有提供 source_name 或 target_name,则需要 node_dict
|
||||
if source_name is None or target_name is None:
|
||||
assert node_dict is not None, "node_dict is required when source_name or target_name is None"
|
||||
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
|
||||
|
||||
relationship_type = relationship._properties.get("type", "unknown")
|
||||
if relationship_type == "unknown":
|
||||
relationship_type = relationship.type
|
||||
|
||||
edge_info = {
|
||||
"id": rel_id,
|
||||
"type": relationship_type,
|
||||
"source_id": source_id,
|
||||
"target_id": target_id,
|
||||
"source_name": source_name,
|
||||
"target_name": target_name,
|
||||
}
|
||||
|
||||
node_info = [
|
||||
{"id": source_id, "name": source_name},
|
||||
{"id": target_id, "name": target_name},
|
||||
]
|
||||
|
||||
return node_info, edge_info
|
||||
|
||||
def format_general_results(self, results):
|
||||
formatted_results = {"nodes": [], "edges": []}
|
||||
nodes = []
|
||||
edges = []
|
||||
|
||||
for item in results:
|
||||
relationship = item[1]
|
||||
source_name = item[0]._properties.get("name", "unknown")
|
||||
target_name = item[2]._properties.get("name", "unknown") if len(item) > 2 else "unknown"
|
||||
|
||||
node_info, edge_info = self._extract_relationship_info(relationship, source_name, target_name)
|
||||
if node_info is None or edge_info is None:
|
||||
continue
|
||||
|
||||
for node in node_info:
|
||||
if node["id"] not in [n["id"] for n in formatted_results["nodes"]]:
|
||||
formatted_results["nodes"].append(node)
|
||||
|
||||
formatted_results["edges"].append(edge_info)
|
||||
nodes.extend([item['h'], item['t']])
|
||||
edges.append(item['r'])
|
||||
|
||||
formatted_results = {"nodes": nodes, "edges": edges}
|
||||
return formatted_results
|
||||
|
||||
def format_query_result_to_graph(self, query_results):
|
||||
@ -709,6 +643,45 @@ class GraphDatabase:
|
||||
|
||||
return formatted_results
|
||||
|
||||
def _extract_relationship_info(self, relationship, source_name=None, target_name=None, node_dict=None):
|
||||
"""
|
||||
提取关系信息并返回格式化的节点和边信息
|
||||
"""
|
||||
rel_id = relationship.element_id
|
||||
nodes = relationship.nodes
|
||||
if len(nodes) != 2:
|
||||
return None, None
|
||||
|
||||
source, target = nodes
|
||||
source_id = source.element_id
|
||||
target_id = target.element_id
|
||||
|
||||
# 如果没有提供 source_name 或 target_name,则需要 node_dict
|
||||
if source_name is None or target_name is None:
|
||||
assert node_dict is not None, "node_dict is required when source_name or target_name is None"
|
||||
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
|
||||
|
||||
relationship_type = relationship._properties.get("type", "unknown")
|
||||
if relationship_type == "unknown":
|
||||
relationship_type = relationship.type
|
||||
|
||||
edge_info = {
|
||||
"id": rel_id,
|
||||
"type": relationship_type,
|
||||
"source_id": source_id,
|
||||
"target_id": target_id,
|
||||
"source_name": source_name,
|
||||
"target_name": target_name,
|
||||
}
|
||||
|
||||
node_info = [
|
||||
{"id": source_id, "name": source_name},
|
||||
{"id": target_id, "name": target_name},
|
||||
]
|
||||
|
||||
return node_info, edge_info
|
||||
|
||||
def clean_triples_embedding(triples):
|
||||
for item in triples:
|
||||
if hasattr(item[0], '_properties'):
|
||||
|
||||
Loading…
Reference in New Issue
Block a user