kgrag/neo4j_indexing.py
Defeng b919725117 更新 neo4j_indexing.py
检索索引构建冗余代码
2026-07-07 16:18:57 +08:00

418 lines
14 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

# neo4j-admin database load --from-path="neo4j.dump文件所在目录" --overwrite-destination=true neo4j --verbose
import os
import re
import jieba
import logging
import numpy as np
from concurrent.futures import ThreadPoolExecutor
from neo4j import GraphDatabase
from neo4j_graphrag.indexes import (
create_vector_index,
upsert_vectors,
create_fulltext_index,
)
from modelsAPI.model_api import OpenaiAPI
from dotenv import load_dotenv
from typing import List, Set, Tuple, Dict, Optional
from config import LLM_CONFIG, EMBEDDING_CONFIG, NEO4J_CONFIG, SEARCH_CONFIG
# 加载 .env 文件中的环境变量
load_dotenv()
# 配置控制台日志
logger = logging.getLogger("indexing")
logger.setLevel(logging.INFO)
if not logger.handlers:
formatter = logging.Formatter("[%(levelname)s]%(asctime)s: %(message)s")
handler = logging.StreamHandler()
handler.setFormatter(formatter)
logger.addHandler(handler)
VECTOR_DIMENSION = SEARCH_CONFIG["vector_dimension"]
vector_dim = SEARCH_CONFIG["vector_dimension"] # OpenAI bge-m3 模型的嵌入向量维度
embed_batch_size = SEARCH_CONFIG["embed_batch_size"] # 嵌入向量计算批次大小
EMBED_MAX_WORKERS = SEARCH_CONFIG["embed_thread_num"] # embedding 批次并行线程数(受下游 OpenaiAPI 限流约束,勿过大)
NODE_FETCH_PAGE = SEARCH_CONFIG["node_fetch_page"] # 单次拉取待生成 embedding 节点的分页大小
# 全局索引名称(覆盖所有节点标签)
DEFAULT_URI = NEO4J_CONFIG["uri"]
DEFAULT_AUTH = (
NEO4J_CONFIG["username"],
NEO4J_CONFIG["password"],
)
NAME_PROPERTY = SEARCH_CONFIG["name_property"]
FULLTEXT_PROPERTY = SEARCH_CONFIG["fulltext_property"]
EMBEDDING_PROPERTY = SEARCH_CONFIG["embedding_property"]
SEARCH_LABEL = SEARCH_CONFIG["search_label"]
VECTOR_INDEX_NAME = SEARCH_CONFIG["vector_index_name"]
FULLTEXT_INDEX_NAME = SEARCH_CONFIG["fulltext_index_name"]
FULLTEXT_FIELD_INDEX_NAME = SEARCH_CONFIG["fulltext_field_index_name"]
EXCLUDED_BUSINESS_LABELS: Set[str] = set(SEARCH_CONFIG["excluded_business_labels"])
OLD_INDEX_NAMES: List[str] = [
VECTOR_INDEX_NAME,
FULLTEXT_INDEX_NAME,
FULLTEXT_FIELD_INDEX_NAME,
]
def get_all_labels(driver):
"""
获取知识图谱中所有的节点标签
Args:
driver: Neo4j driver
Returns:
list: 所有节点标签列表
"""
query = "CALL db.labels()"
result = driver.execute_query(query)
labels = [record["label"] for record in result.records]
return labels
def get_label_properties(driver, label):
"""
获取指定标签的所有属性键(通过采样多个节点获取所有可能的属性)
Args:
driver: Neo4j driver
label: 节点标签
Returns:
list: 属性键列表(排除内部属性如 embedding, fulltext
"""
# 采样多个节点来获取所有可能的属性(因为不同节点可能有不同的属性集)
query = f"""
MATCH (n:{label})
WITH n LIMIT 100
UNWIND keys(n) AS prop
RETURN DISTINCT prop AS properties
ORDER BY properties
"""
result = driver.execute_query(query)
if result.records:
properties = [record["properties"] for record in result.records]
# 排除内部索引属性
excluded_props = {"embedding", "fulltext", "elementId", "id"}
properties = [prop for prop in properties if prop not in excluded_props]
return properties
else:
logger.warning(f"标签 {label} 没有节点,无法获取属性列表")
return []
def build_text_from_all_properties(record_dict):
"""
从节点的所有属性构建文本(用于全文索引)
Args:
record_dict: 节点属性字典(已排除 id
Returns:
str: 构建的文本字符串
"""
text_parts = []
for key, value in record_dict.items():
if value is not None:
# 将值转换为字符串
if isinstance(value, (list, dict)):
value_str = str(value)
else:
value_str = str(value).strip()
if value_str:
text_parts.append(f"{key}{value_str}")
return " ".join(text_parts) if text_parts else None
def _compute_embeddings_parallel(names: List[str]) -> List[list]:
"""
并行计算一批名称的 embedding。
将 names 切成多个 batch用线程池并发调用 OpenaiAPI.batch_embeddings
再按原始顺序拼回。相比串行可显著缩短大批量节点的 embedding 耗时。
Args:
names: 待计算 embedding 的文本列表
Returns:
list: 与 names 等长、顺序一致的 embedding 向量列表
Raises:
任一 batch 计算失败时向上抛出异常
"""
# 切分批次,保留批次索引以便按序拼回
batches = [
(idx, names[i:i + embed_batch_size])
for idx, i in enumerate(range(0, len(names), embed_batch_size))
]
results: Dict[int, list] = {}
def _run_batch(item):
idx, batch_names = item
batch_result = OpenaiAPI.batch_embeddings(batch_names)
# batch_result 返回的是字典列表,每个字典包含 'embedding' 字段
return idx, [row["embedding"] for row in batch_result]
workers = min(EMBED_MAX_WORKERS, len(batches)) or 1
with ThreadPoolExecutor(max_workers=workers) as executor:
for idx, batch_embeddings in executor.map(_run_batch, batches):
results[idx] = batch_embeddings
# 按批次顺序拼回,保证与输入 names 对齐
all_embeddings: List[list] = []
for idx in range(len(batches)):
all_embeddings.extend(results[idx])
return all_embeddings
def drop_constraint(driver):
"""删除所有约束"""
records = driver.execute_query("show constraints").records
for record in records:
driver.execute_query(f"drop constraint {record['name']} if exists")
def drop_index_without_constraint(driver):
"""删除所有没有约束的索引"""
records = driver.execute_query("show index").records
for record in records:
if not record["owningConstraint"]:
driver.execute_query(f"drop index {record['name']} if exists")
def print_neo4j_version(driver) -> None:
"""打印当前 Neo4j 实例的版本和版本类型。"""
with driver.session() as session:
rec = session.run(
"CALL dbms.components() YIELD name, versions, edition "
"WHERE name='Neo4j Kernel' RETURN versions[0] AS v, edition"
).single()
if rec:
logger.info(f"Neo4j 版本: {rec['v']} ({rec['edition']})")
def migrate_remove_entity_label(driver, dry_run: bool = False) -> None:
"""
从所有节点上移除 :Entity 标签。
Parameters
----------
driver : Neo4j GraphDatabase driver
dry_run : 若为 True只统计数量不执行移除
"""
with driver.session() as session:
cnt = session.run("MATCH (n:Entity) RETURN count(n) AS c").single()["c"]
logger.info(f"当前带 :Entity 标签的节点数: {cnt}")
if cnt == 0:
return
if dry_run:
logger.info("[dry_run] 未执行移除")
return
session.run("MATCH (n:Entity) REMOVE n:Entity")
logger.info(f"已从 {cnt} 个节点上移除 :Entity")
def tag_searchable_nodes(driver, excluded: Set[str]) -> int:
"""
为满足以下条件的节点打上 :Searchable 标签:
- 拥有 `名称` 属性
- 拥有 `embedding` 属性
- 尚未持有 :Searchable 标签
- 至少有一个非排除业务标签
注意:排除标签列表通过查询参数传入,避免 Cypher 字符串注入。
Returns
-------
int : 本次新打标签的节点数
"""
# 标签名无法参数化,但排除标签集合的“值”可以参数化,避免拼接注入
query = f"""
MATCH (n)
WHERE n.`{NAME_PROPERTY}` IS NOT NULL
AND n.{EMBEDDING_PROPERTY} IS NOT NULL
AND NOT n:{SEARCH_LABEL}
AND any(lbl IN labels(n) WHERE NOT lbl IN $excluded)
SET n:{SEARCH_LABEL}
RETURN count(n) AS added
"""
excluded_list = list(excluded)
with driver.session() as session:
added = session.run(query, excluded=excluded_list).single()["added"]
logger.info(f"本次新打上 :{SEARCH_LABEL} 的节点数: {added}")
with driver.session() as session:
total = session.run(
f"MATCH (n:{SEARCH_LABEL}) RETURN count(n) AS c"
).single()["c"]
logger.info(f"当前带 :{SEARCH_LABEL} 标签的节点总数: {total}")
return added
def drop_old_indexes(driver, names: List[str]) -> None:
"""删除指定名称的索引(不存在时忽略)。"""
with driver.session() as session:
for name in names:
try:
session.run(f"DROP INDEX {name} IF EXISTS")
logger.info(f"已删除旧索引: {name}")
except Exception as e:
logger.warning(f"删除索引 {name} 失败(可忽略): {e}")
def create_indexes(driver) -> None:
"""
创建三个索引:
1. 向量索引 —— 用于语义相似度检索
2. 名称全文索引 —— 用于节点名称关键词检索
3. fulltext 字段全文索引 —— 用于节点 fulltext 属性关键词检索
"""
with driver.session() as session:
# 1. 向量索引
session.run(f"""
CREATE VECTOR INDEX {VECTOR_INDEX_NAME} IF NOT EXISTS
FOR (n:{SEARCH_LABEL}) ON (n.{EMBEDDING_PROPERTY})
OPTIONS {{
indexConfig: {{
`vector.dimensions`: {VECTOR_DIMENSION},
`vector.similarity_function`: 'cosine'
}}
}}
""")
logger.info(f"向量索引 '{VECTOR_INDEX_NAME}' 已就绪")
# 2. 名称全文索引
session.run(f"""
CREATE FULLTEXT INDEX {FULLTEXT_INDEX_NAME} IF NOT EXISTS
FOR (n:{SEARCH_LABEL}) ON EACH [n.`{NAME_PROPERTY}`]
""")
logger.info(f"全文索引(名称) '{FULLTEXT_INDEX_NAME}' 已就绪")
# 3. fulltext 字段全文索引
session.run(f"""
CREATE FULLTEXT INDEX {FULLTEXT_FIELD_INDEX_NAME} IF NOT EXISTS
FOR (n:{SEARCH_LABEL}) ON EACH [n.`{FULLTEXT_PROPERTY}`]
""")
logger.info(f"全文索引fulltext 字段) '{FULLTEXT_FIELD_INDEX_NAME}' 已就绪")
# ================== 连接解析工具 ==================
def _resolve_connection(
uri: Optional[str],
user: Optional[str],
password: Optional[str],
) -> Tuple[str, Tuple[str, str]]:
"""
按优先级解析连接参数:传参 > 环境变量 > 默认常量。
Returns
-------
(uri, (user, password))
"""
resolved_uri = uri or os.getenv("NEO4J_URI", DEFAULT_URI)
resolved_user = user or os.getenv("NEO4J_USER", DEFAULT_AUTH[0])
resolved_pw = password or os.getenv("NEO4J_PASSWORD", DEFAULT_AUTH[1])
return resolved_uri, (resolved_user, resolved_pw)
# ================== 主入口函数 ==================
def build_hybrid_indexes(
driver=None,
*,
uri: Optional[str] = None,
user: Optional[str] = None,
password: Optional[str] = None,
remove_entity_label: bool = False,
dry_run_entity: bool = False,
drop_old: bool = True,
) -> dict:
"""
完整的混合索引构建流程。
连接参数优先级:传参 > 环境变量(NEO4J_URI/NEO4J_USER/NEO4J_PASSWORD) > 默认常量
Parameters
----------
driver : 已有的 Neo4j driver传入则直接使用不再创建新连接
uri : Neo4j 连接地址,如 "bolt://host:7687"
user : 数据库用户名
password : 数据库密码
remove_entity_label : 是否执行移除 :Entity 标签的迁移步骤
dry_run_entity : remove_entity_label=True 时,是否仅统计而不实际移除
drop_old : 是否在创建前先删除旧索引(建议保持 True
Returns
-------
dict : {"success": bool, "message": str, "searchable_added": int}
失败时 success=False 且 message 含错误信息
"""
# 若未传入 driver则根据参数/环境变量/默认值自动创建
_owns_driver = driver is None
if _owns_driver:
resolved_uri, resolved_auth = _resolve_connection(uri, user, password)
logger.info(f"连接 Neo4j: {resolved_uri} 用户: {resolved_auth[0]}")
driver = GraphDatabase.driver(resolved_uri, auth=resolved_auth)
try:
logger.info("=" * 60)
logger.info("开始构建混合索引")
logger.info("=" * 60)
print_neo4j_version(driver)
if remove_entity_label:
migrate_remove_entity_label(driver, dry_run=dry_run_entity)
searchable_added = tag_searchable_nodes(driver, EXCLUDED_BUSINESS_LABELS)
if drop_old:
drop_old_indexes(driver, OLD_INDEX_NAMES)
create_indexes(driver)
return {
"success": True,
"message": "混合索引构建完成",
"searchable_added": searchable_added,
}
except Exception as e:
error_msg = f"构建混合索引过程中发生错误: {e}"
logger.error(error_msg, exc_info=True)
return {
"success": False,
"message": error_msg,
"searchable_added": 0,
}
finally:
# 仅关闭由本函数自己创建的 driver
if _owns_driver:
driver.close()
# 向后兼容别名:保留旧拼写,避免其他模块仍引用 build_hyrid_indexes 时报错
build_hyrid_indexes = build_hybrid_indexes
if __name__ == "__main__":
result = build_hybrid_indexes(
drop_old=True
)
if result["success"]:
logger.info("混合索引创建成功")
else:
logger.error(result["message"])