ForcePilot/src/core/graphbase.py

697 lines
27 KiB
Python
Raw Normal View History

import os
2024-07-15 18:58:34 +08:00
import json
import warnings
2025-03-19 01:21:51 +08:00
import traceback
2024-07-20 20:30:32 +08:00
import torch
2024-07-16 18:14:27 +08:00
from neo4j import GraphDatabase as GD
2024-07-20 20:30:32 +08:00
from src import config
2025-02-28 00:38:59 +08:00
from src.utils import logger
2024-09-06 12:54:17 +08:00
warnings.filterwarnings("ignore", category=UserWarning)
2024-09-09 17:07:03 +08:00
2024-07-24 18:32:14 +08:00
UIE_MODEL = None
2024-07-16 18:14:27 +08:00
class GraphDatabase:
def __init__(self):
2024-07-16 18:14:27 +08:00
self.driver = None
self.files = []
self.status = "closed"
self.kgdb_name = "neo4j"
2025-02-23 16:39:52 +08:00
self.embed_model_name = None
self.work_dir = os.path.join(config.save_dir, "knowledge_graph", self.kgdb_name)
2025-02-28 00:38:59 +08:00
os.makedirs(self.work_dir, exist_ok=True)
2025-02-28 02:40:34 +08:00
# 尝试加载已保存的图数据库信息
2025-03-11 14:18:48 +08:00
if not self.load_graph_info():
2025-04-28 22:53:13 +08:00
logger.debug(f"未找到已保存的图数据库信息,将创建新的配置")
2025-02-28 02:40:34 +08:00
2025-02-28 00:38:59 +08:00
self.start()
2024-07-16 18:14:27 +08:00
def start(self):
if not config.enable_knowledge_graph or not config.enable_knowledge_base:
return
2024-09-11 01:08:13 +08:00
uri = os.environ.get("NEO4J_URI", "bolt://localhost:7687")
username = os.environ.get("NEO4J_USERNAME", "neo4j")
password = os.environ.get("NEO4J_PASSWORD", "0123456789")
logger.info(f"Connecting to Neo4j at {uri}/{self.kgdb_name}")
2024-10-08 22:16:17 +08:00
try:
self.driver = GD.driver(f"{uri}/{self.kgdb_name}", auth=(username, password))
self.status = "open"
logger.info(f"Connected to Neo4j at {uri}/{self.kgdb_name}, {self.get_graph_info(self.kgdb_name)}")
2025-02-28 02:40:34 +08:00
# 连接成功后保存图数据库信息
self.save_graph_info(self.kgdb_name)
2024-10-08 22:16:17 +08:00
except Exception as e:
logger.error(f"Failed to connect to Neo4j: {e}, {uri}, {self.kgdb_name}, {username}, {password}")
self.config.enable_knowledge_graph = False
2024-07-15 18:58:34 +08:00
def close(self):
"""关闭数据库连接"""
self.driver.close()
def is_running(self):
"""检查图数据库是否正在运行"""
if not config.enable_knowledge_graph or not config.enable_knowledge_base:
return False
else:
return self.status == "open"
2024-09-06 12:54:17 +08:00
def get_sample_nodes(self, kgdb_name='neo4j', num=50):
2024-09-09 17:07:03 +08:00
"""获取指定数据库的 num 个节点信息"""
2024-09-06 12:54:17 +08:00
self.use_database(kgdb_name)
def query(tx, num):
result = tx.run("MATCH (n)-[r]->(m) RETURN n, r, m LIMIT $num", num=int(num))
return result.values()
with self.driver.session() as session:
return session.execute_read(query, num)
2024-07-15 18:58:34 +08:00
def create_graph_database(self, kgdb_name):
"""创建新的数据库,如果已存在则返回已有数据库的名称"""
with self.driver.session() as session:
existing_databases = session.run("SHOW DATABASES")
existing_db_names = [db['name'] for db in existing_databases]
2024-07-16 18:14:27 +08:00
2024-07-15 18:58:34 +08:00
if existing_db_names:
print(f"已存在数据库: {existing_db_names[0]}")
return existing_db_names[0] # 返回所有已有数据库名称
2024-07-16 18:14:27 +08:00
2024-07-15 18:58:34 +08:00
session.run(f"CREATE DATABASE {kgdb_name}")
print(f"数据库 '{kgdb_name}' 创建成功.")
return kgdb_name # 返回创建的数据库名称
2024-07-16 18:14:27 +08:00
2024-09-14 02:44:53 +08:00
def use_database(self, kgdb_name="neo4j"):
2024-07-15 18:58:34 +08:00
"""切换到指定数据库"""
2024-09-14 02:44:53 +08:00
assert kgdb_name == self.kgdb_name, f"传入的数据库名称 '{kgdb_name}' 与当前实例的数据库名称 '{self.kgdb_name}' 不一致"
if self.status == "closed":
self.start()
2024-07-15 18:58:34 +08:00
2024-07-17 18:50:01 +08:00
def txt_add_entity(self, triples, kgdb_name='neo4j'):
2024-07-15 18:58:34 +08:00
"""添加实体三元组"""
self.use_database(kgdb_name)
def create(tx, triples):
for triple in triples:
h = triple['h']
t = triple['t']
r = triple['r']
query = (
"MERGE (a:Entity {name: $h}) "
"MERGE (b:Entity {name: $t}) "
"MERGE (a)-[:" + r.replace(" ", "_") + "]->(b)"
)
tx.run(query, h=h, t=t)
with self.driver.session() as session:
session.execute_write(create, triples)
2025-04-23 21:18:39 +08:00
async def txt_add_vector_entity(self, triples, kgdb_name='neo4j'):
2024-07-20 20:30:32 +08:00
"""添加实体三元组"""
self.use_database(kgdb_name)
def _index_exists(tx, index_name):
2025-02-28 00:38:59 +08:00
"""检查索引是否存在"""
2024-07-20 20:30:32 +08:00
result = tx.run("SHOW INDEXES")
for record in result:
if record["name"] == index_name:
return True
return False
2025-02-28 00:38:59 +08:00
2024-07-20 20:30:32 +08:00
def _create_graph(tx, data):
2025-02-28 00:38:59 +08:00
"""添加一个三元组"""
2024-07-20 20:30:32 +08:00
for entry in data:
tx.run("""
MERGE (h:Entity {name: $h})
MERGE (t:Entity {name: $t})
MERGE (h)-[r:RELATION {type: $r}]->(t)
""", h=entry['h'], t=entry['t'], r=entry['r'])
2025-02-28 00:38:59 +08:00
2024-09-06 12:54:17 +08:00
def _create_vector_index(tx, dim):
2025-02-28 00:38:59 +08:00
"""创建向量索引"""
# NOTE 这里是否是会重复构建索引?
2024-09-06 12:54:17 +08:00
index_name = "entityEmbeddings"
2024-07-20 20:30:32 +08:00
if not _index_exists(tx, index_name):
tx.run(f"""
CREATE VECTOR INDEX {index_name}
FOR (n: Entity) ON (n.embedding)
OPTIONS {{indexConfig: {{
2024-09-06 12:54:17 +08:00
`vector.dimensions`: {dim},
2024-07-20 20:30:32 +08:00
`vector.similarity_function`: 'cosine'
}} }};
""")
2024-09-06 12:54:17 +08:00
def _get_nodes_without_embedding(tx, entity_names):
"""获取没有embedding的节点列表"""
# 构建参数字典,将列表转换为"param0"、"param1"等键值对形式
params = {f"param{i}": name for i, name in enumerate(entity_names)}
# 构建查询参数列表
param_placeholders = ", ".join([f"${key}" for key in params.keys()])
# 执行查询
result = tx.run(f"""
MATCH (n:Entity)
WHERE n.name IN [{param_placeholders}] AND n.embedding IS NULL
RETURN n.name AS name
""", params)
return [record["name"] for record in result]
2025-04-23 21:18:39 +08:00
def _batch_set_embeddings(tx, entity_embedding_pairs):
"""批量设置实体的嵌入向量"""
for entity_name, embedding in entity_embedding_pairs:
tx.run("""
MATCH (e:Entity {name: $name})
CALL db.create.setNodeVectorProperty(e, 'embedding', $embedding)
""", name=entity_name, embedding=embedding)
2025-02-28 00:38:59 +08:00
# 判断模型名称是否匹配
cur_embed_info = config.embed_model_names[config.embed_model]
2025-02-28 00:38:59 +08:00
self.embed_model_name = self.embed_model_name or cur_embed_info.get('name')
assert self.embed_model_name == cur_embed_info.get('name') or self.embed_model_name is None, \
f"embed_model_name={self.embed_model_name}, {cur_embed_info.get('name')=}"
2024-07-20 20:30:32 +08:00
with self.driver.session() as session:
2025-02-28 00:38:59 +08:00
logger.info(f"Adding entity to {kgdb_name}")
2024-07-20 20:30:32 +08:00
session.execute_write(_create_graph, triples)
logger.info(f"Creating vector index for {kgdb_name} with {config.embed_model}")
2025-02-28 00:38:59 +08:00
session.execute_write(_create_vector_index, cur_embed_info['dimension'])
2025-04-23 21:18:39 +08:00
# 收集所有需要处理的实体名称,去重
all_entities = []
for entry in triples:
if entry['h'] not in all_entities:
all_entities.append(entry['h'])
if entry['t'] not in all_entities:
all_entities.append(entry['t'])
# 筛选出没有embedding的节点
nodes_without_embedding = session.execute_read(_get_nodes_without_embedding, all_entities)
if not nodes_without_embedding:
logger.info(f"所有实体已有embedding无需重新计算")
return
logger.info(f"需要为{len(nodes_without_embedding)}/{len(all_entities)}个实体计算embedding")
2025-04-23 21:18:39 +08:00
# 批量处理实体
max_batch_size = 1024 # 限制此部分的主要是内存大小 1024 * 1024 * 4 / 1024 / 1024 = 4GB
total_entities = len(nodes_without_embedding)
2025-04-23 21:18:39 +08:00
for i in range(0, total_entities, max_batch_size):
batch_entities = nodes_without_embedding[i:i+max_batch_size]
2025-04-23 21:18:39 +08:00
logger.debug(f"Processing entities batch {i//max_batch_size + 1}/{(total_entities-1)//max_batch_size + 1} ({len(batch_entities)} entities)")
# 批量获取嵌入向量
batch_embeddings = await self.aget_embedding(batch_entities)
# 将实体名称和嵌入向量配对
entity_embedding_pairs = list(zip(batch_entities, batch_embeddings))
# 批量写入数据库
session.execute_write(_batch_set_embeddings, entity_embedding_pairs)
2024-07-20 20:30:32 +08:00
2025-02-28 02:40:34 +08:00
# 数据添加完成后保存图信息
self.save_graph_info()
2025-04-23 21:18:39 +08:00
async def jsonl_file_add_entity(self, file_path, kgdb_name='neo4j'):
2024-07-24 18:32:14 +08:00
self.status = "processing"
2024-10-14 16:51:20 +08:00
kgdb_name = kgdb_name or 'neo4j'
2024-07-20 20:30:32 +08:00
self.use_database(kgdb_name) # 切换到指定数据库
2025-02-28 00:38:59 +08:00
logger.info(f"Start adding entity to {kgdb_name} with {file_path}")
2024-07-24 18:32:14 +08:00
2024-07-20 20:30:32 +08:00
def read_triples(file_path):
with open(file_path, 'r', encoding='utf-8') as file:
for line in file:
2025-03-30 11:09:46 +08:00
if line.strip():
yield json.loads(line.strip())
triples = list(read_triples(file_path))
2025-04-23 21:18:39 +08:00
await self.txt_add_vector_entity(triples, kgdb_name)
2024-09-06 12:54:17 +08:00
2024-07-24 18:32:14 +08:00
self.status = "open"
2025-02-28 02:40:34 +08:00
# 更新并保存图数据库信息
self.save_graph_info()
2024-07-20 20:30:32 +08:00
return kgdb_name
2024-07-15 18:58:34 +08:00
def delete_entity(self, entity_name=None, kgdb_name="neo4j"):
"""删除数据库中的指定实体三元组, 参数entity_name为空则删除全部实体"""
self.use_database(kgdb_name)
with self.driver.session() as session:
if entity_name:
session.execute_write(self._delete_specific_entity, entity_name)
else:
session.execute_write(self._delete_all_entities)
def _delete_specific_entity(self, tx, entity_name):
query = """
MATCH (n {name: $entity_name})
DETACH DELETE n
"""
tx.run(query, entity_name=entity_name)
def _delete_all_entities(self, tx):
query = """
MATCH (n)
DETACH DELETE n
"""
tx.run(query)
2025-04-13 21:06:23 +08:00
def query_node(self, entity_name, threshold=0.9, kgdb_name='neo4j', hops=2, max_entities=5, **kwargs):
2025-05-07 00:38:33 +08:00
"""知识图谱查询节点的入口:"""
2024-09-14 02:44:53 +08:00
# TODO 添加判断节点数量为 0 停止检索
# 判断是否启动
if not self.is_running():
raise Exception("图数据库未启动")
2025-02-28 02:40:34 +08:00
2024-07-15 18:58:34 +08:00
self.use_database(kgdb_name)
def _index_exists(tx, index_name):
"""检查索引是否存在"""
result = tx.run("SHOW INDEXES")
for record in result:
if record["name"] == index_name:
return True
return False
2024-09-14 02:44:53 +08:00
def query(tx, text):
# 首先检查索引是否存在
if not _index_exists(tx, "entityEmbeddings"):
raise Exception("向量索引不存在,请先创建索引")
2024-09-14 02:44:53 +08:00
embedding = self.get_embedding(text)
result = tx.run("""
CALL db.index.vector.queryNodes('entityEmbeddings', 10, $embedding)
YIELD node AS similarEntity, score
RETURN similarEntity.name AS name, score
""", embedding=embedding)
2024-07-15 18:58:34 +08:00
return result.values()
try:
with self.driver.session() as session:
results = session.execute_read(query, entity_name)
except Exception as e:
if "向量索引不存在" in str(e):
logger.error(f"向量索引不存在,请先创建索引: {e}, {traceback.format_exc()}")
return []
raise e
2025-02-28 00:38:59 +08:00
# 筛选出分数高于阈值的实体
2025-04-13 21:06:23 +08:00
qualified_entities = [result[0] for result in results[:max_entities] if result[1] > threshold]
2025-02-28 00:38:59 +08:00
logger.debug(f"Graph Query Entities: {entity_name}, {qualified_entities=}")
# 对每个合格的实体进行查询
all_query_results = []
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)
return all_query_results
2024-07-15 18:58:34 +08:00
2025-04-13 21:06:23 +08:00
def query_specific_entity(self, entity_name, kgdb_name='neo4j', hops=2, limit=100):
2025-02-28 00:38:59 +08:00
"""查询指定实体三元组信息(无向关系)"""
2025-04-13 21:06:23 +08:00
if not entity_name:
logger.warning("实体名称为空")
return []
2024-07-15 18:58:34 +08:00
self.use_database(kgdb_name)
2025-04-13 21:06:23 +08:00
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
LIMIT $limit
"""
result = tx.run(query_str, entity_name=entity_name, limit=limit)
if not result:
logger.info(f"未找到实体 {entity_name} 的相关信息")
return []
values = result.values()
# 安全地处理embedding属性
values = clean_triples_embedding(values)
return values
except Exception as e:
logger.error(f"查询实体 {entity_name} 失败: {str(e)}")
return []
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 []
2024-07-15 18:58:34 +08:00
2024-09-14 02:44:53 +08:00
def query_all_nodes_and_relationships(self, kgdb_name='neo4j', hops = 2):
2025-04-13 21:06:23 +08:00
"""查询图数据库中所有三元组信息 NEVER USE"""
2024-09-14 02:44:53 +08:00
self.use_database(kgdb_name)
def query(tx, hops):
result = tx.run(f"""
MATCH (n)-[r*1..{hops}]->(m)
2025-04-13 21:06:23 +08:00
RETURN n AS n, r, m AS m
2024-09-14 02:44:53 +08:00
""")
2025-04-13 21:06:23 +08:00
values = result.values()
values = clean_triples_embedding(values)
return values
2024-09-14 02:44:53 +08:00
with self.driver.session() as session:
return session.execute_read(query, hops)
2024-07-17 18:50:01 +08:00
def query_by_relationship_type(self, relationship_type, kgdb_name='neo4j', hops = 2):
2025-04-13 21:06:23 +08:00
"""查询指定关系三元组信息 NEVER USE"""
2024-07-15 18:58:34 +08:00
self.use_database(kgdb_name)
2024-07-16 12:27:46 +08:00
def query(tx, relationship_type, hops):
2024-07-15 18:58:34 +08:00
result = tx.run(f"""
2024-07-16 12:27:46 +08:00
MATCH (n)-[r:`{relationship_type}`*1..{hops}]->(m)
2025-04-13 21:06:23 +08:00
RETURN n AS n, r, m AS m
2024-07-15 18:58:34 +08:00
""")
2025-04-13 21:06:23 +08:00
values = result.values()
values = clean_triples_embedding(values)
return values
2024-07-15 18:58:34 +08:00
with self.driver.session() as session:
2024-07-17 19:34:05 +08:00
return session.execute_read(query, relationship_type, hops)
2024-07-15 18:58:34 +08:00
2024-07-17 18:50:01 +08:00
def query_entity_like(self, keyword, kgdb_name='neo4j', hops = 2):
2025-04-13 21:06:23 +08:00
"""模糊查询 NEVER USE"""
2024-07-15 18:58:34 +08:00
self.use_database(kgdb_name)
2024-07-16 12:27:46 +08:00
def query(tx, keyword, hops):
result = tx.run(f"""
2024-07-15 18:58:34 +08:00
MATCH (n:Entity)
WHERE n.name CONTAINS $keyword
2024-07-16 12:27:46 +08:00
MATCH (n)-[r*1..{hops}]->(m)
2025-04-13 21:06:23 +08:00
RETURN n AS n, r, m AS m
2024-07-15 18:58:34 +08:00
""", keyword=keyword)
2025-04-13 21:06:23 +08:00
values = result.values()
values = clean_triples_embedding(values)
return values
2024-07-15 18:58:34 +08:00
with self.driver.session() as session:
2024-07-17 19:34:05 +08:00
return session.execute_read(query, keyword, hops)
2024-07-24 18:32:14 +08:00
2024-07-17 18:50:01 +08:00
def query_node_info(self, node_name, kgdb_name='neo4j', hops = 2):
2025-04-13 21:06:23 +08:00
"""查询指定节点的详细信息返回信息 NEVER USE"""
2024-07-15 18:58:34 +08:00
self.use_database(kgdb_name) # 切换到指定数据库
2024-07-16 12:27:46 +08:00
def query(tx, node_name, hops):
result = tx.run(f"""
2024-07-16 18:14:27 +08:00
MATCH (n {{name: $node_name}})
OPTIONAL MATCH (n)-[r*1..{hops}]->(m)
2025-04-13 21:06:23 +08:00
RETURN n AS n, r, m AS m
2024-07-15 18:58:34 +08:00
""", node_name=node_name)
2025-04-13 21:06:23 +08:00
values = result.values()
values = clean_triples_embedding(values)
return values
2024-07-15 18:58:34 +08:00
with self.driver.session() as session:
2024-07-17 19:34:05 +08:00
return session.execute_read(query, node_name, hops)
2024-07-24 18:32:14 +08:00
2025-04-23 21:18:39 +08:00
async def aget_embedding(self, text):
from src import knowledge_base
if isinstance(text, list):
outputs = await knowledge_base.embed_model.abatch_encode(text, batch_size=40)
return outputs
else:
outputs = await knowledge_base.embed_model.aencode(text)
return outputs
2024-07-20 20:30:32 +08:00
def get_embedding(self, text):
2025-04-23 21:18:39 +08:00
from src import knowledge_base
if isinstance(text, list):
outputs = knowledge_base.embed_model.batch_encode(text, batch_size=40)
return outputs
else:
outputs = knowledge_base.embed_model.encode([text])[0]
2025-02-28 00:38:59 +08:00
return outputs
2024-07-24 18:32:14 +08:00
2024-07-20 20:30:32 +08:00
def set_embedding(self, tx, entity_name, embedding):
tx.run("""
MATCH (e:Entity {name: $name})
CALL db.create.setNodeVectorProperty(e, 'embedding', $embedding)
""", name=entity_name, embedding=embedding)
2024-07-24 18:32:14 +08:00
def get_graph_info(self, graph_name="neo4j"):
self.use_database(graph_name)
def query(tx):
entity_count = tx.run("MATCH (n) RETURN count(n) AS count").single()["count"]
relationship_count = tx.run("MATCH ()-[r]->() RETURN count(r) AS count").single()["count"]
triples_count = tx.run("MATCH (n)-[r]->(m) RETURN count(n) AS count").single()["count"]
# 获取所有标签
labels = tx.run("CALL db.labels() YIELD label RETURN collect(label) AS labels").single()["labels"]
return {
"graph_name": graph_name,
"entity_count": entity_count,
"relationship_count": relationship_count,
"triples_count": triples_count,
"labels": labels,
2025-02-28 02:40:34 +08:00
"status": self.status,
"embed_model_name": self.embed_model_name,
"unindexed_node_count": self.query_nodes_without_embedding(graph_name)
2025-02-28 02:40:34 +08:00
}
try:
if self.status == "open" and self.driver and self.is_running():
# 获取数据库信息
with self.driver.session() as session:
graph_info = session.execute_read(query)
2025-04-01 18:37:30 +08:00
# 添加时间戳
from datetime import datetime
graph_info["last_updated"] = datetime.now().isoformat()
return graph_info
except Exception as e:
logger.error(f"获取图数据库信息失败:{e}, {traceback.format_exc()}")
return None
def save_graph_info(self, graph_name="neo4j"):
"""
将图数据库的基本信息保存到工作目录中的JSON文件
保存的信息包括数据库名称状态嵌入模型名称等
"""
try:
graph_info = self.get_graph_info(graph_name)
if graph_info is None:
logger.error(f"图数据库信息为空,无法保存")
return False
2025-02-28 02:40:34 +08:00
info_file_path = os.path.join(self.work_dir, "graph_info.json")
with open(info_file_path, 'w', encoding='utf-8') as f:
json.dump(graph_info, f, ensure_ascii=False, indent=2)
2025-04-14 00:15:30 +08:00
# logger.info(f"图数据库信息已保存到:{info_file_path}")
2025-02-28 02:40:34 +08:00
return True
except Exception as e:
logger.error(f"保存图数据库信息失败:{e}")
return False
2025-03-08 22:49:13 +08:00
def query_nodes_without_embedding(self, kgdb_name='neo4j'):
"""查询没有嵌入向量的节点
Returns:
list: 没有嵌入向量的节点列表
"""
self.use_database(kgdb_name)
def query(tx):
result = tx.run("""
MATCH (n:Entity)
WHERE n.embedding IS NULL
RETURN n.name AS name
""")
return [record["name"] for record in result]
with self.driver.session() as session:
return session.execute_read(query)
2025-02-28 02:40:34 +08:00
def load_graph_info(self):
"""
从工作目录中的JSON文件加载图数据库的基本信息
返回True表示加载成功False表示加载失败
"""
try:
info_file_path = os.path.join(self.work_dir, "graph_info.json")
if not os.path.exists(info_file_path):
2025-04-28 22:53:13 +08:00
logger.debug(f"图数据库信息文件不存在:{info_file_path}")
2025-02-28 02:40:34 +08:00
return False
with open(info_file_path, 'r', encoding='utf-8') as f:
graph_info = json.load(f)
# 更新对象属性
if graph_info.get("embed_model_name"):
self.embed_model_name = graph_info["embed_model_name"]
# 如果需要,可以加载更多信息
# 注意这里不更新self.kgdb_name因为它是在初始化时设置的
logger.info(f"已加载图数据库信息,最后更新时间:{graph_info.get('last_updated')}")
return True
except Exception as e:
logger.error(f"加载图数据库信息失败:{e}")
return False
2025-03-08 22:49:13 +08:00
def add_embedding_to_nodes(self, node_names=None, kgdb_name='neo4j'):
"""为节点添加嵌入向量
Args:
node_names (list, optional): 要添加嵌入向量的节点名称列表None表示所有没有嵌入向量的节点
kgdb_name (str, optional): 图数据库名称默认为'neo4j'
Returns:
int: 成功添加嵌入向量的节点数量
"""
self.use_database(kgdb_name)
# 如果node_names为None则获取所有没有嵌入向量的节点
if node_names is None:
node_names = self.query_nodes_without_embedding(kgdb_name)
count = 0
with self.driver.session() as session:
for node_name in node_names:
try:
embedding = self.get_embedding(node_name)
session.execute_write(self.set_embedding, node_name, embedding)
count += 1
except Exception as e:
2025-03-19 01:21:51 +08:00
logger.error(f"为节点 '{node_name}' 添加嵌入向量失败: {e}, {traceback.format_exc()}")
2025-03-08 22:49:13 +08:00
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 = 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": []}
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)
return formatted_results
def format_query_result_to_graph(self, query_results):
"""将检索到的结果转换为 {"nodes": [], "edges": []} 的格式
例如
{
"nodes": [
{
"id": "4:5efbff88-72ef-44f9-b867-6c0e164a4a13:103",
"name": "张若锦"
},
{
"id": "4:5efbff88-72ef-44f9-b867-6c0e164a4a13:20",
"name": "贾宝玉"
},
....
],
"edges": [
{
"id": "5:5efbff88-72ef-44f9-b867-6c0e164a4a13:71",
"type": "奴仆",
"source_id": "4:5efbff88-72ef-44f9-b867-6c0e164a4a13:88",
"target_id": "4:5efbff88-72ef-44f9-b867-6c0e164a4a13:20",
"source_name": "宋嬷嬷",
"target_name": "贾宝玉"
},
....
]
}
"""
formatted_results = {"nodes": [], "edges": []}
node_dict = {}
edge_dict = {}
for item in query_results:
# 检查数据格式
if len(item) < 2 or not isinstance(item[1], list):
continue
node_dict[item[0].element_id] = dict(id=item[0].element_id, name=item[0]._properties.get("name", "Unknown"))
node_dict[item[2].element_id] = dict(id=item[2].element_id, name=item[2]._properties.get("name", "Unknown"))
# 处理关系列表中的每个关系
for i, relationship in enumerate(item[1]):
try:
# 提取关系信息
node_info, edge_info = self._extract_relationship_info(relationship, node_dict=node_dict)
if node_info is None or edge_info is None:
continue
# 添加边
edge_dict[edge_info["id"]] = edge_info
except Exception as e:
logger.error(f"处理关系时出错: {e}, 关系: {relationship}, {traceback.format_exc()}")
continue
# 将节点字典转换为列表
formatted_results["nodes"] = list(node_dict.values())
formatted_results["edges"] = list(edge_dict.values())
return formatted_results
2025-04-13 21:06:23 +08:00
def clean_triples_embedding(triples):
for item in triples:
if hasattr(item[0], '_properties'):
item[0]._properties['embedding'] = None
if hasattr(item[2], '_properties'):
item[2]._properties['embedding'] = None
return triples
2024-07-15 18:58:34 +08:00
if __name__ == "__main__":
2025-03-30 11:09:46 +08:00
pass