808 lines
29 KiB
Python
808 lines
29 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 节点的分页大小
|
||
|
||
# 全局索引名称(覆盖所有节点标签)
|
||
GLOBAL_VECTOR_INDEX_NAME = SEARCH_CONFIG["global_entity_embedding"]
|
||
GLOBAL_FULLTEXT_INDEX_NAME = SEARCH_CONFIG["global_entity_content_search"]
|
||
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,
|
||
GLOBAL_VECTOR_INDEX_NAME,
|
||
GLOBAL_FULLTEXT_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 create_global_indexes(driver, force_refresh=False):
|
||
"""
|
||
创建覆盖所有节点标签的全局向量索引和全文索引。
|
||
与按标签创建的索引互补,支持跨标签统一检索。
|
||
|
||
Args:
|
||
driver: Neo4j driver
|
||
force_refresh: True=先删除旧索引再重建,False=IF NOT EXISTS 跳过已有索引
|
||
"""
|
||
if force_refresh:
|
||
for index_name in (GLOBAL_VECTOR_INDEX_NAME, GLOBAL_FULLTEXT_INDEX_NAME):
|
||
try:
|
||
driver.execute_query(f"DROP INDEX {index_name} IF EXISTS")
|
||
logger.info(f"强制刷新:已删除全局索引 {index_name}")
|
||
except Exception as e:
|
||
logger.warning(f"删除全局索引 {index_name} 时出错(可能不存在): {e}")
|
||
|
||
# 全局向量索引(作用于所有节点的 embedding 属性)
|
||
try:
|
||
driver.execute_query(f"""
|
||
CREATE VECTOR INDEX {GLOBAL_VECTOR_INDEX_NAME} IF NOT EXISTS
|
||
FOR (n) ON n.embedding
|
||
OPTIONS {{indexConfig: {{
|
||
`vector.dimensions`: {vector_dim},
|
||
`vector.similarity_function`: 'cosine'
|
||
}}}}
|
||
""")
|
||
logger.info(f"全局向量索引 '{GLOBAL_VECTOR_INDEX_NAME}' 已就绪")
|
||
except Exception as e:
|
||
logger.error(f"创建全局向量索引失败: {e}")
|
||
raise
|
||
|
||
# 全局全文索引(作用于所有节点的 名称 和 内容 属性)
|
||
try:
|
||
driver.execute_query(f"""
|
||
CREATE FULLTEXT INDEX {GLOBAL_FULLTEXT_INDEX_NAME} IF NOT EXISTS
|
||
FOR (n) ON EACH [n.名称, n.内容]
|
||
""")
|
||
logger.info(f"全局全文索引 '{GLOBAL_FULLTEXT_INDEX_NAME}' 已就绪")
|
||
except Exception as e:
|
||
logger.error(f"创建全局全文索引失败: {e}")
|
||
raise
|
||
|
||
|
||
# --------- 创建向量索引 ---------
|
||
def vector_indexing(driver, label, force_refresh=False):
|
||
"""
|
||
创建向量索引(如果节点没有 embedding,使用"名称"属性生成 embedding)
|
||
|
||
说明:
|
||
- 待生成 embedding 的节点通过分页循环全部处理,不再受单次 LIMIT 限制,
|
||
避免节点数超过一万时被静默丢弃。
|
||
- embedding 计算采用并行批次以缩短耗时。
|
||
|
||
Args:
|
||
driver: Neo4j driver
|
||
label: 节点标签
|
||
force_refresh: 是否强制刷新(如果索引已存在,force_refresh=False 时会跳过,True 时会先删除再创建)
|
||
"""
|
||
|
||
# === 分页处理所有缺失 embedding 的节点 ===
|
||
# 每轮取一页未生成 embedding 的节点并写回,直到没有剩余节点为止。
|
||
# 这样即使待处理节点远超一万也能全部覆盖。
|
||
total_generated = 0
|
||
while True:
|
||
query_no_embedding = f"""
|
||
MATCH (n:{label})
|
||
WHERE n.embedding IS NULL
|
||
AND n.名称 IS NOT NULL
|
||
RETURN elementId(n) AS id, n.名称 AS 名称
|
||
LIMIT {NODE_FETCH_PAGE}
|
||
"""
|
||
result_no_embedding = driver.execute_query(query_no_embedding)
|
||
nodes_to_process = [
|
||
(record["id"], record["名称"]) for record in result_no_embedding.records
|
||
]
|
||
|
||
if not nodes_to_process:
|
||
break
|
||
|
||
logger.info(f"标签 {label} 本轮发现 {len(nodes_to_process)} 个节点需要生成 embedding")
|
||
|
||
ids, names = zip(*nodes_to_process)
|
||
ids = list(ids)
|
||
names = list(names)
|
||
|
||
# 并行计算 embedding
|
||
logger.info(f"开始为 {len(names)} 个节点计算 embedding(并行批次)...")
|
||
try:
|
||
all_embeddings = _compute_embeddings_parallel(names)
|
||
except Exception as e:
|
||
logger.error(f"批量嵌入向量计算失败: {e}")
|
||
raise
|
||
|
||
# 转换为 numpy 数组并做 L2 归一化
|
||
embeddings = np.array(all_embeddings)
|
||
norms = np.linalg.norm(embeddings, axis=1, keepdims=True)
|
||
norms = np.where(norms == 0, 1, norms) # 避免除零
|
||
embeddings = embeddings / norms
|
||
|
||
# 保存 embedding 到节点
|
||
logger.info(f"写入 {len(ids)} 个节点的 embedding...")
|
||
upsert_vectors(
|
||
driver,
|
||
ids=ids,
|
||
embedding_property="embedding",
|
||
embeddings=embeddings,
|
||
)
|
||
total_generated += len(ids)
|
||
logger.info(f"本轮成功为 {len(ids)} 个节点生成并保存 embedding")
|
||
|
||
# 不足一页说明已处理完,提前结束
|
||
if len(nodes_to_process) < NODE_FETCH_PAGE:
|
||
break
|
||
|
||
if total_generated:
|
||
logger.info(f"标签 {label} 共生成 {total_generated} 个节点的 embedding")
|
||
|
||
# 检查是否有节点存在 embedding(包括刚生成的)
|
||
check_query = f"MATCH (n:{label}) WHERE n.embedding IS NOT NULL RETURN count(n) AS count"
|
||
result = driver.execute_query(check_query)
|
||
node_count = result.records[0]["count"] if result.records else 0
|
||
|
||
if node_count == 0:
|
||
logger.warning(f"标签 {label} 没有节点包含 embedding 属性,跳过向量索引创建")
|
||
return
|
||
|
||
# 创建向量索引
|
||
index_name = f"{label.lower()}_vector"
|
||
|
||
# 如果强制刷新且索引已存在,先删除它
|
||
if force_refresh:
|
||
try:
|
||
driver.execute_query(f"DROP INDEX {index_name} IF EXISTS")
|
||
logger.info(f"强制刷新模式:已删除旧索引 {index_name}")
|
||
except Exception as e:
|
||
logger.warning(f"删除索引 {index_name} 时出错(可能不存在): {e}")
|
||
|
||
try:
|
||
create_vector_index(
|
||
driver,
|
||
name=index_name, # 索引的唯一名称
|
||
label=label, # 要索引的节点标签
|
||
embedding_property="embedding", # 包含嵌入向量值的节点属性键
|
||
dimensions=vector_dim, # 向量嵌入维度,1024与使用的 bge-m3 嵌入模型一致
|
||
similarity_fn="cosine", # 向量相似度函数,可选 "euclidean" 或 "cosine"
|
||
)
|
||
logger.info(f"成功为标签 {label} 创建向量索引(共 {node_count} 个节点)")
|
||
except Exception as e:
|
||
# 如果是增量更新模式且索引已存在,这是正常的,只记录警告
|
||
if not force_refresh and "already exists" in str(e).lower():
|
||
logger.info(f"标签 {label} 的向量索引已存在,跳过创建")
|
||
else:
|
||
logger.error(f"为标签 {label} 创建向量索引失败: {e}")
|
||
raise
|
||
|
||
|
||
# --------- 创建全文索引 ---------
|
||
def fulltext_indexing(driver, label, force_refresh=False):
|
||
"""
|
||
创建全文索引,并添加节点属性(自动获取所有属性)
|
||
|
||
全文索引:词 ---》 id
|
||
后续检索,分词后,根据词反向找到节点
|
||
|
||
Args:
|
||
driver: Neo4j driver
|
||
label: 节点标签
|
||
force_refresh: 是否强制刷新所有节点的全文索引(True=清空后重新生成,False=只处理缺失的节点)
|
||
"""
|
||
# 排除内部属性
|
||
excluded_props = {"embedding", "fulltext", "elementId", "id"}
|
||
|
||
# 根据 force_refresh 决定查询条件
|
||
if force_refresh:
|
||
# 强制刷新:处理所有节点(包括已有 fulltext 的节点)
|
||
query = f"""MATCH (n:{label})
|
||
WITH n, elementId(n) AS id, properties(n) AS props
|
||
RETURN id, props"""
|
||
else:
|
||
# 增量更新:只处理没有 fulltext 的节点
|
||
query = f"""MATCH (n:{label})
|
||
WHERE n.fulltext IS NULL
|
||
WITH n, elementId(n) AS id, properties(n) AS props
|
||
RETURN id, props"""
|
||
|
||
# 执行查询
|
||
records = driver.execute_query(query).records
|
||
|
||
if not records:
|
||
if force_refresh:
|
||
logger.info(f"{label} 没有节点需要处理")
|
||
else:
|
||
logger.info(f"{label} 所有节点皆存在全文索引属性")
|
||
return
|
||
|
||
# 构建文本描述(使用所有属性)
|
||
record_list = []
|
||
for r in records:
|
||
node_id = r["id"]
|
||
props = r["props"]
|
||
# 过滤掉内部属性
|
||
filtered_props = {k: v for k, v in props.items() if k not in excluded_props}
|
||
# 使用所有属性构建文本
|
||
text = build_text_from_all_properties(filtered_props)
|
||
if text:
|
||
record_list.append({"id": node_id, "text": text})
|
||
|
||
if not record_list:
|
||
logger.info(f"{label} 没有可索引的文本内容")
|
||
return
|
||
|
||
# 创建全文索引(在写入数据之前创建索引)
|
||
try:
|
||
create_fulltext_index(
|
||
driver,
|
||
name=f"{label.lower()}_fulltext", # 索引的唯一名称
|
||
label=label, # 要创建索引的节点标签
|
||
node_properties=["fulltext"], # 要创建全文索引的节点属性列表
|
||
)
|
||
except Exception as e:
|
||
logger.warning(f"创建全文索引可能已存在(继续处理): {e}")
|
||
|
||
# 文本分词,作为全文索引属性
|
||
logger.info(f"计算 {label} ({len(record_list)}) 的全文索引")
|
||
pattern = re.compile(r"[a-zA-Z0-9\u4e00-\u9fa5]+") # 匹配英文字母、数字和中文字符
|
||
fulltext_tuple_list = []
|
||
for id_, text in [(r["id"], r["text"]) for r in record_list]:
|
||
# 使用 jieba 分词
|
||
words = jieba.lcut(text)
|
||
# 过滤并连接
|
||
filtered_words = [
|
||
word.strip()
|
||
for word in words
|
||
if word.strip() and pattern.fullmatch(word.strip())
|
||
]
|
||
fulltext_value = " ".join(filtered_words)
|
||
if fulltext_value:
|
||
fulltext_tuple_list.append((id_, fulltext_value))
|
||
|
||
if not fulltext_tuple_list:
|
||
logger.warning(f"{label} 分词后没有有效内容")
|
||
return
|
||
|
||
ids, fulltexts = zip(*fulltext_tuple_list)
|
||
ids = list(ids)
|
||
fulltexts = list(fulltexts)
|
||
|
||
# 按 elementId 添加全文索引属性
|
||
logger.info(f"写入 {label} ({len(fulltexts)}) 的全文索引")
|
||
insert_batch_size = 1000
|
||
for i in range(0, len(fulltext_tuple_list), insert_batch_size):
|
||
batch_rows = [
|
||
{"id": id_, "fulltext": ft}
|
||
for id_, ft in zip(
|
||
ids[i: i + insert_batch_size],
|
||
fulltexts[i: i + insert_batch_size],
|
||
)
|
||
]
|
||
|
||
# UNWIND:将列表数据展开为多行记录
|
||
driver.execute_query(
|
||
"UNWIND $rows AS row " # 将传入的 rows 列表展开,每项作为一行数据,命名为 row
|
||
"MATCH (n) "
|
||
"WHERE elementId(n) = row.id "
|
||
"SET n.fulltext = row.fulltext ",
|
||
{"rows": batch_rows},
|
||
)
|
||
|
||
|
||
def create_all_indexes(neo4j_uri=None, neo4j_username=None, neo4j_password=None, force_refresh=False, driver=None):
|
||
"""
|
||
创建所有索引的主函数(可被外部调用)
|
||
|
||
注意:Neo4j 连接配置优先从 .env 文件中读取(NEO4J_URI, NEO4J_USERNAME, NEO4J_PASSWORD)
|
||
如果 .env 文件中没有配置,则使用参数中的值(如果提供)或默认值
|
||
|
||
Args:
|
||
neo4j_uri: Neo4j URI(可选,优先级低于 .env,提供 driver 时忽略)
|
||
neo4j_username: Neo4j 用户名(可选,优先级低于 .env,提供 driver 时忽略)
|
||
neo4j_password: Neo4j 密码(可选,优先级低于 .env,提供 driver 时忽略)
|
||
force_refresh: 是否强制刷新所有索引(True=清空重建,False=增量更新)
|
||
driver: 可选的 Neo4j driver 实例,提供则复用外部 driver
|
||
|
||
Returns:
|
||
dict: 包含执行结果的字典
|
||
"""
|
||
use_external_driver = driver is not None
|
||
|
||
if not use_external_driver:
|
||
# 优先从 .env 文件读取配置,如果 .env 中没有则使用参数或默认值
|
||
neo4j_uri = os.getenv("NEO4J_URI") or neo4j_uri or "bolt://192.168.0.46:57687"
|
||
neo4j_username = os.getenv("NEO4J_USERNAME") or neo4j_username or "neo4j"
|
||
neo4j_password = os.getenv("NEO4J_PASSWORD") or neo4j_password or "zdht123@"
|
||
logger.info(f"连接到 Neo4j: {neo4j_uri}")
|
||
else:
|
||
logger.info("使用外部提供的 Neo4j driver")
|
||
|
||
# 定义内部函数来执行索引创建逻辑(避免代码重复)
|
||
def _execute_indexing(driver):
|
||
# 1、根据 force_refresh 决定是否清空索引和约束
|
||
if force_refresh:
|
||
logger.info("强制刷新模式:清空所有约束...")
|
||
drop_constraint(driver)
|
||
logger.info("强制刷新模式:清空所有索引...")
|
||
drop_index_without_constraint(driver)
|
||
logger.info("强制刷新模式:清空所有节点的 fulltext 属性...")
|
||
driver.execute_query("MATCH (n) WHERE n.fulltext IS NOT NULL REMOVE n.fulltext")
|
||
else:
|
||
logger.info("增量更新模式:保留现有索引和约束")
|
||
|
||
# 2、自动获取所有节点标签并创建向量索引
|
||
logger.info("=" * 50)
|
||
logger.info("开始创建向量索引...")
|
||
logger.info("=" * 50)
|
||
|
||
all_labels = get_all_labels(driver)
|
||
logger.info(f"发现 {len(all_labels)} 个节点标签: {all_labels}")
|
||
|
||
vector_index_errors = []
|
||
for label in all_labels:
|
||
logger.info(f"处理标签: {label}")
|
||
try:
|
||
vector_indexing(driver, label, force_refresh=force_refresh)
|
||
except Exception as e:
|
||
error_msg = f"为标签 {label} 创建向量索引时出错: {e}"
|
||
logger.error(error_msg)
|
||
vector_index_errors.append(error_msg)
|
||
continue
|
||
|
||
# 3、创建全文索引(自动获取所有属性)
|
||
logger.info("=" * 50)
|
||
logger.info("开始创建全文索引...")
|
||
logger.info("=" * 50)
|
||
|
||
fulltext_index_errors = []
|
||
for label in all_labels:
|
||
logger.info(f"处理标签: {label}")
|
||
try:
|
||
fulltext_indexing(driver, label, force_refresh=force_refresh)
|
||
except Exception as e:
|
||
error_msg = f"为标签 {label} 创建全文索引时出错: {e}"
|
||
logger.error(error_msg)
|
||
fulltext_index_errors.append(error_msg)
|
||
continue
|
||
|
||
logger.info("=" * 50)
|
||
logger.info("所有索引创建完成!")
|
||
logger.info("=" * 50)
|
||
|
||
has_errors = len(vector_index_errors) > 0 or len(fulltext_index_errors) > 0
|
||
mode_text = "强制刷新" if force_refresh else "增量更新"
|
||
return {
|
||
"success": not has_errors,
|
||
"message": f"{mode_text}模式:索引创建完成" if not has_errors
|
||
else f"{mode_text}模式:索引创建完成,但部分标签出现错误",
|
||
"mode": "force_refresh" if force_refresh else "incremental",
|
||
"labels": all_labels,
|
||
"vector_index_errors": vector_index_errors,
|
||
"fulltext_index_errors": fulltext_index_errors,
|
||
}
|
||
|
||
try:
|
||
if use_external_driver:
|
||
# 使用外部 driver,直接执行,不关闭 driver
|
||
return _execute_indexing(driver)
|
||
else:
|
||
# 使用内部 driver,使用上下文管理器自动关闭
|
||
with GraphDatabase.driver(neo4j_uri, auth=(neo4j_username, neo4j_password)) as driver:
|
||
return _execute_indexing(driver)
|
||
except Exception as e:
|
||
error_msg = f"创建索引过程中发生错误: {e}"
|
||
logger.error(error_msg, exc_info=True)
|
||
return {
|
||
"success": False,
|
||
"message": error_msg,
|
||
"mode": "error",
|
||
"labels": [],
|
||
"vector_index_errors": [],
|
||
"fulltext_index_errors": [],
|
||
}
|
||
|
||
|
||
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__":
|
||
# 可以通过命令行参数或环境变量控制是否强制刷新
|
||
import sys
|
||
|
||
force_refresh = "--force" in sys.argv or os.getenv("FORCE_REFRESH", "false").lower() == "true"
|
||
|
||
result = create_all_indexes(force_refresh=force_refresh)
|
||
if result["success"]:
|
||
logger.info(f"索引创建成功完成(模式: {result['mode']})")
|
||
else:
|
||
logger.error(f"索引创建完成但有错误: {result['message']}")
|