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:
+1
-1
@@ -1,5 +1,5 @@
|
||||
# Text-to-SQL Multi-Agent System
|
||||
# 基于 CAMEL AI + DeepSeek + Qwen3-Embedding 的智能SQL生成系统
|
||||
# 基于 CAMEL AI + DeepSeek + 向量 Embedding 的智能 SQL 生成系统
|
||||
|
||||
__version__ = "0.1.0"
|
||||
__author__ = "Backman Team"
|
||||
|
||||
Binary file not shown.
Binary file not shown.
+102
-19
@@ -64,7 +64,7 @@ class Text2SQLOrchestrator:
|
||||
schema_manager: Schema管理器实例
|
||||
deepseek_api_key: DeepSeek API密钥(也可通过环境变量DEEPSEEK_API_KEY)
|
||||
deepseek_config: DeepSeek配置对象(优先于api_key)
|
||||
embedding_model_path: Qwen3-Embedding模型路径
|
||||
embedding_model_path: 本地 Embedding 模型目录(仅 USE_LOCAL_EMBEDDING=true)
|
||||
vector_db_path: 向量数据库路径
|
||||
max_retry: 最大重试次数(包含首次生成)
|
||||
use_vector_search: 是否使用向量检索粗筛
|
||||
@@ -96,12 +96,22 @@ class Text2SQLOrchestrator:
|
||||
|
||||
if self.fewshot_enabled:
|
||||
try:
|
||||
path = fewshot_samples_path or os.getenv(
|
||||
"FEWSHOT_DATA_PATH",
|
||||
"./data/experiences/all_samples.jsonl"
|
||||
use_chroma = os.getenv("FEWSHOT_USE_CHROMA", "false").lower() in (
|
||||
"1",
|
||||
"true",
|
||||
"yes",
|
||||
)
|
||||
if fewshot_samples_path is not None:
|
||||
path = str(fewshot_samples_path).strip()
|
||||
else:
|
||||
env_p = os.getenv("FEWSHOT_DATA_PATH")
|
||||
if env_p is not None:
|
||||
path = env_p.strip()
|
||||
else:
|
||||
# Chroma 优先时默认不再依赖 JSONL;否则保留原默认路径
|
||||
path = "" if use_chroma else "./data/experiences/all_samples.jsonl"
|
||||
self.fewshot_selector = FewShotSelector(
|
||||
path,
|
||||
path or None,
|
||||
embedding_model_path=self._embedding_model_path,
|
||||
)
|
||||
logger.info(
|
||||
@@ -119,6 +129,23 @@ class Text2SQLOrchestrator:
|
||||
+ (", en→zh=on" if self.translate_english_to_zh else ", en→zh=off")
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _merge_dialog_for_model(
|
||||
dialog_context: str, question: str, max_len: int
|
||||
) -> str:
|
||||
"""拼接上文与当前问句,控制总长,优先保留当前问句完整。"""
|
||||
dc = (dialog_context or "").strip()
|
||||
q = (question or "").strip()
|
||||
if not dc:
|
||||
return q
|
||||
tail = "\n\n【当前用户问题】\n" + q
|
||||
if len(dc) + len(tail) <= max_len:
|
||||
return dc + tail
|
||||
room = max_len - len(tail)
|
||||
if room < 80:
|
||||
return tail[-max_len:]
|
||||
return dc[:room].rstrip() + tail
|
||||
|
||||
def _get_vector_index(self) -> SchemaIndexer:
|
||||
"""获取或创建向量索引(懒加载)"""
|
||||
if self._vector_index is None:
|
||||
@@ -291,6 +318,7 @@ class Text2SQLOrchestrator:
|
||||
schema_str: str,
|
||||
dialect: str = "tsql",
|
||||
validation_feedback: Optional[str] = None,
|
||||
dialog_context: Optional[str] = None,
|
||||
) -> str:
|
||||
"""
|
||||
SQL生成(SQL Generator Agent)
|
||||
@@ -301,17 +329,24 @@ class Text2SQLOrchestrator:
|
||||
dialect: SQL方言
|
||||
validation_feedback: 非空时附加到用户提示(重试时传入上次校验错误)
|
||||
|
||||
dialog_context: 前几轮对话摘要;与 ``question`` 一并供指代消解与续问。
|
||||
|
||||
Returns:
|
||||
SQL语句
|
||||
"""
|
||||
from config.prompts import SQL_GENERATOR_SYSTEM, SQL_GENERATOR_USER
|
||||
from utils.sql_parser import normalize_sql_for_dialect
|
||||
|
||||
dc = (dialog_context or "").strip()
|
||||
fewshot_question = question
|
||||
if dc:
|
||||
fewshot_question = f"{dc}\n\n【当前问】{question}"
|
||||
|
||||
# Few-shot 增强
|
||||
if self.fewshot_enabled and self.fewshot_selector:
|
||||
try:
|
||||
examples = self.fewshot_selector.select(
|
||||
question=question,
|
||||
question=fewshot_question,
|
||||
top_k=self.fewshot_top_k,
|
||||
min_rating=self.fewshot_min_rating
|
||||
)
|
||||
@@ -329,7 +364,15 @@ class Text2SQLOrchestrator:
|
||||
if dialect == "tsql":
|
||||
dialect_label = "Microsoft SQL Server (T-SQL)"
|
||||
|
||||
user_content = SQL_GENERATOR_USER.format(
|
||||
prefix = ""
|
||||
if dc:
|
||||
prefix = (
|
||||
"【对话上文】(用于理解「这/那/同样/上面/刚才」等指代及续问条件;"
|
||||
"请结合下文「当前用户问题」生成 SQL。)\n"
|
||||
f"{dc}\n\n"
|
||||
)
|
||||
|
||||
user_content = prefix + SQL_GENERATOR_USER.format(
|
||||
schema=schema_str,
|
||||
question=question,
|
||||
dialect=dialect_label,
|
||||
@@ -380,6 +423,7 @@ class Text2SQLOrchestrator:
|
||||
schema_str: str,
|
||||
dialect: str = "tsql",
|
||||
question: str = "",
|
||||
dialog_context: Optional[str] = None,
|
||||
) -> Tuple[bool, List[str], List[str], Optional[int], Optional[str]]:
|
||||
"""
|
||||
SQL验证(程序验证 + 库探针 + 按探针分支的 LLM)
|
||||
@@ -389,6 +433,7 @@ class Text2SQLOrchestrator:
|
||||
schema_str: Schema描述
|
||||
dialect: 与生成一致的 SQL 方言(sqlglot 名,默认 tsql)
|
||||
question: 用户自然语言(探针为 0 时用于生成补充说明)
|
||||
dialog_context: 会话上文;探针 0 时与 question 一并传入说明模型
|
||||
|
||||
Returns:
|
||||
(是否通过, 错误列表, 警告列表, 库执行探针状态, 无数据时的用户说明)
|
||||
@@ -433,11 +478,11 @@ class Text2SQLOrchestrator:
|
||||
db_execution_status, db_probe_err = probe_sql_execution_status_ex(sql)
|
||||
if db_execution_status == -1:
|
||||
msg = (
|
||||
"数据库执行验证失败:SQL 在目标库执行报错(探针状态 -1),"
|
||||
"将据此重新生成 SQL。"
|
||||
"【库探针结果:-1 执行失败】SQL 在目标库执行报错,"
|
||||
"系统将依据下列错误**自动重新生成** SQL(请等待重试结果)。"
|
||||
)
|
||||
if db_probe_err:
|
||||
msg += f" 数据库返回:{db_probe_err}"
|
||||
msg += f"\n数据库返回:{db_probe_err}"
|
||||
errors.append(msg)
|
||||
elif db_execution_status is None:
|
||||
warnings.append(
|
||||
@@ -482,23 +527,35 @@ class Text2SQLOrchestrator:
|
||||
# === 阶段2b:探针 0 时生成用户可读补充说明(仍返回 SQL,由 API/CLI 一并展示) ===
|
||||
if db_execution_status == 0 and len(errors) == 0:
|
||||
prefix = (
|
||||
"该 SQL 已在数据库成功执行,但返回的数据行数为 0(未查到匹配记录)。"
|
||||
"请将下方 SQL 与说明一并核对;若不符合预期,请补充或调整条件后再次提问。"
|
||||
"【库探针结果:0 行】该 SQL 已在数据库成功执行,但**返回数据行数为 0**(未查到匹配记录)。"
|
||||
"下方已附带完整 SQL 与原因分析,请一并阅读。"
|
||||
)
|
||||
fb_q = question
|
||||
dc = (dialog_context or "").strip()
|
||||
if dc:
|
||||
fb_q = f"{dc}\n\n【当前用户问题】\n{question}"
|
||||
try:
|
||||
llm_fb = self.deepseek.empty_result_user_feedback(
|
||||
question=question,
|
||||
question=fb_q,
|
||||
sql=sql,
|
||||
schema=schema_str,
|
||||
)
|
||||
empty_feedback = f"{prefix}\n\n【分析与建议】\n{llm_fb}"
|
||||
empty_feedback = f"{prefix}\n\n【问题分析】\n{llm_fb}"
|
||||
except Exception as e:
|
||||
logger.warning(f"无数据说明生成失败: {e}")
|
||||
empty_feedback = (
|
||||
f"{prefix}\n\n【分析与建议】\n"
|
||||
f"{prefix}\n\n【问题分析】\n"
|
||||
"未能自动生成详细分析。请补充时间范围、筛选条件或业务对象后重新提问。"
|
||||
)
|
||||
|
||||
follow = (
|
||||
"\n\n【追问 — 请补充后再次提问以重新生成 SQL】\n"
|
||||
"1. 请根据上述分析,尽量具体地补充或修正:**时间范围**、**业务对象**(账户/合约/代码等)、"
|
||||
"**筛选口径** 或 **您认为 SQL 中不合理的条件**。\n"
|
||||
"2. 补充说明后请**重新发起一次自然语言提问**(无需粘贴 SQL),系统会结合您的新描述**重新生成**查询。"
|
||||
)
|
||||
empty_feedback = (empty_feedback or prefix) + follow
|
||||
|
||||
is_valid = len(errors) == 0
|
||||
return is_valid, errors, warnings, db_execution_status, empty_feedback
|
||||
|
||||
@@ -507,7 +564,8 @@ class Text2SQLOrchestrator:
|
||||
question: str,
|
||||
dialect: str = "tsql",
|
||||
top_k_candidates: int = 20,
|
||||
include_schema_in_result: bool = False
|
||||
include_schema_in_result: bool = False,
|
||||
dialog_context: Optional[str] = None,
|
||||
) -> GenerationResult:
|
||||
"""
|
||||
主生成流程
|
||||
@@ -517,6 +575,7 @@ class Text2SQLOrchestrator:
|
||||
dialect: SQL方言(默认 tsql;与 sqlglot 一致)
|
||||
top_k_candidates: 粗筛候选表数量
|
||||
include_schema_in_result: 结果中是否包含使用的Schema字符串
|
||||
dialog_context: 前几轮对话可读摘要;选表、向量粗筛、SQL 生成与无数据说明会参考
|
||||
|
||||
Returns:
|
||||
GenerationResult对象
|
||||
@@ -543,6 +602,9 @@ class Text2SQLOrchestrator:
|
||||
logger.warning("[GEN] 英译中失败,使用原文: %s", e)
|
||||
|
||||
question = work_question
|
||||
dc_raw = (dialog_context or "").strip()
|
||||
retrieval_question = self._merge_dialog_for_model(dc_raw, question, 4000)
|
||||
linker_question = self._merge_dialog_for_model(dc_raw, question, 6000)
|
||||
logger.info(f"[GEN] 开始生成SQL:{question[:50]}...")
|
||||
|
||||
attempt = 0
|
||||
@@ -558,14 +620,18 @@ class Text2SQLOrchestrator:
|
||||
# === Step 1: Schema筛选(仅首次) ===
|
||||
if attempt == 0:
|
||||
# 1.1 粗筛
|
||||
candidate_tables = self._coarse_filter(question, top_k=top_k_candidates)
|
||||
candidate_tables = self._coarse_filter(
|
||||
retrieval_question, top_k=top_k_candidates
|
||||
)
|
||||
|
||||
# 1.2 LLM精筛
|
||||
relevant_tables, reasoning = self._llm_select_tables(
|
||||
question,
|
||||
linker_question,
|
||||
candidate_tables,
|
||||
)
|
||||
relevant_tables = self._prioritize_broker_tables(question, relevant_tables)
|
||||
relevant_tables = self._prioritize_broker_tables(
|
||||
linker_question, relevant_tables
|
||||
)
|
||||
|
||||
# 1.3 外键扩展
|
||||
expanded_tables = self._expand_relations(relevant_tables)
|
||||
@@ -613,6 +679,7 @@ class Text2SQLOrchestrator:
|
||||
filtered_schema_str,
|
||||
dialect,
|
||||
validation_feedback=feedback,
|
||||
dialog_context=dc_raw or None,
|
||||
)
|
||||
last_sql = sql
|
||||
except Exception as e:
|
||||
@@ -626,6 +693,7 @@ class Text2SQLOrchestrator:
|
||||
filtered_schema_str,
|
||||
dialect=dialect,
|
||||
question=question,
|
||||
dialog_context=dc_raw or None,
|
||||
)
|
||||
if db_probe is not None:
|
||||
last_db_execution_status = db_probe
|
||||
@@ -643,6 +711,19 @@ class Text2SQLOrchestrator:
|
||||
meta["db_execution_status"] = db_probe
|
||||
if empty_feedback:
|
||||
meta["db_empty_feedback"] = empty_feedback
|
||||
if db_probe == 1:
|
||||
try:
|
||||
meta["sql_delivery_message"] = (
|
||||
self.deepseek.sql_probe_success_delivery_message(
|
||||
question=question,
|
||||
sql=sql,
|
||||
)
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning("[GEN] 探针1交付说明生成失败: %s", e)
|
||||
meta["sql_delivery_message"] = None
|
||||
if dc_raw:
|
||||
meta["dialog_context_chars"] = len(dc_raw)
|
||||
|
||||
result = GenerationResult(
|
||||
sql=sql,
|
||||
@@ -663,6 +744,8 @@ class Text2SQLOrchestrator:
|
||||
fail_meta: Dict = dict(translation_meta)
|
||||
if last_db_execution_status is not None:
|
||||
fail_meta["db_execution_status"] = last_db_execution_status
|
||||
if dc_raw:
|
||||
fail_meta["dialog_context_chars"] = len(dc_raw)
|
||||
return GenerationResult(
|
||||
sql=last_sql or "",
|
||||
valid=False,
|
||||
|
||||
Binary file not shown.
Binary file not shown.
@@ -283,6 +283,24 @@ EMPTY_RESULT_FEEDBACK_USER = """用户原始问题:
|
||||
|
||||
请直接输出给终端用户阅读的说明文字(纯文本)。"""
|
||||
|
||||
# 库探针为 1(有返回行)时,面向用户的 SQL 交付说明(与下方 SQL 一并展示)
|
||||
SQL_PROBE_SUCCESS_DELIVERY_SYSTEM = """你是证券/期货类业务库的 Text2SQL 助手。
|
||||
用户的自然语言问题已转为 SQL,且在目标库**试执行成功且至少返回一行数据**。
|
||||
|
||||
请用 1~3 句简洁中文向用户说明:
|
||||
- 该 SQL 大致在查询或统计什么(业务语义);
|
||||
- 可提示用户可在下方查看完整 SQL 并自行执行或导出。
|
||||
|
||||
不要编造具体数据值;不要逐列复述;不要输出 Markdown 代码块或 JSON。"""
|
||||
|
||||
SQL_PROBE_SUCCESS_DELIVERY_USER = """用户问题:
|
||||
{question}
|
||||
|
||||
已通过库探针(有数据行)的 SQL:
|
||||
{sql}
|
||||
|
||||
请输出面向终端用户的简短说明(纯文本)。"""
|
||||
|
||||
|
||||
# ========== Few-Shot 示例 ==========
|
||||
FEW_SHOT_EXAMPLES: Dict[str, str] = {
|
||||
@@ -350,3 +368,28 @@ TRANSLATE_NL_TO_ZH_USER = """原句:
|
||||
{question}
|
||||
|
||||
仅输出一句中文:"""
|
||||
|
||||
# ========== 对话意图分类(LLM,与规则分类配合)==========
|
||||
DIALOG_INTENT_CLASSIFIER_SYSTEM = """你是证券/期货类 Text2SQL 产品的「意图分类」模块。
|
||||
|
||||
只输出**一行**严格 JSON(不要 markdown、不要解释),格式:
|
||||
{"intent":"text2sql"|"conversation","reply_zh":"..."}
|
||||
|
||||
字段含义:
|
||||
- intent=text2sql:用户本轮在**提出新的或可执行的数据查询/统计需求**(含对上一轮的**具体补充条件**,如「再加上经纪商维度」「改成按日」),需要走 SQL 生成。
|
||||
- intent=conversation:用户在做**寒暄致谢**、**元问题**(你是谁)、**对结果的情绪反馈但缺少可执行信息**(如「不对呀」「结果错了」「重新算」却未说清要改什么)、**纯抱怨或否定而无新条件**——此时不应生成 SQL;reply_zh 用简短中文引导用户**具体说明**要查什么或错在哪里。
|
||||
- reply_zh:当 intent=conversation 时必填,为直接展示给用户的友好中文(1~4 句);intent=text2sql 时填空字符串 ""。
|
||||
|
||||
注意:若「上一轮助手刚返回过数据/SQL」而用户只说结果不对、未给出新的筛选/维度/时间,判为 conversation。"""
|
||||
|
||||
DIALOG_INTENT_CLASSIFIER_USER = """上一轮助手是否为「数据查询/SQL 结果」:{last_turn_label}
|
||||
|
||||
可选会话摘要(可能为空):
|
||||
---
|
||||
{context_snip}
|
||||
---
|
||||
|
||||
用户本轮输入:
|
||||
{user_message}
|
||||
|
||||
只输出 JSON:"""
|
||||
|
||||
@@ -17,7 +17,7 @@ class Settings(BaseSettings):
|
||||
max_tokens: int = 4096
|
||||
|
||||
# Embedding配置
|
||||
embedding_model_path: str = "./data/models/Qwen3-Embedding-0.6B"
|
||||
embedding_model_path: str = "" # 仅本地 Embedding 时使用;默认走远程 API
|
||||
vector_dim: int = 2048
|
||||
use_local_embedding: bool = True
|
||||
|
||||
|
||||
Binary file not shown.
@@ -253,6 +253,43 @@ class DeepSeekClient:
|
||||
text = (msg.content or "").strip()
|
||||
return text
|
||||
|
||||
def sql_probe_success_delivery_message(
|
||||
self,
|
||||
question: str,
|
||||
sql: str,
|
||||
*,
|
||||
max_sql_chars: int = 4000,
|
||||
**kwargs: Any,
|
||||
) -> str:
|
||||
"""
|
||||
库探针为 1(执行成功且至少一行数据)时,生成面向用户的简短交付说明。
|
||||
"""
|
||||
from config.prompts import (
|
||||
SQL_PROBE_SUCCESS_DELIVERY_SYSTEM,
|
||||
SQL_PROBE_SUCCESS_DELIVERY_USER,
|
||||
)
|
||||
|
||||
q = (question or "").strip() or "(无)"
|
||||
s = (sql or "").strip()
|
||||
if len(s) > max_sql_chars:
|
||||
s = s[: max_sql_chars - 20].rstrip() + "\n-- …(已截断)"
|
||||
|
||||
messages = [
|
||||
{"role": "system", "content": SQL_PROBE_SUCCESS_DELIVERY_SYSTEM},
|
||||
{
|
||||
"role": "user",
|
||||
"content": SQL_PROBE_SUCCESS_DELIVERY_USER.format(
|
||||
question=q,
|
||||
sql=s,
|
||||
),
|
||||
},
|
||||
]
|
||||
kwargs.setdefault("temperature", 0.2)
|
||||
kwargs.setdefault("top_p", 1.0)
|
||||
kwargs.setdefault("max_tokens", 320)
|
||||
msg = self.chat(messages, **kwargs)
|
||||
return (msg.content or "").strip()
|
||||
|
||||
def select_tables(
|
||||
self,
|
||||
question: str,
|
||||
|
||||
+16
-15
@@ -27,9 +27,6 @@ logging.basicConfig(
|
||||
)
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_DEFAULT_EMBEDDING_PATH = "./data/models/Qwen3-Embedding-0.6B"
|
||||
|
||||
|
||||
def _repo_root() -> Path:
|
||||
"""仓库根目录(含 data/、.env、api_server.py 的目录)。"""
|
||||
return Path(__file__).resolve().parent.parent
|
||||
@@ -41,7 +38,8 @@ def _load_project_env():
|
||||
|
||||
|
||||
def _embedding_model_path() -> str:
|
||||
return os.getenv("EMBEDDING_MODEL_PATH", _DEFAULT_EMBEDDING_PATH).strip()
|
||||
"""仅 ``USE_LOCAL_EMBEDDING=true`` 时需要;远程 Embedding 可为空。"""
|
||||
return os.getenv("EMBEDDING_MODEL_PATH", "").strip()
|
||||
|
||||
|
||||
def resolve_sql_dialect(name: str) -> str:
|
||||
@@ -55,20 +53,23 @@ def resolve_sql_dialect(name: str) -> str:
|
||||
def setup_environment():
|
||||
"""环境检查"""
|
||||
_load_project_env()
|
||||
use_local_emb = os.getenv("USE_LOCAL_EMBEDDING", "true").strip().lower() in (
|
||||
use_local_emb = os.getenv("USE_LOCAL_EMBEDDING", "false").strip().lower() in (
|
||||
"1", "true", "yes", "on",
|
||||
)
|
||||
if use_local_emb:
|
||||
# 与 .env 中 EMBEDDING_MODEL_PATH 及 utils.embedding 一致
|
||||
model_path = Path(_embedding_model_path())
|
||||
mp_str = _embedding_model_path()
|
||||
if not mp_str:
|
||||
logger.warning(
|
||||
"USE_LOCAL_EMBEDDING=true 但未设置 EMBEDDING_MODEL_PATH(本地模型目录)"
|
||||
)
|
||||
return False
|
||||
model_path = Path(mp_str)
|
||||
if not model_path.exists():
|
||||
logger.warning(f"Embedding模型不存在: {model_path}")
|
||||
logger.info("请先下载模型:")
|
||||
logger.info(" modelscope download --model 'Qwen/Qwen3-Embedding-0.6B' "
|
||||
f"--local_dir '{model_path}'")
|
||||
logger.info("或使用远程 Embedding API:USE_LOCAL_EMBEDDING=false,并配置 "
|
||||
"MODELSCOPE_API_KEY、或 OPENAI_API_KEY+OPENAI_EMBEDDING_MODEL"
|
||||
"(及可选 OPENAI_BASE_URL)、或 DASHSCOPE_*(百炼)")
|
||||
logger.warning("本地 Embedding 目录不存在: %s", model_path)
|
||||
logger.info(
|
||||
"请设置正确的 EMBEDDING_MODEL_PATH,或改用 USE_LOCAL_EMBEDDING=false,"
|
||||
"并配置 OPENAI_* / MODELSCOPE_* / DASHSCOPE_* 远程 Embedding"
|
||||
)
|
||||
return False
|
||||
else:
|
||||
ms_key = os.getenv("MODELSCOPE_API_KEY", "").strip()
|
||||
@@ -206,7 +207,7 @@ def single_query(orchestrator, question: str, dialect: str = "tsql"):
|
||||
|
||||
logger.info(f"[Q] 问题: {question}")
|
||||
|
||||
classified = classify_dialog(question)
|
||||
classified = classify_dialog(question, llm_client=orchestrator.deepseek)
|
||||
if classified.intent == DialogIntent.CONVERSATION:
|
||||
reply = classified.reply_suggestion or ""
|
||||
logger.info("[Q] 意图: conversation(跳过 SQL 生成)")
|
||||
|
||||
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,
|
||||
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -1,17 +1,27 @@
|
||||
"""
|
||||
用户输入意图分类:区分「自然语言查数 / Text2SQL」与「寒暄、致谢、元问题」等不适合直接生成 SQL 的对话。
|
||||
|
||||
支持环境变量 ``DIALOG_INTENT_CLASSIFIER``:
|
||||
- ``rules``:仅用关键词与短语规则(无 LLM 调用)。
|
||||
- ``hybrid``(默认):明显查数词/寒暄/元问题走规则;其余交 LLM 判断(需调用方传入 ``llm_client``)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import os
|
||||
import re
|
||||
import unicodedata
|
||||
from enum import Enum
|
||||
from typing import NamedTuple, Optional
|
||||
from typing import TYPE_CHECKING, NamedTuple, Optional
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from llm.deepseek_client import DeepSeekClient
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_INTENT_CLASSIFIER_ENV = "DIALOG_INTENT_CLASSIFIER"
|
||||
|
||||
|
||||
class DialogIntent(str, Enum):
|
||||
TEXT2SQL = "text2sql"
|
||||
@@ -100,12 +110,24 @@ def _strip_trailing_punct(t: str) -> str:
|
||||
return re.sub(r"[!!。.??,,;;:~~…、]+$", "", t).strip()
|
||||
|
||||
|
||||
def classify_dialog(user_text: str) -> DialogClassifyResult:
|
||||
"""
|
||||
对用户一轮输入做粗分类。
|
||||
def _intent_classifier_mode() -> str:
|
||||
m = (os.getenv(_INTENT_CLASSIFIER_ENV) or "hybrid").strip().lower()
|
||||
if m not in ("rules", "hybrid"):
|
||||
logger.warning(
|
||||
"[dialog] unknown DIALOG_INTENT_CLASSIFIER=%r, use hybrid", m
|
||||
)
|
||||
return "hybrid"
|
||||
return m
|
||||
|
||||
策略:优先用「查询/业务」关键词锁定 TEXT2SQL;否则对短寒暄、致谢、元问题判为 CONVERSATION;
|
||||
其余默认 TEXT2SQL,避免漏判真实查询。
|
||||
|
||||
def _classify_dialog_rules(
|
||||
user_text: str,
|
||||
*,
|
||||
last_turn_was_data_query: bool = False,
|
||||
) -> DialogClassifyResult:
|
||||
"""
|
||||
规则分类(与历史行为一致):优先查数关键词;寒暄/元问题为 conversation;
|
||||
若上一轮为数据查询且非寒暄/元问题,则倾向 text2sql。
|
||||
"""
|
||||
t = _normalize(user_text)
|
||||
if not t:
|
||||
@@ -131,5 +153,135 @@ def classify_dialog(user_text: str) -> DialogClassifyResult:
|
||||
reply_suggestion=DEFAULT_CONVERSATION_REPLY,
|
||||
)
|
||||
|
||||
if last_turn_was_data_query:
|
||||
logger.debug("[dialog] intent=text2sql (follow-up after data query)")
|
||||
return DialogClassifyResult(DialogIntent.TEXT2SQL, None)
|
||||
|
||||
logger.debug("[dialog] intent=text2sql (default)")
|
||||
return DialogClassifyResult(DialogIntent.TEXT2SQL, None)
|
||||
|
||||
|
||||
def _classify_dialog_llm(
|
||||
client: DeepSeekClient,
|
||||
user_text: str,
|
||||
*,
|
||||
last_turn_was_data_query: bool,
|
||||
dialog_context: Optional[str],
|
||||
) -> DialogClassifyResult:
|
||||
from config.prompts import (
|
||||
DIALOG_INTENT_CLASSIFIER_SYSTEM,
|
||||
DIALOG_INTENT_CLASSIFIER_USER,
|
||||
)
|
||||
|
||||
last_label = "是" if last_turn_was_data_query else "否"
|
||||
ctx = (dialog_context or "").strip()
|
||||
if len(ctx) > 2400:
|
||||
ctx = ctx[:2400].rstrip() + "\n…(已截断)"
|
||||
if not ctx:
|
||||
ctx = "(无)"
|
||||
|
||||
raw = client.chat_with_json(
|
||||
[
|
||||
{"role": "system", "content": DIALOG_INTENT_CLASSIFIER_SYSTEM},
|
||||
{
|
||||
"role": "user",
|
||||
"content": DIALOG_INTENT_CLASSIFIER_USER.format(
|
||||
last_turn_label=last_label,
|
||||
context_snip=ctx,
|
||||
user_message=_normalize(user_text),
|
||||
),
|
||||
},
|
||||
],
|
||||
temperature=0.0,
|
||||
max_tokens=256,
|
||||
top_p=1.0,
|
||||
)
|
||||
if raw.get("_json_decode_failed"):
|
||||
raise ValueError("intent JSON decode failed")
|
||||
|
||||
intent_s = (raw.get("intent") or "").strip().lower()
|
||||
reply = (raw.get("reply_zh") or "").strip()
|
||||
|
||||
if intent_s == "conversation":
|
||||
if not reply:
|
||||
reply = (
|
||||
"若上一版 SQL 或结果不符合预期,请具体说明:希望增加/修改哪些条件、"
|
||||
"时间范围或统计维度,以便重新生成。"
|
||||
)
|
||||
logger.info("[dialog] intent=conversation (LLM)")
|
||||
return DialogClassifyResult(DialogIntent.CONVERSATION, reply_suggestion=reply)
|
||||
|
||||
if intent_s == "text2sql":
|
||||
logger.info("[dialog] intent=text2sql (LLM)")
|
||||
return DialogClassifyResult(DialogIntent.TEXT2SQL, None)
|
||||
|
||||
raise ValueError(f"unexpected intent field: {intent_s!r}")
|
||||
|
||||
|
||||
def classify_dialog(
|
||||
user_text: str,
|
||||
*,
|
||||
last_turn_was_data_query: bool = False,
|
||||
dialog_context: Optional[str] = None,
|
||||
llm_client: Optional[DeepSeekClient] = None,
|
||||
) -> DialogClassifyResult:
|
||||
"""
|
||||
对用户一轮输入做分类。
|
||||
|
||||
``DIALOG_INTENT_CLASSIFIER``:
|
||||
- ``rules``:仅规则。
|
||||
- ``hybrid``:规则快速命中查数词/寒暄/元问题后返回;否则在有 ``llm_client`` 时用 LLM,
|
||||
失败或无客户端时回退规则。
|
||||
|
||||
Args:
|
||||
user_text: 用户输入。
|
||||
last_turn_was_data_query: 上一轮助手是否为数据查询(供规则与 LLM 参考)。
|
||||
dialog_context: 会话摘要,供 LLM 参考(可选)。
|
||||
llm_client: DeepSeek 客户端;hybrid 下 LLM 分支需要。
|
||||
"""
|
||||
mode = _intent_classifier_mode()
|
||||
|
||||
if mode == "rules":
|
||||
return _classify_dialog_rules(
|
||||
user_text, last_turn_was_data_query=last_turn_was_data_query
|
||||
)
|
||||
|
||||
# hybrid
|
||||
t = _normalize(user_text)
|
||||
if not t:
|
||||
return DialogClassifyResult(
|
||||
DialogIntent.CONVERSATION, reply_suggestion=_EMPTY_INPUT_REPLY
|
||||
)
|
||||
|
||||
if _SQL_OR_QUERY_HINT_RE.search(t):
|
||||
logger.debug("[dialog] intent=text2sql (query/business hint, hybrid fast)")
|
||||
return DialogClassifyResult(DialogIntent.TEXT2SQL, None)
|
||||
|
||||
core = _strip_trailing_punct(t)
|
||||
if core.casefold() in _CHITCHAT_KEYS:
|
||||
logger.debug("[dialog] intent=conversation (chitchat phrase, hybrid fast)")
|
||||
return DialogClassifyResult(
|
||||
DialogIntent.CONVERSATION, reply_suggestion=DEFAULT_CONVERSATION_REPLY
|
||||
)
|
||||
|
||||
if _META_QUESTION_RE.search(t):
|
||||
logger.debug("[dialog] intent=conversation (meta question, hybrid fast)")
|
||||
return DialogClassifyResult(
|
||||
DialogIntent.CONVERSATION,
|
||||
reply_suggestion=DEFAULT_CONVERSATION_REPLY,
|
||||
)
|
||||
|
||||
if llm_client is not None:
|
||||
try:
|
||||
return _classify_dialog_llm(
|
||||
llm_client,
|
||||
user_text,
|
||||
last_turn_was_data_query=last_turn_was_data_query,
|
||||
dialog_context=dialog_context,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning("[dialog] LLM intent failed, fallback rules: %s", e)
|
||||
|
||||
return _classify_dialog_rules(
|
||||
user_text, last_turn_was_data_query=last_turn_was_data_query
|
||||
)
|
||||
|
||||
@@ -0,0 +1,130 @@
|
||||
"""
|
||||
从 Lite NL 会话消息构造 Text2SQL 可用的上文摘要,并判断上一轮是否为数据查询(用于续问意图分类)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
from typing import Any, Dict, List, Tuple
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_MAX_ASSISTANT_SQL_CHARS = 1400
|
||||
_MAX_ASSISTANT_TEXT_CHARS = 900
|
||||
_MAX_BLOCK_CHARS = 7500
|
||||
|
||||
|
||||
def _truncate(s: str, max_len: int) -> str:
|
||||
s = (s or "").strip()
|
||||
if len(s) <= max_len:
|
||||
return s
|
||||
return s[: max_len - 1].rstrip() + "…"
|
||||
|
||||
|
||||
def summarize_assistant_nl_payload(content: str) -> str:
|
||||
"""将落库的 assistant JSON 转为一小段可读摘要。"""
|
||||
raw = (content or "").strip()
|
||||
if not raw:
|
||||
return ""
|
||||
try:
|
||||
data: Dict[str, Any] = json.loads(raw)
|
||||
except json.JSONDecodeError:
|
||||
return _truncate(raw, _MAX_ASSISTANT_TEXT_CHARS)
|
||||
|
||||
br = data.get("branch_result") or {}
|
||||
if isinstance(br.get("answer"), str) and br["answer"].strip():
|
||||
return "[助手] " + _truncate(br["answer"].strip(), _MAX_ASSISTANT_TEXT_CHARS)
|
||||
|
||||
sql = (br.get("sql") or "").strip() if isinstance(br.get("sql"), str) else ""
|
||||
if sql:
|
||||
parts = ["[上轮 SQL] " + _truncate(sql, _MAX_ASSISTANT_SQL_CHARS)]
|
||||
if br.get("follow_up_required"):
|
||||
parts.append("[状态] 上轮为库探针0行,需补充条件后重问")
|
||||
for key in ("db_empty_feedback", "sql_explain"):
|
||||
v = br.get(key)
|
||||
if isinstance(v, str) and v.strip():
|
||||
parts.append("[说明] " + _truncate(v.strip(), 500))
|
||||
break
|
||||
return "\n".join(parts)
|
||||
|
||||
intent = (data.get("intent") or {}).get("intent")
|
||||
return _truncate(f"[助手] intent={intent}", _MAX_ASSISTANT_TEXT_CHARS)
|
||||
|
||||
|
||||
def last_assistant_was_data_query(items: List[Dict[str, Any]]) -> bool:
|
||||
"""最后一条 assistant 消息是否为数据查询分支(含 SQL 或 DATA_QUERY intent)。"""
|
||||
for m in reversed(items or []):
|
||||
if m.get("role") != "assistant":
|
||||
continue
|
||||
raw = (m.get("content") or "").strip()
|
||||
if not raw:
|
||||
return False
|
||||
try:
|
||||
data = json.loads(raw)
|
||||
except json.JSONDecodeError:
|
||||
return False
|
||||
intent = (data.get("intent") or {}).get("intent")
|
||||
if intent == "DATA_QUERY":
|
||||
return True
|
||||
br = data.get("branch_result") or {}
|
||||
sql = (br.get("sql") or "").strip() if isinstance(br.get("sql"), str) else ""
|
||||
return bool(sql)
|
||||
return False
|
||||
|
||||
|
||||
def messages_to_text2sql_context(
|
||||
items: List[Dict[str, Any]],
|
||||
*,
|
||||
max_chars: int = _MAX_BLOCK_CHARS,
|
||||
max_pairs: int = 8,
|
||||
) -> Tuple[str, int]:
|
||||
"""
|
||||
将会话消息列表转为供模型阅读的「上文」文本(不含本轮用户输入)。
|
||||
|
||||
Returns:
|
||||
(context_text, num_user_turns_included)
|
||||
"""
|
||||
if not items:
|
||||
return "", 0
|
||||
|
||||
pairs: List[Tuple[str, str]] = []
|
||||
i = 0
|
||||
n = len(items)
|
||||
while i < n:
|
||||
u = items[i]
|
||||
if u.get("role") != "user":
|
||||
i += 1
|
||||
continue
|
||||
user_text = (u.get("content") or "").strip()
|
||||
asst_text = ""
|
||||
if i + 1 < n and items[i + 1].get("role") == "assistant":
|
||||
asst_text = summarize_assistant_nl_payload(str(items[i + 1].get("content") or ""))
|
||||
i += 2
|
||||
else:
|
||||
i += 1
|
||||
if user_text or asst_text:
|
||||
pairs.append((user_text, asst_text))
|
||||
|
||||
if not pairs:
|
||||
return "", 0
|
||||
|
||||
pairs = pairs[-max_pairs:]
|
||||
lines: List[str] = []
|
||||
total = 0
|
||||
user_count = 0
|
||||
for idx, (uq, aq) in enumerate(pairs, start=1):
|
||||
chunk_parts = [f"--- 第{idx}轮 ---"]
|
||||
if uq:
|
||||
chunk_parts.append(f"用户:{uq}")
|
||||
if aq:
|
||||
chunk_parts.append(f"{aq}")
|
||||
chunk = "\n".join(chunk_parts)
|
||||
if total + len(chunk) + 2 > max_chars:
|
||||
break
|
||||
lines.append(chunk)
|
||||
total += len(chunk) + 2
|
||||
if uq:
|
||||
user_count += 1
|
||||
|
||||
return "\n\n".join(lines).strip(), user_count
|
||||
+19
-22
@@ -1,6 +1,6 @@
|
||||
"""
|
||||
Embedding 封装:本地 Qwen3-Embedding,或兼容 OpenAI /v1/embeddings 的远程 API
|
||||
(ModelScope 推理、阿里云 DashScope 等)。
|
||||
Embedding 封装:兼容 OpenAI /v1/embeddings 的远程 API(默认),
|
||||
或可选本地 HuggingFace 目录(USE_LOCAL_EMBEDDING=true + EMBEDDING_MODEL_PATH)。
|
||||
"""
|
||||
|
||||
import os
|
||||
@@ -112,10 +112,10 @@ def _remote_embedding_from_env() -> _RemoteEmbeddingEnv:
|
||||
|
||||
class Qwen3Embedding:
|
||||
"""
|
||||
Qwen3-Embedding-0.6B 向量化封装
|
||||
本地 HuggingFace 格式 Embedding 模型(Mean Pooling + L2,用于向量检索)。
|
||||
|
||||
使用 Mean Pooling 将token embeddings聚合为句子向量,
|
||||
并进行L2归一化以支持余弦相似度计算。
|
||||
仅在 ``USE_LOCAL_EMBEDDING=true`` 时使用;路径由 ``model_path`` 或环境变量
|
||||
``EMBEDDING_MODEL_PATH`` 指定,**不再内置默认目录**。
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
@@ -128,7 +128,7 @@ class Qwen3Embedding:
|
||||
初始化 embedding 模型
|
||||
|
||||
Args:
|
||||
model_path: 本地模型路径,若为None则从环境变量或默认路径加载
|
||||
model_path: 本地模型目录;None 或空字符串时读 ``EMBEDDING_MODEL_PATH``
|
||||
device: 推理设备('cpu', 'cuda', 'cuda:0'等),None则自动选择
|
||||
use_fp16: 是否使用FP16混合精度(GPU可用时建议开启,速度更快)
|
||||
"""
|
||||
@@ -138,22 +138,19 @@ class Qwen3Embedding:
|
||||
"pip install transformers torch sentencepiece accelerate"
|
||||
)
|
||||
|
||||
# 确定模型路径
|
||||
if model_path is None:
|
||||
model_path = os.getenv(
|
||||
"EMBEDDING_MODEL_PATH",
|
||||
"./data/models/Qwen3-Embedding-0.6B"
|
||||
resolved = (model_path or "").strip() or os.getenv("EMBEDDING_MODEL_PATH", "").strip()
|
||||
if not resolved:
|
||||
raise ValueError(
|
||||
"已启用本地 Embedding(USE_LOCAL_EMBEDDING=true),但未设置有效模型路径。"
|
||||
"请在 .env 中设置 EMBEDDING_MODEL_PATH 指向本地模型目录,"
|
||||
"或设置 USE_LOCAL_EMBEDDING=false 使用 OPENAI_* / MODELSCOPE_* / DASHSCOPE_* 远程接口。"
|
||||
)
|
||||
|
||||
model_path = Path(model_path)
|
||||
model_path = Path(resolved)
|
||||
if not model_path.exists():
|
||||
raise FileNotFoundError(
|
||||
f"模型目录不存在:{model_path}\n"
|
||||
"请先下载模型:\n"
|
||||
" modelscope download --model 'Qwen/Qwen3-Embedding-0.6B' "
|
||||
f"--local_dir '{model_path}'\n"
|
||||
"或从Hugging Face下载:git lfs install && git clone "
|
||||
f"https://huggingface.co/Qwen/Qwen3-Embedding-0.6B {model_path}"
|
||||
f"本地 Embedding 模型目录不存在:{model_path}\n"
|
||||
"请修正 EMBEDDING_MODEL_PATH,或改用 USE_LOCAL_EMBEDDING=false。"
|
||||
)
|
||||
|
||||
# 确定设备
|
||||
@@ -161,7 +158,7 @@ class Qwen3Embedding:
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
|
||||
self.device = device
|
||||
logger.info(f"加载Qwen3-Embedding模型:{model_path},设备:{device}")
|
||||
logger.info("加载本地 Embedding 模型:%s,设备:%s", model_path, device)
|
||||
|
||||
# 加载 tokenizer:fast(Rust) 解析 tokenizer.json 需较新 tokenizers;
|
||||
# 旧版本会报 ModelWrapper / untagged enum,回退到慢速 tokenizer 可恢复。
|
||||
@@ -499,13 +496,13 @@ def get_embedder(
|
||||
force_reload: bool = False,
|
||||
) -> Any:
|
||||
"""
|
||||
获取 Embedding 单例:USE_LOCAL_EMBEDDING=true 时用本地 Qwen3,否则用远程
|
||||
OpenAI 兼容 API(优先级见 OpenAICompatibleRemoteEmbedding)。
|
||||
获取 Embedding 单例:默认 ``USE_LOCAL_EMBEDDING=false``,使用远程
|
||||
OpenAI 兼容 / ModelScope / DashScope;为 true 时用本地目录(EMBEDDING_MODEL_PATH)。
|
||||
"""
|
||||
global _embedding_instance
|
||||
|
||||
if force_reload or _embedding_instance is None:
|
||||
if _env_flag("USE_LOCAL_EMBEDDING", "true"):
|
||||
if _env_flag("USE_LOCAL_EMBEDDING", "false"):
|
||||
_embedding_instance = Qwen3Embedding(
|
||||
model_path=model_path,
|
||||
device=device,
|
||||
|
||||
@@ -0,0 +1,214 @@
|
||||
"""
|
||||
Few-shot 经验样本的 Chroma 向量库存储与检索。
|
||||
|
||||
与 SchemaIndexer 一致:使用项目统一 ``get_embedder()``,持久化目录默认
|
||||
``./data/embeddings/chroma_fewshot``,集合名 ``fewshot_samples``。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, List, Optional, Tuple
|
||||
|
||||
import chromadb
|
||||
from chromadb.config import Settings as ChromaSettings
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
DEFAULT_FEWSHOT_CHROMA_DIR = "./data/embeddings/chroma_fewshot"
|
||||
COLLECTION_NAME = "fewshot_samples"
|
||||
# Chroma metadata 单值不宜过大,SQL/说明超长时截断
|
||||
_MAX_SQL_META = 16000
|
||||
_MAX_EXPLAIN_META = 6000
|
||||
_MAX_QEN_META = 2000
|
||||
|
||||
|
||||
class FewShotChromaStore:
|
||||
"""Few-shot JSONL → Chroma 写入与按向量检索。"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
embedder: Any,
|
||||
persist_dir: Optional[str] = None,
|
||||
collection_name: Optional[str] = None,
|
||||
):
|
||||
self.embedder = embedder
|
||||
# Chroma 的 collection 名;显式参数优先,否则读 FEWSHOT_CHROMA_COLLECTION,再回退默认
|
||||
explicit = (collection_name or "").strip()
|
||||
from_env = (os.getenv("FEWSHOT_CHROMA_COLLECTION") or "").strip()
|
||||
self.collection_name = explicit or from_env or COLLECTION_NAME
|
||||
self.persist_dir = Path(
|
||||
(persist_dir or os.getenv("FEWSHOT_CHROMA_PATH") or DEFAULT_FEWSHOT_CHROMA_DIR).strip()
|
||||
)
|
||||
self.persist_dir.mkdir(parents=True, exist_ok=True)
|
||||
self.client = chromadb.PersistentClient(
|
||||
path=str(self.persist_dir),
|
||||
settings=ChromaSettings(anonymized_telemetry=False),
|
||||
)
|
||||
self.collection = self.client.get_or_create_collection(
|
||||
name=self.collection_name,
|
||||
metadata={"hnsw:space": "cosine"},
|
||||
)
|
||||
logger.info(
|
||||
"[OK] FewShotChromaStore: path=%s collection=%s count=%s",
|
||||
self.persist_dir,
|
||||
self.collection_name,
|
||||
self.collection.count(),
|
||||
)
|
||||
|
||||
def count(self) -> int:
|
||||
return int(self.collection.count())
|
||||
|
||||
def clear(self) -> None:
|
||||
"""删除集合内全部文档(用于 force rebuild)。"""
|
||||
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("[OK] Few-shot Chroma 集合已清空并重建")
|
||||
|
||||
def build_from_samples(
|
||||
self,
|
||||
samples: List["ExperienceSample"],
|
||||
*,
|
||||
force_rebuild: bool = False,
|
||||
batch_size: int = 32,
|
||||
) -> int:
|
||||
"""
|
||||
将已解析的样本列表写入 Chroma(向量由 question_zh 编码)。
|
||||
|
||||
Returns:
|
||||
写入条数
|
||||
"""
|
||||
if not samples:
|
||||
logger.warning("Few-shot Chroma: 无样本,跳过构建")
|
||||
return 0
|
||||
|
||||
if force_rebuild and self.count() > 0:
|
||||
self.clear()
|
||||
|
||||
if not force_rebuild and self.count() > 0:
|
||||
logger.info("Few-shot Chroma 已有 %s 条,跳过构建(加 --force 可重建)", self.count())
|
||||
return self.count()
|
||||
|
||||
texts = [(s.question_zh or "").strip() or " " for s in samples]
|
||||
# Chroma 要求 ids 为唯一字符串
|
||||
ids = [str(s.qid) for s in samples]
|
||||
metadatas: List[Dict[str, Any]] = []
|
||||
for s in samples:
|
||||
metadatas.append(
|
||||
{
|
||||
"qid": str(s.qid),
|
||||
"question_en": ((s.question_en or "")[:_MAX_QEN_META]),
|
||||
"sql": (s.sql or "")[:_MAX_SQL_META],
|
||||
"explanation": (s.explanation or "")[:_MAX_EXPLAIN_META],
|
||||
"rating": int(s.rating) if s.rating is not None else -1,
|
||||
"difficulty": (s.difficulty or "medium")[:32],
|
||||
"tags_json": json.dumps(s.tags or [], ensure_ascii=False)[:4000],
|
||||
}
|
||||
)
|
||||
|
||||
logger.info("Few-shot Chroma: 计算 %s 条 embedding...", len(texts))
|
||||
embeddings = self.embedder.encode(
|
||||
texts,
|
||||
batch_size=min(batch_size, len(texts)),
|
||||
normalize=True,
|
||||
show_progress=True,
|
||||
)
|
||||
|
||||
self.collection.add(
|
||||
embeddings=embeddings.tolist(),
|
||||
documents=texts,
|
||||
metadatas=metadatas,
|
||||
ids=ids,
|
||||
)
|
||||
logger.info("[OK] Few-shot Chroma 索引完成: %s 条", len(ids))
|
||||
return len(ids)
|
||||
|
||||
def search_raw(
|
||||
self,
|
||||
query: str,
|
||||
top_k: int = 20,
|
||||
) -> List[Tuple[float, Dict[str, Any], str]]:
|
||||
"""
|
||||
Returns:
|
||||
(score, metadata_dict, document) 列表,score 同 SchemaIndexer 为 1-distance
|
||||
"""
|
||||
n = self.count()
|
||||
if n == 0:
|
||||
return []
|
||||
q_emb = self.embedder.encode(
|
||||
[(query or "").strip() or " "],
|
||||
batch_size=1,
|
||||
normalize=True,
|
||||
show_progress=False,
|
||||
)
|
||||
k = min(max(top_k, 1), n)
|
||||
results = self.collection.query(
|
||||
query_embeddings=q_emb.tolist(),
|
||||
n_results=k,
|
||||
)
|
||||
out: List[Tuple[float, Dict[str, Any], str]] = []
|
||||
if not results["ids"] or not results["ids"][0]:
|
||||
return out
|
||||
for tid, dist, meta, doc in zip(
|
||||
results["ids"][0],
|
||||
results["distances"][0],
|
||||
results["metadatas"][0],
|
||||
results["documents"][0],
|
||||
):
|
||||
score = 1.0 - float(dist)
|
||||
m = dict(meta or {})
|
||||
m["qid"] = tid
|
||||
out.append((score, m, doc or ""))
|
||||
return out
|
||||
|
||||
def get_all_as_samples(self) -> List[Any]:
|
||||
"""导出集合中全部样本(用于统计/按标签列举;条数大时慎用)。"""
|
||||
if self.count() == 0:
|
||||
return []
|
||||
data = self.collection.get(include=["metadatas", "documents"])
|
||||
ids = data.get("ids") or []
|
||||
metas = data.get("metadatas") or []
|
||||
docs = data.get("documents") or []
|
||||
out: List[Any] = []
|
||||
for i, qid in enumerate(ids):
|
||||
meta = dict(metas[i] or {}) if i < len(metas) else {}
|
||||
meta["qid"] = qid
|
||||
doc = docs[i] if i < len(docs) else ""
|
||||
out.append(sample_from_chroma_metadata(meta, doc))
|
||||
return out
|
||||
|
||||
|
||||
def sample_from_chroma_metadata(meta: Dict[str, Any], document: str) -> "ExperienceSample":
|
||||
from utils.fewshot_selector import ExperienceSample
|
||||
|
||||
tags_raw = meta.get("tags_json") or "[]"
|
||||
try:
|
||||
tags = json.loads(tags_raw) if isinstance(tags_raw, str) else []
|
||||
except json.JSONDecodeError:
|
||||
tags = []
|
||||
if not isinstance(tags, list):
|
||||
tags = []
|
||||
r = meta.get("rating")
|
||||
try:
|
||||
ri = int(r) if r is not None else -1
|
||||
except (TypeError, ValueError):
|
||||
ri = -1
|
||||
rating = ri if ri >= 0 else None
|
||||
qen = meta.get("question_en")
|
||||
return ExperienceSample(
|
||||
qid=str(meta.get("qid", "")),
|
||||
question_zh=document or str(meta.get("question_zh", "")),
|
||||
question_en=qen if qen else None,
|
||||
sql=str(meta.get("sql", "")),
|
||||
explanation=str(meta.get("explanation", "")),
|
||||
rating=rating,
|
||||
tags=tags,
|
||||
difficulty=str(meta.get("difficulty", "medium") or "medium"),
|
||||
)
|
||||
@@ -7,7 +7,9 @@ Few-shot示例选择器 - 基于经验数据集动态选择相关示例
|
||||
用法:
|
||||
from utils.fewshot_selector import FewShotSelector
|
||||
|
||||
selector = FewShotSelector("data/experiences/all_samples.jsonl")
|
||||
# Chroma 模式(FEWSHOT_USE_CHROMA=true):可不传 JSONL,路径传 None 或 ""
|
||||
selector = FewShotSelector(None)
|
||||
# 或传统:FewShotSelector("data/experiences/all_samples.jsonl")
|
||||
examples = selector.select(question="查询2024年1月的销售额", top_k=3)
|
||||
|
||||
# 在Prompt中使用
|
||||
@@ -24,7 +26,8 @@ import logging
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_DEFAULT_LOCAL_EMBED_PATH = "./data/models/Qwen3-Embedding-0.6B"
|
||||
# 仅 USE_LOCAL_EMBEDDING=true 时通过 EMBEDDING_MODEL_PATH 使用;远程模式留空即可
|
||||
_DEFAULT_LOCAL_EMBED_PATH = ""
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -80,50 +83,128 @@ class FewShotSelector:
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
samples_path: str,
|
||||
samples_path: Optional[str] = None,
|
||||
embedding_model_path: Optional[str] = None,
|
||||
use_cache: bool = True,
|
||||
use_chroma: Optional[bool] = None,
|
||||
chroma_persist_dir: Optional[str] = None,
|
||||
):
|
||||
"""
|
||||
Args:
|
||||
samples_path: 样本JSONL文件路径
|
||||
samples_path: 样本 JSONL 路径;**空字符串**表示不读文件(仅当 ``FEWSHOT_USE_CHROMA=true``
|
||||
且 Chroma 中已有数据时可用)。构建索引仍请用 ``scripts/build_fewshot_chroma_index.py --samples``。
|
||||
embedding_model_path: 本地 Embedding 模型目录;None 时用环境变量
|
||||
EMBEDDING_MODEL_PATH(仅 USE_LOCAL_EMBEDDING=true 时有效)
|
||||
use_cache: 是否缓存样本向量(按向量维度分文件,换模型会自动重建)
|
||||
use_chroma: 是否使用 Chroma 持久化向量库;None 时读环境变量 FEWSHOT_USE_CHROMA
|
||||
chroma_persist_dir: Chroma 目录;None 时用 FEWSHOT_CHROMA_PATH 或默认 chroma_fewshot
|
||||
"""
|
||||
self.samples_path = Path(samples_path)
|
||||
sp = (samples_path or "").strip()
|
||||
self.samples_path: Optional[Path] = Path(sp) if sp else None
|
||||
self.samples: List[ExperienceSample] = []
|
||||
self._embedder = None
|
||||
self.embeddings: Optional[np.ndarray] = None
|
||||
self.use_cache = use_cache
|
||||
self.cache_path: Optional[Path] = None
|
||||
self._chroma_store = None
|
||||
self.use_chroma = (
|
||||
use_chroma
|
||||
if use_chroma is not None
|
||||
else os.getenv("FEWSHOT_USE_CHROMA", "false").lower() in ("1", "true", "yes")
|
||||
)
|
||||
self._chroma_persist_dir = chroma_persist_dir
|
||||
self._embedding_model_path = (
|
||||
embedding_model_path
|
||||
if embedding_model_path is not None
|
||||
else os.getenv("EMBEDDING_MODEL_PATH", _DEFAULT_LOCAL_EMBED_PATH).strip()
|
||||
)
|
||||
|
||||
self._load_samples()
|
||||
self._build_index()
|
||||
if self.use_chroma:
|
||||
self._init_chroma_mode()
|
||||
else:
|
||||
self._load_samples()
|
||||
self._init_embedder_and_vector_index()
|
||||
|
||||
def _init_chroma_mode(self) -> None:
|
||||
"""Chroma 模式:先连向量库;库非空则不再读 JSONL。"""
|
||||
from utils.embedding import get_embedder
|
||||
from utils.fewshot_chroma_store import FewShotChromaStore
|
||||
|
||||
self._embedder = get_embedder(self._embedding_model_path)
|
||||
self._chroma_store = FewShotChromaStore(
|
||||
self._embedder,
|
||||
persist_dir=self._chroma_persist_dir,
|
||||
)
|
||||
n = self._chroma_store.count()
|
||||
if n > 0:
|
||||
logger.info(
|
||||
"[Few-shot] 已从向量库加载(Chroma %s 条,%s)",
|
||||
n,
|
||||
self._chroma_store.persist_dir,
|
||||
)
|
||||
else:
|
||||
self._load_samples()
|
||||
if self.samples:
|
||||
logger.info(
|
||||
"Few-shot Chroma 库为空,正从 JSONL 写入向量索引: %s",
|
||||
self.samples_path,
|
||||
)
|
||||
self._chroma_store.build_from_samples(self.samples, force_rebuild=False)
|
||||
elif self.samples_path is None:
|
||||
logger.warning(
|
||||
"[Few-shot] Chroma 库为空且未配置 JSONL;请运行 "
|
||||
"scripts/build_fewshot_chroma_index.py 或设置 FEWSHOT_DATA_PATH"
|
||||
)
|
||||
elif not self.samples_path.exists():
|
||||
logger.warning(
|
||||
"[Few-shot] Chroma 库为空且 JSONL 不存在: %s;请先灌库",
|
||||
self.samples_path,
|
||||
)
|
||||
else:
|
||||
logger.warning(
|
||||
"[Few-shot] Chroma 库为空且 JSONL 无有效样本: %s",
|
||||
self.samples_path,
|
||||
)
|
||||
|
||||
self.embeddings = None
|
||||
self.cache_path = None
|
||||
logger.info("[OK] Few-shot 使用 Chroma(%s 条)", self._chroma_store.count())
|
||||
|
||||
def _load_samples(self):
|
||||
"""加载样本数据"""
|
||||
"""从 JSONL 加载样本到内存(非 Chroma 模式必需;Chroma 空库时用于首次灌库)。"""
|
||||
if self.samples_path is None:
|
||||
if self.use_chroma:
|
||||
logger.info("[Few-shot] 未配置 JSONL 路径,运行时仅从 Chroma 检索")
|
||||
return
|
||||
raise ValueError("未指定样本 JSONL 路径且未启用 Chroma(FEWSHOT_USE_CHROMA)")
|
||||
|
||||
if not self.samples_path.exists():
|
||||
if self.use_chroma:
|
||||
logger.warning(
|
||||
"[Few-shot] JSONL 不存在 %s,跳过文件加载,仅从 Chroma 检索",
|
||||
self.samples_path,
|
||||
)
|
||||
return
|
||||
raise FileNotFoundError(f"样本文件不存在: {self.samples_path}")
|
||||
|
||||
logger.info(f"加载样本: {self.samples_path}")
|
||||
with self.samples_path.open('r', encoding='utf-8') as f:
|
||||
logger.info("从 JSONL 加载样本: %s", self.samples_path)
|
||||
with self.samples_path.open("r", encoding="utf-8") as f:
|
||||
for line in f:
|
||||
data = json.loads(line.strip())
|
||||
self.samples.append(ExperienceSample.from_dict(data))
|
||||
|
||||
logger.info(f"[OK] 加载 {len(self.samples)} 个样本")
|
||||
logger.info("[OK] 自 JSONL 加载 %s 个样本", len(self.samples))
|
||||
|
||||
def _build_index(self):
|
||||
"""用项目统一 Embedder 构建语义索引"""
|
||||
def _init_embedder_and_vector_index(self) -> None:
|
||||
"""非 Chroma:初始化 Embedder 与内存 numpy 索引。"""
|
||||
from utils.embedding import get_embedder
|
||||
|
||||
self._embedder = get_embedder(self._embedding_model_path)
|
||||
self._build_numpy_index()
|
||||
|
||||
def _build_numpy_index(self) -> None:
|
||||
"""内存向量 + 可选 .npy 缓存(与历史行为一致)。"""
|
||||
assert self._embedder is not None
|
||||
probe = self._embedder.encode(
|
||||
[" "],
|
||||
batch_size=1,
|
||||
@@ -168,6 +249,43 @@ class FewShotSelector:
|
||||
np.save(self.cache_path, self.embeddings)
|
||||
logger.info(f"[OK] Few-shot 向量已缓存: {self.cache_path}")
|
||||
|
||||
def _select_chroma(
|
||||
self,
|
||||
question: str,
|
||||
top_k: int,
|
||||
min_rating: Optional[int],
|
||||
required_tags: Optional[List[str]],
|
||||
max_difficulty: str,
|
||||
exclude_qids: Optional[List[str]],
|
||||
) -> List[ExperienceSample]:
|
||||
from utils.fewshot_chroma_store import sample_from_chroma_metadata
|
||||
|
||||
assert self._chroma_store is not None
|
||||
over_fetch = max(top_k * 12, 48)
|
||||
rows = self._chroma_store.search_raw(question, top_k=over_fetch)
|
||||
candidates: List[tuple[float, ExperienceSample]] = []
|
||||
for score, meta, doc in rows:
|
||||
s = sample_from_chroma_metadata(meta, doc)
|
||||
if exclude_qids and s.qid in exclude_qids:
|
||||
continue
|
||||
if min_rating and s.rating is not None and s.rating < min_rating:
|
||||
continue
|
||||
if max_difficulty == "easy" and s.difficulty != "easy":
|
||||
continue
|
||||
if max_difficulty == "medium" and s.difficulty == "hard":
|
||||
continue
|
||||
if required_tags and not all(tag in s.tags for tag in required_tags):
|
||||
continue
|
||||
candidates.append((score, s))
|
||||
|
||||
candidates.sort(key=lambda x: -x[0])
|
||||
selected = [s for _, s in candidates[:top_k]]
|
||||
logger.info(
|
||||
f"Few-shot(Chroma)选择: 问题='{question[:30]}...' → 选中{len(selected)}个示例 "
|
||||
f"(top_k={top_k}, min_rating={min_rating})"
|
||||
)
|
||||
return selected
|
||||
|
||||
def select(
|
||||
self,
|
||||
question: str,
|
||||
@@ -191,6 +309,16 @@ class FewShotSelector:
|
||||
Returns:
|
||||
排序后的示例列表(最相关优先)
|
||||
"""
|
||||
if self._chroma_store is not None:
|
||||
return self._select_chroma(
|
||||
question,
|
||||
top_k,
|
||||
min_rating,
|
||||
required_tags,
|
||||
max_difficulty,
|
||||
exclude_qids,
|
||||
)
|
||||
|
||||
if self._embedder is None or self.embeddings is None:
|
||||
logger.error("Few-shot 索引未初始化")
|
||||
return []
|
||||
@@ -266,10 +394,18 @@ class FewShotSelector:
|
||||
|
||||
return "\n".join(lines)
|
||||
|
||||
def _all_samples_for_aggregation(self) -> List[ExperienceSample]:
|
||||
"""内存中的 JSONL 样本,或 Chroma 全量导出(用于统计/按标签列举)。"""
|
||||
if self.samples:
|
||||
return self.samples
|
||||
if self._chroma_store is not None and self._chroma_store.count() > 0:
|
||||
return self._chroma_store.get_all_as_samples()
|
||||
return []
|
||||
|
||||
def get_tagged_examples(self, tags: List[str], top_k_per_tag: int = 2) -> str:
|
||||
"""获取特定标签的示例"""
|
||||
tagged_samples = []
|
||||
for sample in self.samples:
|
||||
for sample in self._all_samples_for_aggregation():
|
||||
if any(tag in sample.tags for tag in tags):
|
||||
tagged_samples.append(sample)
|
||||
|
||||
@@ -287,27 +423,33 @@ class FewShotSelector:
|
||||
|
||||
def get_stats(self) -> Dict:
|
||||
"""获取数据集统计"""
|
||||
src = self._all_samples_for_aggregation()
|
||||
stats = {
|
||||
"total": len(self.samples),
|
||||
"total": len(src),
|
||||
"by_rating": {},
|
||||
"by_difficulty": {},
|
||||
"by_tag": {},
|
||||
"avg_rating": 0.0
|
||||
"avg_rating": 0.0,
|
||||
"source": (
|
||||
"chroma"
|
||||
if (not self.samples and self._chroma_store and self._chroma_store.count() > 0)
|
||||
else ("jsonl" if self.samples else "empty")
|
||||
),
|
||||
}
|
||||
|
||||
ratings = [s.rating for s in self.samples if s.rating]
|
||||
ratings = [s.rating for s in src if s.rating]
|
||||
if ratings:
|
||||
stats["avg_rating"] = sum(ratings) / len(ratings)
|
||||
for r in range(1, 11):
|
||||
stats["by_rating"][r] = sum(1 for s in self.samples if s.rating == r)
|
||||
stats["by_rating"][r] = sum(1 for s in src if s.rating == r)
|
||||
|
||||
for diff in ["easy", "medium", "hard"]:
|
||||
stats["by_difficulty"][diff] = sum(
|
||||
1 for s in self.samples if s.difficulty == diff
|
||||
1 for s in src if s.difficulty == diff
|
||||
)
|
||||
|
||||
tag_counts = {}
|
||||
for s in self.samples:
|
||||
for s in src:
|
||||
for tag in s.tags:
|
||||
tag_counts[tag] = tag_counts.get(tag, 0) + 1
|
||||
stats["by_tag"] = tag_counts
|
||||
|
||||
Reference in New Issue
Block a user