329 lines
11 KiB
Python
329 lines
11 KiB
Python
"""
|
||
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. 基于问题的表检索
|
||
3. 默认使用 Chroma ``PersistentClient``,数据落在 ``persist_dir``(与 ``VECTOR_DB_PATH`` 一致);
|
||
仅当环境变量 ``SCHEMA_INDEXER_EPHEMERAL=true`` 或构造参数 ``use_ephemeral=True`` 时使用内存客户端。
|
||
"""
|
||
|
||
def __init__(
|
||
self,
|
||
embedder: Any,
|
||
persist_dir: str = "./data/embeddings/chroma",
|
||
collection_name: str = "schema_tables",
|
||
*,
|
||
use_ephemeral: Optional[bool] = None,
|
||
):
|
||
"""
|
||
初始化索引器
|
||
|
||
Args:
|
||
embedder: Embedding模型实例
|
||
persist_dir: Chroma 持久化根目录(磁盘路径)
|
||
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
|
||
|
||
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),
|
||
)
|
||
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(
|
||
f"[OK] 初始化SchemaIndexer({backend_desc}): collection={self.collection_name}, "
|
||
f"count={self.collection.count()}"
|
||
)
|
||
|
||
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,
|
||
})
|
||
|
||
top = [(x["table_name"], round(float(x.get("score", 0.0)), 4)) for x in formatted[:15]]
|
||
logger.info(
|
||
"Schema 向量检索: query_chars=%s 命中=%s(阈值=%s)top=%s",
|
||
len(query or ""),
|
||
len(formatted),
|
||
score_threshold,
|
||
top,
|
||
)
|
||
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),
|
||
"chroma_mode": "memory" if self._chroma_ephemeral else "persistent",
|
||
}
|