Files
2026-04-16 10:53:10 +08:00

329 lines
11 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
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",
}