418 lines
14 KiB
Python
418 lines
14 KiB
Python
# 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"]) |