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:
陈辅元
2026-04-14 18:02:12 +08:00
parent 026d3bf88c
commit ca8bc5e7de
87 changed files with 5808 additions and 6351 deletions
Binary file not shown.
+31 -8
View File
@@ -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,