diff --git a/src/core/retriever.py b/src/core/retriever.py index 2071f3cb..bafa7f84 100644 --- a/src/core/retriever.py +++ b/src/core/retriever.py @@ -1,5 +1,6 @@ from src.models.embedding import Reranker from src.utils.logging_config import setup_logger + logger = setup_logger("server-common") @@ -37,12 +38,14 @@ class Retriever: # 解析图数据库的结果 db_res = refs.get("graph_base").get("results", []) if len(db_res["nodes"]) > 0: - db_text = '\n'.join([f"{edge['source_name']}和{edge['target_name']}的关系是{edge['type']}" for edge in db_res['edges']]) + db_text = "\n".join( + [f"{edge['source_name']}和{edge['target_name']}的关系是{edge['type']}" for edge in db_res["edges"]] + ) external += f"图数据库信息: \n\n{db_text}" # 构造查询 if len(external) > 0: - query = f"以下是参考资料:\n\n\n{external}\n\n\n请根据前面的知识回答:{query}" + query = f"参考资料:\n\n\n{external}\n\n\n请根据前面的知识回答问题。\n\n问题:{query}\n\n回答:" return query @@ -70,7 +73,12 @@ class Retriever: kb_res = [] final_res = [] if not refs["meta"].get("db_name") or not self.config.enable_knowledge_base: - return {"results": final_res, "all_results": kb_res, "rw_query": query, "message": "Knowledge base is disabled"} + return { + "results": final_res, + "all_results": kb_res, + "rw_query": query, + "message": "Knowledge base is disabled", + } rw_query = self.rewrite_query(query, history, refs) @@ -85,7 +93,7 @@ class Retriever: if self.config.enable_reranker: RERANK_THRESHOLD = 0.001 for r in kb_res: - r["rerank_score"] = self.reranker.compute_score([query, r["entity"]["text"]], normalize=True) + r["rerank_score"] = self.reranker.compute_score([rw_query, r["entity"]["text"]], normalize=True) kb_res.sort(key=lambda x: x["rerank_score"], reverse=True) final_res = [_res for _res in kb_res if _res["rerank_score"] > RERANK_THRESHOLD] @@ -101,7 +109,8 @@ class Retriever: rewritten_query = query else: from src.utils.prompts import rewritten_query_prompt_template - history_query = [entry['content'] for entry in history if entry['role'] == 'user'] if history else "" + + history_query = [entry["content"] for entry in history if entry["role"] == "user"] if history else "" rewritten_query_prompt = rewritten_query_prompt_template.format(history=history_query, query=query) rewritten_query = self.model.predict(rewritten_query_prompt).content @@ -118,54 +127,70 @@ class Retriever: entities = [] if refs["meta"].get("use_graph"): from src.utils.prompts import entity_extraction_prompt_template + entity_extraction_prompt = entity_extraction_prompt_template.format(text=query) entities = self.model.predict(entity_extraction_prompt).content.split(",") - entities = [entity for entity in entities if all(char.isalnum() or char in '汉字' for char in entity)] + entities = [entity for entity in entities if all(char.isalnum() or char in "汉字" for char in entity)] return entities - def foramt_general_results(self, results): + def _extract_relationship_info(self, relationship, source_name, target_name): + """ + 提取关系信息并返回格式化的节点和边信息 + """ + 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 + + 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): logger.debug(f"Formatting general results: {results}") formatted_results = {"nodes": [], "edges": []} for item in results: relationship = item[1] - rel_id = relationship.element_id - nodes = relationship.nodes - if len(nodes) != 2: + 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 - source, target = nodes + for node in node_info: + if node["id"] not in [n["id"] for n in formatted_results["nodes"]]: + formatted_results["nodes"].append(node) - source_id = source.element_id - target_id = target.element_id - source_name = source._properties.get('name', 'unknown') - target_name = target._properties.get('name', 'unknown') - - if source_id not in formatted_results["nodes"]: - formatted_results["nodes"].append({"id": source_id, "name": source_name}) - if target_id not in formatted_results["nodes"]: - formatted_results["nodes"].append({"id": target_id, "name": target_name}) - - relationship_type = relationship._properties.get('type', 'unknown') - if relationship_type == 'unknown': - relationship_type = relationship.type - - formatted_results["edges"].append({ - "id": rel_id, - "type": relationship_type, - "source_id": source_id, - "target_id": target_id, - "source_name": source_name, - "target_name": target_name - }) + formatted_results["edges"].append(edge_info) return formatted_results def format_query_results(self, results): logger.debug(f"Formatting query results: {results}") formatted_results = {"nodes": [], "edges": []} - node_dict = {} for item in results: @@ -173,35 +198,18 @@ class Retriever: continue relationship = item[1][0] - rel_id = relationship.element_id - nodes = relationship.nodes - if len(nodes) != 2: + source_name = item[0] + target_name = item[2] 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 - source, target = nodes + for node in node_info: + if node["id"] not in node_dict: + node_dict[node["id"]] = node - source_id = source.element_id - target_id = target.element_id - source_name = item[0] - target_name = item[2] if len(item) > 2 else 'unknown' - - if source_id not in node_dict: - node_dict[source_id] = {"id": source_id, "name": source_name} - if target_id not in node_dict: - node_dict[target_id] = {"id": target_id, "name": target_name} - - relationship_type = relationship._properties.get('type', 'unknown') - if relationship_type == 'unknown': - relationship_type = relationship.type - - formatted_results["edges"].append({ - "id": rel_id, - "type": relationship_type, - "source_id": source_id, - "target_id": target_id, - "source_name": source_name, - "target_name": target_name - }) + formatted_results["edges"].append(edge_info) formatted_results["nodes"] = list(node_dict.values()) @@ -210,4 +218,4 @@ class Retriever: def __call__(self, query, history, meta): refs = self.retrieval(query, history, meta) query = self.construct_query(query, refs, meta) - return query, refs \ No newline at end of file + return query, refs diff --git a/src/views/database_view.py b/src/views/database_view.py index 98adcfae..5f5057c6 100644 --- a/src/views/database_view.py +++ b/src/views/database_view.py @@ -135,7 +135,7 @@ def get_graph_nodes(): logger.debug(f"Get graph nodes in {kgdb_name} with {num} nodes") result = startup.dbm.graph_base.get_sample_nodes(kgdb_name, num) - return jsonify({'result': startup.retriever.foramt_general_results(result), 'message': 'success'}), 200 + return jsonify({'result': startup.retriever.format_general_results(result), 'message': 'success'}), 200 @db.route('/graph/add', methods=['POST']) def add_graph_entity():