# 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"])