Files
ai-g3sb-backman2.0/backend/schema/indexer.py
T

341 lines
12 KiB
Python
Raw Normal View History

2026-04-10 16:52:07 +08:00
"""
Schema向量索引构建器 - 基于ChromaDB
"""
import logging
import os
from pathlib import Path
from typing import Any, Dict, List, Optional
2026-04-10 16:52:07 +08:00
import chromadb
from chromadb.config import Settings as ChromaSettings
from schema.manager import SchemaManager
logger = logging.getLogger(__name__)
class SchemaIndexer:
"""
Schema向量索引器
功能:
1. 为表结构构建向量索引
2. 基于问题的表检索
<<<<<<< HEAD
3. 默认使用 Chroma ``PersistentClient``,数据落在 ``persist_dir``(与 ``VECTOR_DB_PATH`` 一致);
仅当环境变量 ``SCHEMA_INDEXER_EPHEMERAL=true`` 或构造参数 ``use_ephemeral=True`` 时使用内存客户端。
=======
2026-04-15 13:07:39 +08:00
3. Chroma 使用持久化(PersistentClient),向量与元数据写入 ``persist_dir``,进程重启后可复用
>>>>>>> 9369563942d8de442508613191670123d5ae9d08
2026-04-10 16:52:07 +08:00
"""
def __init__(
self,
embedder: Any,
persist_dir: str = "./data/embeddings/chroma",
2026-04-10 16:52:07 +08:00
collection_name: str = "schema_tables",
*,
use_ephemeral: Optional[bool] = None,
2026-04-10 16:52:07 +08:00
):
"""
初始化索引器
Args:
embedder: Embedding模型实例
<<<<<<< HEAD
persist_dir: Chroma 持久化目录(默认与编排器 ``vector_db_path`` / ``VECTOR_DB_PATH`` 一致)
=======
2026-04-15 13:07:39 +08:00
persist_dir: Chroma 持久化根目录(磁盘路径)
>>>>>>> 9369563942d8de442508613191670123d5ae9d08
2026-04-10 16:52:07 +08:00
collection_name: 集合名称
use_ephemeral: 为 True 时使用内存 Chroma;为 None 时读环境变量 SCHEMA_INDEXER_EPHEMERAL
2026-04-10 16:52:07 +08:00
"""
self.embedder = embedder
self.persist_dir = Path(persist_dir)
2026-04-15 13:07:39 +08:00
self.persist_dir.mkdir(parents=True, exist_ok=True)
self.collection_name = collection_name
2026-04-10 16:52:07 +08:00
<<<<<<< HEAD
env_ephemeral = os.getenv("SCHEMA_INDEXER_EPHEMERAL", "").lower() in (
"1",
"true",
"yes",
=======
2026-04-15 13:07:39 +08:00
# 持久化 Chroma:数据落盘至 persist_dir
self.client = chromadb.PersistentClient(
path=str(self.persist_dir),
2026-04-10 16:52:07 +08:00
settings=ChromaSettings(anonymized_telemetry=False),
>>>>>>> 9369563942d8de442508613191670123d5ae9d08
2026-04-10 16:52:07 +08:00
)
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}"
2026-04-10 16:52:07 +08:00
# 获取或创建集合
self.collection = self.client.get_or_create_collection(
name=self.collection_name,
2026-04-10 16:52:07 +08:00
metadata={"hnsw:space": "cosine"}, # 使用余弦相似度
)
logger.info(
<<<<<<< HEAD
f"[OK] 初始化SchemaIndexer({backend_desc}): collection={self.collection_name}, "
f"count={self.collection.count()}"
=======
2026-04-15 13:07:39 +08:00
f"[OK] 初始化SchemaIndexer(Chroma持久化): collection={self.collection_name}, "
f"path={self.persist_dir}"
>>>>>>> 9369563942d8de442508613191670123d5ae9d08
)
2026-04-10 16:52:07 +08:00
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)
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)}条旧记录")
# 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...")
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": {...}
},
...
]
"""
# 编码查询文本(与建库时一致: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):
"""清空索引"""
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"])
metas = result.get("metadatas") or []
2026-04-10 16:52:07 +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),
<<<<<<< HEAD
"chroma_mode": "memory" if self._chroma_ephemeral else "persistent",
=======
2026-04-15 13:07:39 +08:00
"chroma_mode": "persistent",
>>>>>>> 9369563942d8de442508613191670123d5ae9d08
2026-04-10 16:52:07 +08:00
}