""" Schema向量索引构建器 - 基于ChromaDB """ import chromadb from chromadb.config import Settings as ChromaSettings from typing import Any, List, Dict, Optional import logging from pathlib import Path from schema.manager import SchemaManager logger = logging.getLogger(__name__) class SchemaIndexer: """ Schema向量索引器 功能: 1. 为表结构构建向量索引 2. 基于问题的表检索 3. 持久化存储和加载 """ def __init__( self, embedder: Any, persist_dir: str = "./data/embeddings", collection_name: str = "schema_tables", ): """ 初始化索引器 Args: embedder: Embedding模型实例 persist_dir: 向量数据库持久化目录 collection_name: 集合名称 """ self.embedder = embedder self.persist_dir = Path(persist_dir) self.persist_dir.mkdir(parents=True, exist_ok=True) # 初始化ChromaDB客户端 self.client = chromadb.PersistentClient( path=str(self.persist_dir), settings=ChromaSettings(anonymized_telemetry=False), ) # 获取或创建集合 self.collection = self.client.get_or_create_collection( name=collection_name, metadata={"hnsw:space": "cosine"}, # 使用余弦相似度 ) logger.info(f"[OK] 初始化SchemaIndexer: collection={collection_name}") 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)}条旧记录") self.collection.delete() # 准备数据 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) # 存入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": {...} }, ... ] """ # 编码查询文本 query_embedding = self.embedder.encode([query]) # 执行检索 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): """清空索引""" self.collection.delete() 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"]) total_columns = sum(m.get("column_count", 0) for m in result["metadatas"]) return { "indexed_tables": count, "total_columns": total_columns, "avg_columns": total_columns / count if count > 0 else 0, "persist_dir": str(self.persist_dir), }