253 lines
8.0 KiB
Python
253 lines
8.0 KiB
Python
"""
|
||
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)
|
||
self.collection_name = collection_name
|
||
|
||
# 初始化ChromaDB客户端
|
||
self.client = chromadb.PersistentClient(
|
||
path=str(self.persist_dir),
|
||
settings=ChromaSettings(anonymized_telemetry=False),
|
||
)
|
||
|
||
# 获取或创建集合
|
||
self.collection = self.client.get_or_create_collection(
|
||
name=self.collection_name,
|
||
metadata={"hnsw:space": "cosine"}, # 使用余弦相似度
|
||
)
|
||
|
||
logger.info(f"[OK] 初始化SchemaIndexer: collection={self.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)}条旧记录")
|
||
# 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),
|
||
}
|