fix(graphbase): 修复检索过程中的 limit 不生效的问题 #260

This commit is contained in:
Wenjie Zhang 2025-09-11 03:19:56 +08:00
parent d6dfa2b5e3
commit 45fb5c1b0d
2 changed files with 30 additions and 28 deletions

View File

@ -201,7 +201,7 @@ async def get_neo4j_node(
if not graph_base.is_running(): if not graph_base.is_running():
raise HTTPException(status_code=400, detail="图数据库未启动") raise HTTPException(status_code=400, detail="图数据库未启动")
result = graph_base.query_node(entity_name=entity_name) result = graph_base.query_node(keyword=entity_name)
return {"success": True, "result": result, "message": "success"} return {"success": True, "result": result, "message": "success"}

View File

@ -386,7 +386,7 @@ class GraphDatabase:
tx.run(query) tx.run(query)
def query_node( def query_node(
self, entity_name, threshold=0.9, kgdb_name="neo4j", hops=2, max_entities=5, return_format="graph", **kwargs self, keyword, threshold=0.9, kgdb_name="neo4j", hops=2, max_entities=8, return_format="graph", **kwargs
): ):
"""知识图谱查询节点的入口:""" """知识图谱查询节点的入口:"""
assert self.driver is not None, "Database is not connected" assert self.driver is not None, "Database is not connected"
@ -395,12 +395,12 @@ class GraphDatabase:
self.use_database(kgdb_name) self.use_database(kgdb_name)
# 使用向量索引进行查询 # 使用向量索引进行查询
results_sim = self._query_with_vector_sim(entity_name, kgdb_name, threshold) results_sim = self._query_with_vector_sim(keyword, kgdb_name, threshold)
results_fuzzy = self._query_with_fuzzy_match(entity_name, kgdb_name) results_fuzzy = self._query_with_fuzzy_match(keyword, kgdb_name)
results = results_sim + results_fuzzy results = results_sim + results_fuzzy
qualified_entities = [result[0] for result in results][:max_entities] qualified_entities = [result[0] for result in results][:max_entities]
logger.debug(f"Graph Query Entities: {entity_name}, {qualified_entities=}") logger.debug(f"Graph Query Entities: {keyword}, {qualified_entities=}")
# 对每个合格的实体进行查询 # 对每个合格的实体进行查询
all_query_results = {"nodes": [], "edges": [], "triples": []} all_query_results = {"nodes": [], "edges": [], "triples": []}
@ -484,29 +484,31 @@ class GraphDatabase:
def query(tx, entity_name, hops, limit): def query(tx, entity_name, hops, limit):
try: try:
query_str = """ query_str = """
MATCH (n {name: $entity_name})-[r1]->(m1) WITH [
RETURN // 1跳出边
{id: elementId(n), name: n.name} AS h, [(n {name: $entity_name})-[r1]->(m1) |
{type: r1.type, source_id: elementId(n), target_id: elementId(m1)} AS r, {h: {id: elementId(n), name: n.name},
{id: elementId(m1), name: m1.name} AS t r: {type: r1.type, source_id: elementId(n), target_id: elementId(m1)},
UNION t: {id: elementId(m1), name: m1.name}}],
MATCH (n {name: $entity_name})-[r1]->(m1)-[r2]->(m2) // 2跳出边
RETURN [(n {name: $entity_name})-[r1]->(m1)-[r2]->(m2) |
{id: elementId(m1), name: m1.name} AS h, {h: {id: elementId(m1), name: m1.name},
{type: r2.type, source_id: elementId(m1), target_id: elementId(m2)} AS r, r: {type: r2.type, source_id: elementId(m1), target_id: elementId(m2)},
{id: elementId(m2), name: m2.name} AS t t: {id: elementId(m2), name: m2.name}}],
UNION // 1跳入边
MATCH (m1)-[r1]->(n {name: $entity_name}) [(m1)-[r1]->(n {name: $entity_name}) |
RETURN {h: {id: elementId(m1), name: m1.name},
{id: elementId(m1), name: m1.name} AS h, r: {type: r1.type, source_id: elementId(m1), target_id: elementId(n)},
{type: r1.type, source_id: elementId(m1), target_id: elementId(n)} AS r, t: {id: elementId(n), name: n.name}}],
{id: elementId(n), name: n.name} AS t // 2跳入边
UNION [(m2)-[r2]->(m1)-[r1]->(n {name: $entity_name}) |
MATCH (m2)-[r2]->(m1)-[r1]->(n {name: $entity_name}) {h: {id: elementId(m2), name: m2.name},
RETURN r: {type: r2.type, source_id: elementId(m2), target_id: elementId(m1)},
{id: elementId(m2), name: m2.name} AS h, t: {id: elementId(m1), name: m1.name}}]
{type: r2.type, source_id: elementId(m2), target_id: elementId(m1)} AS r, ] AS all_results
{id: elementId(m1), name: m1.name} AS t UNWIND all_results AS result_list
UNWIND result_list AS item
RETURN item.h AS h, item.r AS r, item.t AS t
LIMIT $limit LIMIT $limit
""" """
results = tx.run(query_str, entity_name=entity_name, limit=limit) results = tx.run(query_str, entity_name=entity_name, limit=limit)