2026-04-10 16:52:07 +08:00
|
|
|
|
"""
|
|
|
|
|
|
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. 基于问题的表检索
|
2026-04-14 18:21:50 +08:00
|
|
|
|
3. Chroma 使用内存模式(EphemeralClient),进程退出后不保留;``persist_dir`` 仅作配置/统计引用
|
2026-04-10 16:52:07 +08:00
|
|
|
|
"""
|
|
|
|
|
|
|
|
|
|
|
|
def __init__(
|
|
|
|
|
|
self,
|
|
|
|
|
|
embedder: Any,
|
|
|
|
|
|
persist_dir: str = "./data/embeddings",
|
|
|
|
|
|
collection_name: str = "schema_tables",
|
|
|
|
|
|
):
|
|
|
|
|
|
"""
|
|
|
|
|
|
初始化索引器
|
|
|
|
|
|
|
|
|
|
|
|
Args:
|
|
|
|
|
|
embedder: Embedding模型实例
|
2026-04-14 18:21:50 +08:00
|
|
|
|
persist_dir: 历史配置中的向量库路径(仅展示与统计,不落盘)
|
2026-04-10 16:52:07 +08:00
|
|
|
|
collection_name: 集合名称
|
|
|
|
|
|
"""
|
|
|
|
|
|
self.embedder = embedder
|
|
|
|
|
|
self.persist_dir = Path(persist_dir)
|
2026-04-14 18:02:12 +08:00
|
|
|
|
self.collection_name = collection_name
|
2026-04-10 16:52:07 +08:00
|
|
|
|
|
2026-04-14 18:21:50 +08:00
|
|
|
|
# 内存 Chroma:不落盘,每次进程需重新 build_index
|
|
|
|
|
|
self.client = chromadb.EphemeralClient(
|
2026-04-10 16:52:07 +08:00
|
|
|
|
settings=ChromaSettings(anonymized_telemetry=False),
|
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
# 获取或创建集合
|
|
|
|
|
|
self.collection = self.client.get_or_create_collection(
|
2026-04-14 18:02:12 +08:00
|
|
|
|
name=self.collection_name,
|
2026-04-10 16:52:07 +08:00
|
|
|
|
metadata={"hnsw:space": "cosine"}, # 使用余弦相似度
|
|
|
|
|
|
)
|
|
|
|
|
|
|
2026-04-14 18:21:50 +08:00
|
|
|
|
logger.info(
|
|
|
|
|
|
f"[OK] 初始化SchemaIndexer(Chroma内存): collection={self.collection_name}, "
|
|
|
|
|
|
f"配置路径引用={self.persist_dir}"
|
|
|
|
|
|
)
|
2026-04-10 16:52:07 +08:00
|
|
|
|
|
|
|
|
|
|
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)}条旧记录")
|
2026-04-14 18:02:12 +08:00
|
|
|
|
# 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"},
|
|
|
|
|
|
)
|
2026-04-10 16:52:07 +08:00
|
|
|
|
|
|
|
|
|
|
# 准备数据
|
|
|
|
|
|
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...")
|
2026-04-14 18:02:12 +08:00
|
|
|
|
embeddings = self.embedder.encode(
|
|
|
|
|
|
table_texts, batch_size=batch_size, normalize=True
|
|
|
|
|
|
)
|
2026-04-10 16:52:07 +08:00
|
|
|
|
|
|
|
|
|
|
# 存入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": {...}
|
|
|
|
|
|
},
|
|
|
|
|
|
...
|
|
|
|
|
|
]
|
|
|
|
|
|
"""
|
2026-04-14 18:02:12 +08:00
|
|
|
|
# 编码查询文本(与建库时一致:L2 归一化 + 余弦空间)
|
|
|
|
|
|
query_embedding = self.embedder.encode([query], normalize=True)
|
2026-04-10 16:52:07 +08:00
|
|
|
|
|
|
|
|
|
|
# 执行检索
|
|
|
|
|
|
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):
|
|
|
|
|
|
"""清空索引"""
|
2026-04-14 18:02:12 +08:00
|
|
|
|
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"},
|
|
|
|
|
|
)
|
2026-04-10 16:52:07 +08:00
|
|
|
|
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"])
|
2026-04-14 18:02:12 +08:00
|
|
|
|
metas = result.get("metadatas") or []
|
2026-04-10 16:52:07 +08:00
|
|
|
|
|
2026-04-14 18:02:12 +08:00
|
|
|
|
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)
|
2026-04-10 16:52:07 +08:00
|
|
|
|
|
|
|
|
|
|
return {
|
|
|
|
|
|
"indexed_tables": count,
|
|
|
|
|
|
"total_columns": total_columns,
|
|
|
|
|
|
"avg_columns": total_columns / count if count > 0 else 0,
|
|
|
|
|
|
"persist_dir": str(self.persist_dir),
|
2026-04-14 18:21:50 +08:00
|
|
|
|
"chroma_mode": "memory",
|
2026-04-10 16:52:07 +08:00
|
|
|
|
}
|