""" Schema向量索引构建器 - 基于ChromaDB """ import logging import os from pathlib import Path from typing import Any, Dict, List, Optional import chromadb from chromadb.config import Settings as ChromaSettings from schema.manager import SchemaManager logger = logging.getLogger(__name__) class SchemaIndexer: """ Schema向量索引器 功能: 1. 为表结构构建向量索引 2. 基于问题的表检索 <<<<<<< HEAD 3. 默认使用 Chroma ``PersistentClient``,数据落在 ``persist_dir``(与 ``VECTOR_DB_PATH`` 一致); 仅当环境变量 ``SCHEMA_INDEXER_EPHEMERAL=true`` 或构造参数 ``use_ephemeral=True`` 时使用内存客户端。 ======= 3. Chroma 使用持久化(PersistentClient),向量与元数据写入 ``persist_dir``,进程重启后可复用 >>>>>>> 9369563942d8de442508613191670123d5ae9d08 """ def __init__( self, embedder: Any, persist_dir: str = "./data/embeddings/chroma", collection_name: str = "schema_tables", *, use_ephemeral: Optional[bool] = None, ): """ 初始化索引器 Args: embedder: Embedding模型实例 <<<<<<< HEAD persist_dir: Chroma 持久化目录(默认与编排器 ``vector_db_path`` / ``VECTOR_DB_PATH`` 一致) ======= persist_dir: Chroma 持久化根目录(磁盘路径) >>>>>>> 9369563942d8de442508613191670123d5ae9d08 collection_name: 集合名称 use_ephemeral: 为 True 时使用内存 Chroma;为 None 时读环境变量 SCHEMA_INDEXER_EPHEMERAL """ self.embedder = embedder self.persist_dir = Path(persist_dir) self.persist_dir.mkdir(parents=True, exist_ok=True) self.collection_name = collection_name <<<<<<< HEAD env_ephemeral = os.getenv("SCHEMA_INDEXER_EPHEMERAL", "").lower() in ( "1", "true", "yes", ======= # 持久化 Chroma:数据落盘至 persist_dir self.client = chromadb.PersistentClient( path=str(self.persist_dir), settings=ChromaSettings(anonymized_telemetry=False), >>>>>>> 9369563942d8de442508613191670123d5ae9d08 ) self._chroma_ephemeral = bool(use_ephemeral) if use_ephemeral is not None else env_ephemeral if self._chroma_ephemeral: self.client = chromadb.EphemeralClient( settings=ChromaSettings(anonymized_telemetry=False), ) backend_desc = "Chroma内存" else: self.persist_dir.mkdir(parents=True, exist_ok=True) self.client = chromadb.PersistentClient( path=str(self.persist_dir), settings=ChromaSettings(anonymized_telemetry=False), ) backend_desc = f"Chroma磁盘 path={self.persist_dir}" # 获取或创建集合 self.collection = self.client.get_or_create_collection( name=self.collection_name, metadata={"hnsw:space": "cosine"}, # 使用余弦相似度 ) logger.info( <<<<<<< HEAD f"[OK] 初始化SchemaIndexer({backend_desc}): collection={self.collection_name}, " f"count={self.collection.count()}" ======= f"[OK] 初始化SchemaIndexer(Chroma持久化): collection={self.collection_name}, " f"path={self.persist_dir}" >>>>>>> 9369563942d8de442508613191670123d5ae9d08 ) def ensure_index_for_schema( self, schema_manager: SchemaManager, batch_size: int = 32, ) -> None: """ 每次使用向量粗筛前调用:若持久化集合中表条数与当前 Schema 一致且非空,则跳过向量化; 若为空则全量构建;若条数不一致则强制重建(Schema 变更后重新加载)。 """ expected = len(schema_manager.get_tables()) indexed = self.count() if expected == 0: if indexed > 0: logger.warning("当前 Schema 无表但向量索引非空,已清空索引") self.clear() return if indexed == expected: logger.info( "Schema 向量索引已就绪(%s 张表),跳过向量化", indexed, ) return if indexed == 0: logger.info( "向量索引为空,开始构建共 %s 张表(向量化)...", expected, ) self.build_index(schema_manager, batch_size=batch_size, force_rebuild=False) return logger.info( "向量索引与当前 Schema 不一致(索引 %s 张,Schema %s 张),重新向量化并加载...", indexed, expected, ) self.build_index(schema_manager, batch_size=batch_size, force_rebuild=True) def build_index( self, schema_manager: SchemaManager, batch_size: int = 32, force_rebuild: bool = False, ) -> bool: """ 为Schema构建向量索引 Args: schema_manager: SchemaManager实例 batch_size: 批处理大小 force_rebuild: 是否强制重建(默认为False,增量更新) Returns: True=成功,False=已存在且未强制重建 """ existing_ids = set(self.collection.get()["ids"]) if self.collection.count() > 0 else set() if not force_rebuild and existing_ids: logger.info(f"索引已存在({len(existing_ids)}条记录),跳过构建") return False if force_rebuild and existing_ids: logger.info(f"强制重建索引,删除{len(existing_ids)}条旧记录") # Chroma 新版本要求 delete 必须带 ids/where;整库清空用删集合再建 self.client.delete_collection(self.collection_name) self.collection = self.client.get_or_create_collection( name=self.collection_name, metadata={"hnsw:space": "cosine"}, ) # 准备数据 table_texts = [] table_ids = [] table_metadatas = [] for table_info in schema_manager.get_table_for_embedding(): table_texts.append(table_info["text"]) table_ids.append(table_info["name"]) table_metadatas.append({ "table_name": table_info["name"], "column_count": len(schema_manager.get_table(table_info["name"]).columns), }) # 批量计算embedding logger.info(f"计算{len(table_texts)}张表的embedding...") embeddings = self.embedder.encode( table_texts, batch_size=batch_size, normalize=True ) # 存入ChromaDB self.collection.add( embeddings=embeddings.tolist(), documents=table_texts, metadatas=table_metadatas, ids=table_ids, ) logger.info(f"[OK] 索引构建完成,共{len(table_ids)}张表") return True def search( self, query: str, top_k: int = 20, score_threshold: float = 0.3, ) -> List[Dict]: """ 检索相关表 Args: query: 查询文本(用户问题) top_k: 返回前K个结果 score_threshold: 相似度阈值(0-1),低于此值的结果会被过滤 Returns: 检索结果列表,按相似度降序排列 [ { "table_name": "表名", "score": 0.95, "document": "表描述文本", "metadata": {...} }, ... ] """ # 编码查询文本(与建库时一致:L2 归一化 + 余弦空间) query_embedding = self.embedder.encode([query], normalize=True) # 执行检索 results = self.collection.query( query_embeddings=query_embedding.tolist(), n_results=min(top_k, self.collection.count()), ) # 格式化结果 formatted = [] if results["ids"] and len(results["ids"][0]) > 0: for idx, (table_id, distance, metadata, document) in enumerate(zip( results["ids"][0], results["distances"][0], results["metadatas"][0], results["documents"][0], )): score = 1 - distance # 余弦相似度(已归一化,转换为0-1,越大越相似) if score >= score_threshold: formatted.append({ "table_name": table_id, "score": float(score), "document": document, "metadata": metadata, "rank": idx + 1, }) logger.debug(f"检索 '{query[:50]}...' -> 找到{len(formatted)}个相关表(阈值={score_threshold})") return formatted def search_by_table_names(self, table_names: List[str]) -> List[Dict]: """ 直接通过表名获取表信息 Args: table_names: 表名列表 Returns: 表信息列表 """ results = [] for name in table_names: try: result = self.collection.get(ids=[name], include=["metadatas", "documents"]) if result["ids"]: results.append({ "table_name": name, "score": 1.0, "document": result["documents"][0], "metadata": result["metadatas"][0], }) except Exception as e: logger.warning(f"获取表信息失败 {name}: {e}") return results def get_all_tables(self) -> List[str]: """获取索引中的所有表名""" result = self.collection.get() return result["ids"] if result["ids"] else [] def delete_table(self, table_name: str): """从索引中删除表""" self.collection.delete(ids=[table_name]) logger.info(f"已删除表索引: {table_name}") def clear(self): """清空索引""" if self.collection.count() > 0: self.client.delete_collection(self.collection_name) self.collection = self.client.get_or_create_collection( name=self.collection_name, metadata={"hnsw:space": "cosine"}, ) logger.info("索引已清空") def count(self) -> int: """获取索引中的表数量""" return self.collection.count() # ========== 统计与调试 ========== def get_statistics(self) -> Dict: """获取索引统计信息""" count = self.collection.count() result = self.collection.get(include=["metadatas"]) metas = result.get("metadatas") or [] def _col_count(m: Optional[Dict]) -> int: if not m: return 0 v = m.get("column_count", 0) try: return int(v) except (TypeError, ValueError): return 0 total_columns = sum(_col_count(m) for m in metas) return { "indexed_tables": count, "total_columns": total_columns, "avg_columns": total_columns / count if count > 0 else 0, "persist_dir": str(self.persist_dir), <<<<<<< HEAD "chroma_mode": "memory" if self._chroma_ephemeral else "persistent", ======= "chroma_mode": "persistent", >>>>>>> 9369563942d8de442508613191670123d5ae9d08 }