class GraphQueryRequest(BaseModel): queryview: Optional[str] = None querylink: Optional[str] = None head: Optional[str] = None tail: Optional[str] = None node: Optional[str] = None # 在这里设置默认值为 3 depth: Optional[int] = 3 @app.post("/api/v2/neo4j/query") async def graph_query(request: GraphQueryRequest): """ 返回结构: { "message": str, "simple_graphData": {...}, # 200节点采样(带 level 属性) "graphData": {...}, # 查询结果图(头尾路径 / 节点展开等) "short_routes": {...} # 最短路径(仅头尾模式有效) } """ queryview = request.queryview querylink= request.querylink head = request.head tail = request.tail node = request.node depth = request.depth # 1. 始终获取采样数据 logger.info("Fetching graph sample | Limit: 200 nodes") try: simple_graphData = await run_in_threadpool(fetch_graph_sample, driver, 200) except Exception as e: logger.exception("Neo4j sample query error") raise HTTPException(status_code=500, detail=f"Sample query failed: {str(e)}") def db_task(): try: with driver.session() as session: # 默认返回值 res = { "graphData": {"nodes": [], "links": []}, "short_routes": {"nodes": [], "links": []}, "message": "" } # =============================== # 情况 1: 头尾节点都有 # =============================== if head and tail: logger.info(f"Mode: HEAD + TAIL | {head} -> {tail}") h_exists = session.run("MATCH (h {名称: $head}) RETURN h", head=head).single() t_exists = session.run("MATCH (t {名称: $tail}) RETURN t", tail=tail).single() if h_exists and t_exists: res["message"] = "头节点和尾节点展开成功" if not h_exists and not t_exists: res["message"] = "头节点和尾节点均不存在" res["graphData"] = {"nodes": [], "links": []} return res elif not h_exists: res["message"] = "头节点不存在" # 返回存在的尾节点(如果存在) nodes_list = [] if t_exists: nodes_list.append(node_to_json(t_exists["t"])) res["graphData"] = {"nodes": nodes_list, "links": []} return res elif not t_exists: res["message"] = "尾节点不存在" nodes_list = [node_to_json(h_exists["h"])] res["graphData"] = {"nodes": nodes_list, "links": []} return res # 两者都存在,查路径 sp_paths = list(session.run(""" MATCH path = allShortestPaths((h {名称: $head})-[*..10]->(t {名称: $tail})) RETURN path AS p """, head=head, tail=tail)) all_paths = list(session.run(""" MATCH path = (h {名称: $head})-[*1..10]->(t {名称: $tail}) UNWIND nodes(path) AS x WITH path, COUNT(DISTINCT x) AS nodeCount WHERE nodeCount = LENGTH(path) + 1 RETURN path AS p """, head=head, tail=tail)) if not sp_paths and not all_paths: # 无连接,返回两个孤立节点 nodes_list = [ node_to_json(h_exists["h"]), node_to_json(t_exists["t"]) ] res["graphData"] = {"nodes": nodes_list, "links": []} res["short_routes"] = {"nodes": [], "links": []} else: res["graphData"] = build_graph_from_paths(all_paths) res["short_routes"] = build_graph_from_paths(sp_paths) # === 新增:为 graphData 中每个节点计算 level(最大无向跳数,≤3)=== nodes_list = res["graphData"]["nodes"] if not nodes_list: return res # 确保每个节点都有 "名称" valid_nodes = [n for n in nodes_list if "name" in n and n["name"]] node_names = [n["name"] for n in valid_nodes] name_to_node = {n["name"]: n for n in valid_nodes} for name in node_names: try: result = session.run(""" MATCH (start {名称: $name}) OPTIONAL MATCH path = (start)-[*1..3]-(end) WHERE start <> end RETURN coalesce(max(length(path)), 0) AS max_level """, name=name).single() if result: level = min(result["max_level"], 3) name_to_node[name]["level"] = level else: name_to_node[name]["level"] = 0 except Exception as e: logger.warning(f"Failed to compute level for {name}: {e}") name_to_node[name]["level"] = 0 return res # =============================== # 情况 2: 指定节点 + 深度 # =============================== elif node and depth is not None: logger.info(f"Mode: NODE + DEPTH | {node} depth={depth}") # # 校验节点是否存在 # n_exists = session.run("MATCH (n {名称: $node}) RETURN n", node=node).single() # if not n_exists: # res["message"] = "指定节点不存在" # return res # 有向:从 node 向外走 depth 跳 paths = session.run(f""" MATCH path = (n {{名称: $node}})-[*1..{depth}]-(m) UNWIND nodes(path) AS x WITH path, COUNT(DISTINCT x) AS nodeCount WHERE nodeCount = LENGTH(path) + 1 RETURN path AS p """, node=node) res["graphData"] = build_graph_from_paths(paths) # === 新增:为 graphData 中每个节点计算 level(最大无向跳数,≤3)=== nodes_list = res["graphData"]["nodes"] if not nodes_list: return res # 确保每个节点都有 "名称" valid_nodes = [n for n in nodes_list if "name" in n and n["name"]] node_names = [n["name"] for n in valid_nodes] name_to_node = {n["name"]: n for n in valid_nodes} for name in node_names: try: result = session.run(""" MATCH (start {名称: $name}) OPTIONAL MATCH path = (start)-[*1..3]-(end) WHERE start <> end RETURN coalesce(max(length(path)), 0) AS max_level """, name=name).single() if result: level = min(result["max_level"], 3) name_to_node[name]["level"] = level else: name_to_node[name]["level"] = 0 except Exception as e: logger.warning(f"Failed to compute level for {name}: {e}") name_to_node[name]["level"] = 0 res["message"] = f"{depth}节点展开成功" return res # =============================== # 情况 3: 只有头节点(有向,最多3跳) # =============================== elif head: logger.info(f"Mode: HEAD ONLY | {head}") h_exists = session.run("MATCH (h {名称: $head}) RETURN h", head=head).single() if not h_exists: res["message"] = "头节点不存在" return res paths = session.run(""" MATCH path = (h {名称: $head})-[*1..3]->(n) UNWIND nodes(path) AS x WITH path, COUNT(DISTINCT x) AS nodeCount WHERE nodeCount = LENGTH(path) + 1 RETURN path AS p """, head=head) paths = list(paths) if not paths: # 3. 如果没有有向路径,返回单个节点信息 res["graphData"] = { "nodes": [node_to_json(h_exists["h"])], "links": [] } else: res["graphData"] = build_graph_from_paths(paths) # === 新增:为 graphData 中每个节点计算 level(最大无向跳数,≤3)=== nodes_list = res["graphData"]["nodes"] if not nodes_list: return res # 确保每个节点都有 "名称" valid_nodes = [n for n in nodes_list if "name" in n and n["name"]] node_names = [n["name"] for n in valid_nodes] name_to_node = {n["name"]: n for n in valid_nodes} for name in node_names: try: result = session.run(""" MATCH (start {名称: $name}) OPTIONAL MATCH path = (start)-[*1..3]-(end) WHERE start <> end RETURN coalesce(max(length(path)), 0) AS max_level """, name=name).single() if result: level = min(result["max_level"], 3) name_to_node[name]["level"] = level else: name_to_node[name]["level"] = 0 except Exception as e: logger.warning(f"Failed to compute level for {name}: {e}") name_to_node[name]["level"] = 0 res["message"] = "头节点展开成功" return res # =============================== # 情况 4: 只有尾节点(应查入边!即 ←) # =============================== elif tail: logger.info(f"Mode: TAIL ONLY | {tail}") t_exists = session.run("MATCH (t {名称: $tail}) RETURN t", tail=tail).single() if not t_exists: res["message"] = "尾节点不存在" return res # 注意:尾节点作为终点,应查找指向它的路径(←) paths = session.run(""" MATCH path = (n)-[*1..3]->(t {名称: $tail}) UNWIND nodes(path) AS x WITH path, COUNT(DISTINCT x) AS nodeCount WHERE nodeCount = LENGTH(path) + 1 RETURN path AS p """, tail=tail) paths = list(paths) if not paths: # 3. 如果没有有向路径,返回单个节点信息 res["graphData"] = { "nodes": [node_to_json(t_exists["t"])], "links": [] } else: res["graphData"] = build_graph_from_paths(paths) # === 新增:为 graphData 中每个节点计算 level(最大无向跳数,≤3)=== nodes_list = res["graphData"]["nodes"] if not nodes_list: return res # 确保每个节点都有 "名称" valid_nodes = [n for n in nodes_list if "name" in n and n["name"]] node_names = [n["name"] for n in valid_nodes] name_to_node = {n["name"]: n for n in valid_nodes} for name in node_names: try: result = session.run(""" MATCH (start {名称: $name}) OPTIONAL MATCH path = (start)-[*1..3]-(end) WHERE start <> end RETURN coalesce(max(length(path)), 0) AS max_level """, name=name).single() if result: level = min(result["max_level"], 3) name_to_node[name]["level"] = level else: name_to_node[name]["level"] = 0 except Exception as e: logger.warning(f"Failed to compute level for {name}: {e}") name_to_node[name]["level"] = 0 res["message"] = "尾节点展开成功" return res # =============================== # 情况 5: 查询节点展开 # =============================== elif queryview and querylink: logger.info(f"Mode: QUERYVIEW + QUERYLINK | nodes={queryview}, links={querylink}") # 解析节点 ID node_ids = [x.strip() for x in queryview.split(",") if x.strip()] link_ids = [x.strip() for x in querylink.split(",") if x.strip()] if not node_ids and not link_ids: res["message"] = "queryview 和 querylink 均为空" return res # === Step 1: 获取 queryview 节点的 3 跳子图(graphData)=== graph_data_nodes = {} graph_data_links = {} if node_ids: nodes_result = session.run( "MATCH (n) WHERE elementId(n) IN $ids RETURN n", ids=node_ids ) seed_nodes = [rec["n"] for rec in nodes_result] if seed_nodes: paths = session.run( """ MATCH (seed) WHERE elementId(seed) IN $ids WITH collect(seed) AS seeds UNWIND seeds AS s MATCH path = (s)-[*0..3]-(m) RETURN path """, ids=node_ids ) for p_rec in paths: path = p_rec["path"] for node_1 in path.nodes: nid = node_1.element_id if nid not in graph_data_nodes: graph_data_nodes[nid] = node_to_json(node_1) for rel in path.relationships: rid = rel.element_id if rid not in graph_data_links: graph_data_links[rid] = rel_to_json(rel) # === Step 2: 获取 querylink 指定的边(用于 short_routes)=== short_links = {} node_ids_from_links = set() if link_ids: # 先查边,并收集两端节点的 elementId rels_result = session.run( """ MATCH ()-[r]->() WHERE elementId(r) IN $link_ids RETURN r, elementId(startNode(r)) AS start_id, elementId(endNode(r)) AS end_id """, link_ids=link_ids ) rel_records = [] for record in rels_result: rel = record["r"] rid = rel.element_id short_links[rid] = rel_to_json(rel) node_ids_from_links.add(record["start_id"]) node_ids_from_links.add(record["end_id"]) rel_records.append(record) # === 关键:用 node_ids_from_links 查询完整节点信息 === if node_ids_from_links: nodes_full_result = session.run( "MATCH (n) WHERE elementId(n) IN $node_ids RETURN n", node_ids=list(node_ids_from_links) ) short_nodes = { node.element_id: node_to_json(node) for record in nodes_full_result for node in [record["n"]] } else: short_nodes = {} else: short_nodes = {} short_links = {} # === Step 3: 合并逻辑(可选:确保 short_routes 节点也在 graphData 中?不强制)=== graph_data = { "nodes": list(graph_data_nodes.values()), "links": list(graph_data_links.values()) } short_routes = { "nodes": list(short_nodes.values()), "links": list(short_links.values()) } # === Step 4: 为 graphData 中的节点计算 level(保持不变)=== valid_nodes = [n for n in graph_data["nodes"] if n.get("name")] name_to_node = {n["name"]: n for n in valid_nodes if n.get("name")} for name in name_to_node: try: result = session.run(""" MATCH (start {名称: $name}) OPTIONAL MATCH path = (start)-[*1..3]-(end) WHERE start <> end RETURN coalesce(max(length(path)), 0) AS max_level """, name=name).single() level = min(result["max_level"], 3) if result else 0 name_to_node[name]["level"] = level except Exception as e: logger.warning(f"Failed to compute level for {name}: {e}") name_to_node[name]["level"] = 0 res["graphData"] = graph_data res["short_routes"] = short_routes res["message"] = f"查询拓展加载成功!" return res # =============================== # 兜底:无任何输入 # =============================== else: res["message"] = "查询成功!" return res except Exception as e: logger.exception("Internal db_task error") return { "graphData": {"nodes": [], "links": []}, "short_routes": {"nodes": [], "links": []}, "message": f"数据库查询异常: {str(e)}" } # 执行数据库任务 query_result = await run_in_threadpool(db_task) # 组装最终响应 if head or tail: simple_graphData=query_result["graphData"] if node and depth: simple_graphData=query_result["graphData"] if queryview and querylink: simple_graphData=query_result["graphData"] return { "message": query_result["message"], "simple_graphData": simple_graphData, "graphData": query_result["graphData"], "short_routes": query_result.get("short_routes", {"nodes": [], "links": []}) }