Enhance dialog context handling for Text2SQL queries by integrating session history. Introduce new methods for summarizing previous assistant messages and determining if the last interaction was a data query. Update environment configuration for embedding options and improve error handling in SQL generation. Add user-facing delivery messages for successful SQL execution. This update supports more coherent follow-up questions and improves user experience in conversational interactions.
This commit is contained in:
Binary file not shown.
@@ -40,6 +40,7 @@ class SchemaIndexer:
|
||||
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(
|
||||
@@ -49,11 +50,11 @@ class SchemaIndexer:
|
||||
|
||||
# 获取或创建集合
|
||||
self.collection = self.client.get_or_create_collection(
|
||||
name=collection_name,
|
||||
name=self.collection_name,
|
||||
metadata={"hnsw:space": "cosine"}, # 使用余弦相似度
|
||||
)
|
||||
|
||||
logger.info(f"[OK] 初始化SchemaIndexer: collection={collection_name}")
|
||||
logger.info(f"[OK] 初始化SchemaIndexer: collection={self.collection_name}")
|
||||
|
||||
def build_index(
|
||||
self,
|
||||
@@ -80,7 +81,12 @@ class SchemaIndexer:
|
||||
|
||||
if force_rebuild and existing_ids:
|
||||
logger.info(f"强制重建索引,删除{len(existing_ids)}条旧记录")
|
||||
self.collection.delete()
|
||||
# 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 = []
|
||||
@@ -97,7 +103,9 @@ class SchemaIndexer:
|
||||
|
||||
# 批量计算embedding
|
||||
logger.info(f"计算{len(table_texts)}张表的embedding...")
|
||||
embeddings = self.embedder.encode(table_texts, batch_size=batch_size)
|
||||
embeddings = self.embedder.encode(
|
||||
table_texts, batch_size=batch_size, normalize=True
|
||||
)
|
||||
|
||||
# 存入ChromaDB
|
||||
self.collection.add(
|
||||
@@ -136,8 +144,8 @@ class SchemaIndexer:
|
||||
...
|
||||
]
|
||||
"""
|
||||
# 编码查询文本
|
||||
query_embedding = self.embedder.encode([query])
|
||||
# 编码查询文本(与建库时一致:L2 归一化 + 余弦空间)
|
||||
query_embedding = self.embedder.encode([query], normalize=True)
|
||||
|
||||
# 执行检索
|
||||
results = self.collection.query(
|
||||
@@ -205,7 +213,12 @@ class SchemaIndexer:
|
||||
|
||||
def clear(self):
|
||||
"""清空索引"""
|
||||
self.collection.delete()
|
||||
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:
|
||||
@@ -218,8 +231,18 @@ class SchemaIndexer:
|
||||
"""获取索引统计信息"""
|
||||
count = self.collection.count()
|
||||
result = self.collection.get(include=["metadatas"])
|
||||
metas = result.get("metadatas") or []
|
||||
|
||||
total_columns = sum(m.get("column_count", 0) for m in result["metadatas"])
|
||||
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,
|
||||
|
||||
Reference in New Issue
Block a user