2026-04-10 16:52:07 +08:00
|
|
|
|
"""
|
|
|
|
|
|
Text2SQL 多智能体编排器
|
|
|
|
|
|
协调 Schema Linker、SQL Generator、Validator 三个Agent
|
|
|
|
|
|
"""
|
|
|
|
|
|
|
|
|
|
|
|
import logging
|
|
|
|
|
|
import os # 新增
|
2026-04-16 13:48:44 +08:00
|
|
|
|
from typing import Callable, Dict, List, Optional, Tuple
|
2026-04-10 16:52:07 +08:00
|
|
|
|
from dataclasses import dataclass, field
|
|
|
|
|
|
|
|
|
|
|
|
from schema.manager import SchemaManager
|
|
|
|
|
|
from schema.indexer import SchemaIndexer
|
|
|
|
|
|
from llm.deepseek_client import DeepSeekClient, DeepSeekConfig
|
2026-04-16 13:48:44 +08:00
|
|
|
|
from utils.fewshot_selector import ExperienceSample, FewShotSelector # 新增
|
2026-04-10 16:52:07 +08:00
|
|
|
|
|
|
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
@dataclass
|
|
|
|
|
|
class GenerationResult:
|
|
|
|
|
|
"""SQL生成结果"""
|
|
|
|
|
|
sql: str
|
|
|
|
|
|
valid: bool
|
|
|
|
|
|
errors: List[str] = field(default_factory=list)
|
|
|
|
|
|
warnings: List[str] = field(default_factory=list)
|
|
|
|
|
|
tables_used: List[str] = field(default_factory=list)
|
|
|
|
|
|
attempts: int = 1
|
|
|
|
|
|
reasoning: Optional[str] = None
|
|
|
|
|
|
metadata: Dict[str, any] = field(default_factory=dict)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
class Text2SQLOrchestrator:
|
|
|
|
|
|
"""
|
|
|
|
|
|
Text2SQL 多智能体编排器
|
|
|
|
|
|
|
|
|
|
|
|
工作流程:
|
2026-04-14 10:28:22 +08:00
|
|
|
|
1. 粗筛候选表
|
|
|
|
|
|
2. Schema Linker:LLM 精筛表
|
|
|
|
|
|
3. 外键扩展 → 拼 Schema 子集
|
|
|
|
|
|
4. SQL Generator:生成 SQL
|
|
|
|
|
|
5. Validator:验证 SQL
|
2026-04-10 16:52:07 +08:00
|
|
|
|
"""
|
|
|
|
|
|
|
|
|
|
|
|
def __init__(
|
|
|
|
|
|
self,
|
|
|
|
|
|
schema_manager: SchemaManager,
|
2026-04-16 09:15:01 +08:00
|
|
|
|
llm_client: Optional[object] = None,
|
2026-04-10 16:52:07 +08:00
|
|
|
|
deepseek_api_key: Optional[str] = None,
|
|
|
|
|
|
deepseek_config: Optional[DeepSeekConfig] = None,
|
|
|
|
|
|
vector_db_path: str = "./data/embeddings/chroma",
|
|
|
|
|
|
max_retry: int = 2,
|
|
|
|
|
|
use_vector_search: bool = True,
|
|
|
|
|
|
# Few-shot配置
|
|
|
|
|
|
fewshot_enabled: bool = True,
|
|
|
|
|
|
fewshot_samples_path: Optional[str] = None,
|
|
|
|
|
|
fewshot_top_k: int = 3,
|
|
|
|
|
|
fewshot_min_rating: int = 7,
|
2026-04-14 10:28:22 +08:00
|
|
|
|
translate_english_to_zh: bool = True,
|
2026-04-10 16:52:07 +08:00
|
|
|
|
):
|
|
|
|
|
|
"""
|
|
|
|
|
|
初始化编排器
|
|
|
|
|
|
|
|
|
|
|
|
Args:
|
|
|
|
|
|
schema_manager: Schema管理器实例
|
2026-04-16 09:15:01 +08:00
|
|
|
|
llm_client: 可选:外部传入的 LLM Client(需具备 chat/chat_with_json 等方法)。
|
2026-04-10 16:52:07 +08:00
|
|
|
|
deepseek_api_key: DeepSeek API密钥(也可通过环境变量DEEPSEEK_API_KEY)
|
|
|
|
|
|
deepseek_config: DeepSeek配置对象(优先于api_key)
|
|
|
|
|
|
vector_db_path: 向量数据库路径
|
|
|
|
|
|
max_retry: 最大重试次数(包含首次生成)
|
|
|
|
|
|
use_vector_search: 是否使用向量检索粗筛
|
2026-04-15 11:25:19 +08:00
|
|
|
|
translate_english_to_zh: 为 True 时(默认)对**所有**非空问句做一次 LLM 归一(temperature=0),
|
|
|
|
|
|
输出一句标准中文供检索与生成;使同一语义的中英文表述对齐,从而 SQL 一致。为 False 时
|
|
|
|
|
|
不做归一(原样英文/中文)。环境变量 ``TRANSLATE_EN_TO_ZH=false`` 可关闭。
|
2026-04-10 16:52:07 +08:00
|
|
|
|
"""
|
|
|
|
|
|
self.schema_manager = schema_manager
|
|
|
|
|
|
self.max_retry = max_retry
|
|
|
|
|
|
self.use_vector_search = use_vector_search
|
2026-04-14 10:28:22 +08:00
|
|
|
|
self.translate_english_to_zh = translate_english_to_zh
|
2026-04-10 16:52:07 +08:00
|
|
|
|
|
2026-04-16 09:15:01 +08:00
|
|
|
|
# 初始化 LLM 客户端(历史属性名保留为 deepseek,避免大范围改动)
|
|
|
|
|
|
if llm_client is not None:
|
|
|
|
|
|
self.deepseek = llm_client
|
2026-04-10 16:52:07 +08:00
|
|
|
|
else:
|
2026-04-16 09:15:01 +08:00
|
|
|
|
if deepseek_config:
|
|
|
|
|
|
self.deepseek = DeepSeekClient(deepseek_config)
|
|
|
|
|
|
else:
|
|
|
|
|
|
self.deepseek = DeepSeekClient(DeepSeekConfig(api_key=deepseek_api_key))
|
2026-04-10 16:52:07 +08:00
|
|
|
|
|
|
|
|
|
|
# 初始化向量索引(延迟加载)
|
|
|
|
|
|
self._vector_index: Optional[SchemaIndexer] = None
|
|
|
|
|
|
self._vector_db_path = vector_db_path
|
|
|
|
|
|
|
|
|
|
|
|
# Few-shot 初始化
|
|
|
|
|
|
self.fewshot_enabled = fewshot_enabled
|
|
|
|
|
|
self.fewshot_top_k = fewshot_top_k
|
|
|
|
|
|
self.fewshot_min_rating = fewshot_min_rating
|
|
|
|
|
|
self.fewshot_selector = None
|
|
|
|
|
|
|
|
|
|
|
|
if self.fewshot_enabled:
|
|
|
|
|
|
try:
|
2026-04-14 18:02:12 +08:00
|
|
|
|
use_chroma = os.getenv("FEWSHOT_USE_CHROMA", "false").lower() in (
|
|
|
|
|
|
"1",
|
|
|
|
|
|
"true",
|
|
|
|
|
|
"yes",
|
2026-04-10 16:52:07 +08:00
|
|
|
|
)
|
2026-04-14 18:02:12 +08:00
|
|
|
|
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"
|
2026-04-15 09:49:18 +08:00
|
|
|
|
self.fewshot_selector = FewShotSelector(path or None)
|
2026-04-10 16:52:07 +08:00
|
|
|
|
logger.info(
|
|
|
|
|
|
f"Few-shot已启用: top_k={fewshot_top_k}, "
|
|
|
|
|
|
f"min_rating={fewshot_min_rating}"
|
|
|
|
|
|
)
|
|
|
|
|
|
except Exception as e:
|
|
|
|
|
|
logger.warning(f"Few-shot加载失败: {e},将使用标准生成")
|
|
|
|
|
|
self.fewshot_enabled = False
|
|
|
|
|
|
|
2026-04-16 13:48:44 +08:00
|
|
|
|
# 若本轮走「Chroma 黄金 SQL 条件适配」,在 metadata 中回传 qid/分数
|
|
|
|
|
|
self._last_fewshot_golden: Optional[Tuple[str, float]] = None
|
|
|
|
|
|
|
2026-04-10 16:52:07 +08:00
|
|
|
|
logger.info(
|
|
|
|
|
|
f"[OK] Text2SQLOrchestrator初始化完成: "
|
|
|
|
|
|
f"max_retry={max_retry}, use_vector_search={use_vector_search}"
|
|
|
|
|
|
+ (f", fewshot=on" if self.fewshot_enabled else "")
|
2026-04-15 11:25:19 +08:00
|
|
|
|
+ (", nl→zh_norm=on" if self.translate_english_to_zh else ", nl→zh_norm=off")
|
2026-04-10 16:52:07 +08:00
|
|
|
|
)
|
|
|
|
|
|
|
2026-04-14 18:02:12 +08:00
|
|
|
|
@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
|
|
|
|
|
|
|
2026-04-10 16:52:07 +08:00
|
|
|
|
def _get_vector_index(self) -> SchemaIndexer:
|
|
|
|
|
|
"""获取或创建向量索引(懒加载)"""
|
|
|
|
|
|
if self._vector_index is None:
|
|
|
|
|
|
from utils.embedding import get_embedder
|
|
|
|
|
|
|
2026-04-15 09:49:18 +08:00
|
|
|
|
embedder = get_embedder()
|
2026-04-10 16:52:07 +08:00
|
|
|
|
self._vector_index = SchemaIndexer(
|
|
|
|
|
|
embedder=embedder,
|
|
|
|
|
|
persist_dir=self._vector_db_path
|
|
|
|
|
|
)
|
|
|
|
|
|
return self._vector_index
|
|
|
|
|
|
|
|
|
|
|
|
def _coarse_filter(
|
|
|
|
|
|
self,
|
|
|
|
|
|
question: str,
|
|
|
|
|
|
top_k: int = 20
|
|
|
|
|
|
) -> List[str]:
|
|
|
|
|
|
"""
|
|
|
|
|
|
阶段1:粗筛(向量检索)
|
|
|
|
|
|
|
|
|
|
|
|
Args:
|
|
|
|
|
|
question: 用户问题
|
|
|
|
|
|
top_k: 返回前K个候选表
|
|
|
|
|
|
|
|
|
|
|
|
Returns:
|
|
|
|
|
|
候选表名列表
|
|
|
|
|
|
"""
|
|
|
|
|
|
if not self.use_vector_search:
|
|
|
|
|
|
# 不使用向量检索时,返回所有表
|
2026-04-14 18:21:50 +08:00
|
|
|
|
logger.info("[Orchestrator] 向量搜索已禁用,使用所有表")
|
2026-04-10 16:52:07 +08:00
|
|
|
|
return self.schema_manager.list_tables()
|
|
|
|
|
|
|
2026-04-14 18:21:50 +08:00
|
|
|
|
try:
|
|
|
|
|
|
indexer = self._get_vector_index()
|
2026-04-15 11:25:19 +08:00
|
|
|
|
indexer.ensure_index_for_schema(self.schema_manager)
|
2026-04-10 16:52:07 +08:00
|
|
|
|
|
2026-04-14 18:21:50 +08:00
|
|
|
|
# 检索
|
2026-04-16 10:53:10 +08:00
|
|
|
|
logger.info(
|
|
|
|
|
|
"[Orchestrator] 开始向量检索: query_chars=%s query_preview=%r",
|
|
|
|
|
|
len(question or ""),
|
|
|
|
|
|
(question or "")[:200] + ("…" if len(question or "") > 200 else ""),
|
|
|
|
|
|
)
|
2026-04-14 18:21:50 +08:00
|
|
|
|
results = indexer.search(
|
|
|
|
|
|
query=question,
|
|
|
|
|
|
top_k=top_k,
|
|
|
|
|
|
score_threshold=0.1 # 降低阈值以提高召回率
|
|
|
|
|
|
)
|
2026-04-10 16:52:07 +08:00
|
|
|
|
|
2026-04-14 18:21:50 +08:00
|
|
|
|
candidate_tables = [r["table_name"] for r in results]
|
2026-04-16 10:53:10 +08:00
|
|
|
|
scored = [
|
|
|
|
|
|
(r["table_name"], round(float(r.get("score", 0.0)), 4))
|
|
|
|
|
|
for r in results[: min(25, len(results))]
|
|
|
|
|
|
]
|
|
|
|
|
|
logger.info(
|
|
|
|
|
|
"[Orchestrator] 向量粗筛: 命中=%s 张(阈值内),表名+分: %s",
|
|
|
|
|
|
len(candidate_tables),
|
|
|
|
|
|
scored,
|
|
|
|
|
|
)
|
2026-04-14 18:21:50 +08:00
|
|
|
|
return candidate_tables
|
|
|
|
|
|
|
|
|
|
|
|
except Exception as e:
|
|
|
|
|
|
logger.error(f"[Orchestrator] 向量检索失败: {e},降级为使用所有表", exc_info=True)
|
|
|
|
|
|
import traceback
|
|
|
|
|
|
logger.error(traceback.format_exc())
|
|
|
|
|
|
# 降级:返回所有表
|
|
|
|
|
|
return self.schema_manager.list_tables()
|
2026-04-10 16:52:07 +08:00
|
|
|
|
|
|
|
|
|
|
def _llm_select_tables(
|
|
|
|
|
|
self,
|
|
|
|
|
|
question: str,
|
|
|
|
|
|
candidate_tables: List[str],
|
2026-04-14 10:28:22 +08:00
|
|
|
|
max_tables: int = 5,
|
2026-04-10 16:52:07 +08:00
|
|
|
|
) -> Tuple[List[str], str]:
|
|
|
|
|
|
"""
|
|
|
|
|
|
阶段2:LLM精筛(Schema Linker Agent)
|
|
|
|
|
|
|
|
|
|
|
|
Args:
|
|
|
|
|
|
question: 用户问题
|
|
|
|
|
|
candidate_tables: 候选表列表
|
|
|
|
|
|
max_tables: 最多选择的表数
|
|
|
|
|
|
|
|
|
|
|
|
Returns:
|
|
|
|
|
|
(相关表列表, 推理理由)
|
|
|
|
|
|
"""
|
|
|
|
|
|
# 构造候选表信息(只显示表名和注释)
|
|
|
|
|
|
table_infos = []
|
|
|
|
|
|
for tbl_name in candidate_tables:
|
|
|
|
|
|
table = self.schema_manager.get_table(tbl_name)
|
|
|
|
|
|
if table:
|
|
|
|
|
|
comment = table.comment or "无描述"
|
|
|
|
|
|
table_infos.append(f"- {tbl_name}: {comment}")
|
|
|
|
|
|
|
|
|
|
|
|
table_list_str = "\n".join(table_infos)
|
|
|
|
|
|
|
|
|
|
|
|
# 使用DeepSeek客户端调用(而非CAMEL Agent,更直接可控)
|
|
|
|
|
|
response = self.deepseek.select_tables(
|
|
|
|
|
|
question=question,
|
2026-04-14 10:28:22 +08:00
|
|
|
|
table_list=table_list_str,
|
2026-04-10 16:52:07 +08:00
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
relevant_tables = response.get("relevant_tables", [])
|
|
|
|
|
|
reasoning = response.get("reasoning", "")
|
|
|
|
|
|
|
|
|
|
|
|
# 限制数量
|
|
|
|
|
|
relevant_tables = relevant_tables[:max_tables]
|
|
|
|
|
|
|
2026-04-16 10:53:10 +08:00
|
|
|
|
rs = (reasoning or "").strip()
|
|
|
|
|
|
logger.info(
|
|
|
|
|
|
"LLM精筛选中表:%s | reasoning_chars=%s reasoning_preview=%r",
|
|
|
|
|
|
relevant_tables,
|
|
|
|
|
|
len(rs),
|
|
|
|
|
|
rs[:600] + ("…" if len(rs) > 600 else ""),
|
|
|
|
|
|
)
|
2026-04-10 16:52:07 +08:00
|
|
|
|
return relevant_tables, reasoning
|
|
|
|
|
|
|
|
|
|
|
|
_BROKER_KEYWORDS_CN = ("对手方", "经纪商", "券商", "對手方")
|
|
|
|
|
|
|
|
|
|
|
|
def _question_implies_broker_dimension(self, question: str) -> bool:
|
|
|
|
|
|
if not question:
|
|
|
|
|
|
return False
|
|
|
|
|
|
if any(k in question for k in self._BROKER_KEYWORDS_CN):
|
|
|
|
|
|
return True
|
|
|
|
|
|
return "broker" in question.lower()
|
|
|
|
|
|
|
|
|
|
|
|
def _prioritize_broker_tables(
|
|
|
|
|
|
self, question: str, relevant_tables: List[str], max_tables: int = 5
|
|
|
|
|
|
) -> List[str]:
|
|
|
|
|
|
"""
|
|
|
|
|
|
问题涉及对手方/经纪商时,优先纳入 TSBBrokerContract 与 MCBroker(若 Schema 中存在),
|
|
|
|
|
|
避免仅选中 VSBHK 报表视图却无 BrokerID,模型又照抄黄金范例列名导致校验失败。
|
|
|
|
|
|
"""
|
|
|
|
|
|
if not self._question_implies_broker_dimension(question):
|
|
|
|
|
|
return relevant_tables[:max_tables]
|
|
|
|
|
|
|
|
|
|
|
|
priority = ["TSBBrokerContract", "MCBroker"]
|
|
|
|
|
|
present = [t for t in priority if self.schema_manager.get_table(t)]
|
|
|
|
|
|
if not present:
|
|
|
|
|
|
return relevant_tables[:max_tables]
|
|
|
|
|
|
|
|
|
|
|
|
seen = set()
|
|
|
|
|
|
merged: List[str] = []
|
|
|
|
|
|
for t in present:
|
|
|
|
|
|
if t not in seen:
|
|
|
|
|
|
merged.append(t)
|
|
|
|
|
|
seen.add(t)
|
|
|
|
|
|
for t in relevant_tables:
|
|
|
|
|
|
if len(merged) >= max_tables:
|
|
|
|
|
|
break
|
|
|
|
|
|
if t not in seen and self.schema_manager.get_table(t):
|
|
|
|
|
|
merged.append(t)
|
|
|
|
|
|
seen.add(t)
|
|
|
|
|
|
logger.info("对手方/经纪商问题:优先纳入 %s,调整后选表:%s", present, merged)
|
|
|
|
|
|
return merged[:max_tables]
|
|
|
|
|
|
|
2026-04-17 15:15:01 +08:00
|
|
|
|
_VC_USER_ACCESSIBLE_FUNCTION = "VCUserAccessibleFunction"
|
|
|
|
|
|
|
|
|
|
|
|
def _prioritize_vc_user_accessible_function(
|
|
|
|
|
|
self, relevant_tables: List[str], max_tables: int = 5
|
|
|
|
|
|
) -> List[str]:
|
|
|
|
|
|
"""
|
|
|
|
|
|
若 Schema 中存在 VCUserAccessibleFunction,则置于选表列表最前,便于模型先根据
|
|
|
|
|
|
FunctionID / DatabaseView 等列定位业务视图,再关联其余表生成 SQL。
|
|
|
|
|
|
"""
|
|
|
|
|
|
vc = self._VC_USER_ACCESSIBLE_FUNCTION
|
|
|
|
|
|
if not self.schema_manager.get_table(vc):
|
|
|
|
|
|
return relevant_tables[:max_tables]
|
|
|
|
|
|
|
|
|
|
|
|
rest = [t for t in relevant_tables if t != vc]
|
|
|
|
|
|
merged = [vc] + rest
|
|
|
|
|
|
logger.info(
|
|
|
|
|
|
"已优先纳入目录视图 %s(置于选表前列),当前选表:%s",
|
|
|
|
|
|
vc,
|
|
|
|
|
|
merged[:max_tables],
|
|
|
|
|
|
)
|
|
|
|
|
|
return merged[:max_tables]
|
|
|
|
|
|
|
2026-04-10 16:52:07 +08:00
|
|
|
|
def _expand_relations(self, table_names: List[str]) -> List[str]:
|
|
|
|
|
|
"""
|
|
|
|
|
|
外键扩展:自动添加关联表
|
|
|
|
|
|
|
|
|
|
|
|
Args:
|
|
|
|
|
|
table_names: 已选中的表名列表
|
|
|
|
|
|
|
|
|
|
|
|
Returns:
|
|
|
|
|
|
扩展后的表名列表
|
|
|
|
|
|
"""
|
|
|
|
|
|
result = set(table_names)
|
|
|
|
|
|
|
|
|
|
|
|
for tbl_name in table_names:
|
|
|
|
|
|
table = self.schema_manager.get_table(tbl_name)
|
|
|
|
|
|
if not table:
|
|
|
|
|
|
continue
|
|
|
|
|
|
|
|
|
|
|
|
# 添加被引用的表(外键指向的表)
|
|
|
|
|
|
for fk in table.foreign_keys:
|
|
|
|
|
|
if fk.ref_table not in result:
|
|
|
|
|
|
result.add(fk.ref_table)
|
|
|
|
|
|
logger.debug(f"外键扩展:添加关联表 {fk.ref_table}")
|
|
|
|
|
|
|
|
|
|
|
|
# 添加引用当前表的表(反向外键)
|
|
|
|
|
|
for other in self.schema_manager.get_tables():
|
|
|
|
|
|
for fk in other.foreign_keys:
|
|
|
|
|
|
if fk.ref_table == tbl_name and other.name not in result:
|
|
|
|
|
|
result.add(other.name)
|
|
|
|
|
|
logger.debug(f"外键扩展:添加引用表 {other.name}")
|
|
|
|
|
|
|
|
|
|
|
|
expanded = list(result)
|
|
|
|
|
|
if len(expanded) > len(table_names):
|
|
|
|
|
|
logger.info(f"外键扩展:{table_names} → {expanded}")
|
|
|
|
|
|
|
|
|
|
|
|
return expanded
|
|
|
|
|
|
|
2026-04-16 13:48:44 +08:00
|
|
|
|
def _sql_chat_completion_text(
|
|
|
|
|
|
self,
|
|
|
|
|
|
messages: List[Dict[str, str]],
|
|
|
|
|
|
*,
|
|
|
|
|
|
sql_stream_callback: Optional[Callable[[str], None]] = None,
|
|
|
|
|
|
) -> str:
|
|
|
|
|
|
"""
|
|
|
|
|
|
SQL 生成相关的一次 LLM 调用:可选流式回调(将原始 completion 文本分片回传)。
|
|
|
|
|
|
无回调或非流式失败时与非流式 chat 行为一致。
|
|
|
|
|
|
"""
|
|
|
|
|
|
ds = self.deepseek
|
|
|
|
|
|
if sql_stream_callback is not None and hasattr(ds, "chat_stream"):
|
|
|
|
|
|
try:
|
|
|
|
|
|
parts: List[str] = []
|
|
|
|
|
|
for piece in ds.chat_stream(messages, temperature=0.0, top_p=1.0):
|
|
|
|
|
|
if piece:
|
|
|
|
|
|
parts.append(piece)
|
|
|
|
|
|
sql_stream_callback(piece)
|
|
|
|
|
|
return "".join(parts).strip()
|
|
|
|
|
|
except Exception as e:
|
|
|
|
|
|
logger.warning("[GEN] chat_stream 失败,回退非流式: %s", e)
|
|
|
|
|
|
msg = ds.chat(messages, temperature=0.0, top_p=1.0)
|
|
|
|
|
|
return (msg.content or "").strip()
|
|
|
|
|
|
|
|
|
|
|
|
def _generate_sql_golden_adapt(
|
|
|
|
|
|
self,
|
|
|
|
|
|
question: str,
|
|
|
|
|
|
schema_str: str,
|
|
|
|
|
|
dialect: str,
|
|
|
|
|
|
golden: ExperienceSample,
|
|
|
|
|
|
golden_score: float,
|
|
|
|
|
|
validation_feedback: Optional[str],
|
|
|
|
|
|
dialog_context: Optional[str],
|
|
|
|
|
|
sql_stream_callback: Optional[Callable[[str], None]] = None,
|
|
|
|
|
|
) -> str:
|
|
|
|
|
|
"""
|
|
|
|
|
|
Chroma few-shot 库内 SQL 视为正确答案:在高分相似命中下,仅让模型调整条件/字面量以匹配当前问题。
|
|
|
|
|
|
"""
|
|
|
|
|
|
from config.prompts import GOLDEN_SQL_ADAPT_SYSTEM, GOLDEN_SQL_ADAPT_USER
|
|
|
|
|
|
from utils.sql_parser import normalize_sql_for_dialect
|
|
|
|
|
|
|
|
|
|
|
|
gsql = (golden.sql or "").strip()
|
|
|
|
|
|
max_sql = int(os.getenv("FEWSHOT_GOLDEN_SQL_PROMPT_MAX", "16000"))
|
|
|
|
|
|
if len(gsql) > max_sql:
|
|
|
|
|
|
gsql = gsql[:max_sql] + "\n-- …(标准答案过长,已截断)"
|
|
|
|
|
|
|
|
|
|
|
|
dialect_label = dialect
|
|
|
|
|
|
if dialect == "tsql":
|
|
|
|
|
|
dialect_label = "Microsoft SQL Server (T-SQL)"
|
|
|
|
|
|
|
|
|
|
|
|
dc = (dialog_context or "").strip()
|
|
|
|
|
|
prefix = ""
|
|
|
|
|
|
if dc:
|
|
|
|
|
|
prefix = (
|
|
|
|
|
|
"【对话上文】(用于理解指代与续问条件;请结合「当前用户问题」调整 WHERE 等。)\n"
|
|
|
|
|
|
f"{dc}\n\n"
|
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
user_content = prefix + GOLDEN_SQL_ADAPT_USER.format(
|
|
|
|
|
|
golden_score=golden_score,
|
|
|
|
|
|
golden_question=(golden.question_zh or "").strip(),
|
|
|
|
|
|
golden_sql=gsql,
|
|
|
|
|
|
schema=schema_str,
|
|
|
|
|
|
question=question,
|
|
|
|
|
|
dialect=dialect_label,
|
|
|
|
|
|
)
|
|
|
|
|
|
if dialect == "tsql":
|
|
|
|
|
|
user_content += (
|
|
|
|
|
|
"\n\n【硬性要求】目标库为 SQL Server(T-SQL):禁止使用 MySQL 反引号 `;"
|
|
|
|
|
|
"标识符如需引用请使用方括号,例如 [TableName]、[ColumnName]。"
|
|
|
|
|
|
"字符串连接使用 `+`(与系统提示中的标准版式范例一致)。"
|
|
|
|
|
|
"「今日」「当天」等与日期列比较时,使用 `CAST(GETDATE() AS DATE)`,"
|
|
|
|
|
|
"**禁止** `CURDATE()`、`NOW()`、`CURRENT_DATE`(MySQL)。"
|
|
|
|
|
|
"条件请使用 T-SQL 惯用写法(例如 IS NOT NULL)。"
|
|
|
|
|
|
"排版仍须遵守:关键字大写、SELECT 每列一行缩进、WHERE 续行以 AND 开头、PascalCase 英文别名。"
|
2026-04-17 15:15:01 +08:00
|
|
|
|
"\n在 `WHERE`/`HAVING`/`JOIN ON` 及 `CASE/WHEN` 的**条件**中,**禁止**用含中日韩的 `'`/`N'…'` 字面量与代码型列比对;"
|
|
|
|
|
|
"业务中文须映射为 Schema 注释中的代码或通过维表 JOIN(勿写 `= '过户费'`)。`SELECT` 或 `CASE … THEN/ELSE` 的展示用中文标签允许。"
|
|
|
|
|
|
"若用户问题与对话上文均未要求按时间筛选,不得在 WHERE 中擅自添加日期列条件(与系统提示 1b 一致)。"
|
2026-04-16 13:48:44 +08:00
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
if validation_feedback:
|
|
|
|
|
|
user_content += (
|
|
|
|
|
|
"\n\n【上次校验未通过】请在保留标准答案主干的前提下修正 SQL;"
|
|
|
|
|
|
"表名、列名必须与「当前 Schema 片段」中完全一致。\n"
|
|
|
|
|
|
f"{validation_feedback}"
|
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
messages = [
|
|
|
|
|
|
{"role": "system", "content": GOLDEN_SQL_ADAPT_SYSTEM},
|
|
|
|
|
|
{"role": "user", "content": user_content},
|
|
|
|
|
|
]
|
|
|
|
|
|
|
|
|
|
|
|
logger.info(
|
|
|
|
|
|
"[GEN] 黄金 few-shot 条件适配: qid=%s score=%.4f",
|
|
|
|
|
|
golden.qid,
|
|
|
|
|
|
golden_score,
|
|
|
|
|
|
)
|
|
|
|
|
|
sql = self._sql_chat_completion_text(
|
|
|
|
|
|
messages, sql_stream_callback=sql_stream_callback
|
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
if "```sql" in sql:
|
|
|
|
|
|
sql = sql[sql.find("```sql") + 6 : sql.find("```", sql.find("```sql") + 6)].strip()
|
|
|
|
|
|
elif "```" in sql:
|
|
|
|
|
|
sql = sql[sql.find("```") + 3 : sql.find("```", sql.find("```") + 3)].strip()
|
|
|
|
|
|
|
|
|
|
|
|
sql = normalize_sql_for_dialect(sql, dialect)
|
|
|
|
|
|
|
|
|
|
|
|
lim = 12000
|
|
|
|
|
|
body = sql if len(sql) <= lim else sql[:lim] + "\n…(日志已截断)"
|
|
|
|
|
|
logger.info("生成的SQL(黄金适配,chars=%s):\n%s", len(sql), body)
|
|
|
|
|
|
return sql
|
|
|
|
|
|
|
2026-04-10 16:52:07 +08:00
|
|
|
|
def _generate_sql(
|
|
|
|
|
|
self,
|
|
|
|
|
|
question: str,
|
|
|
|
|
|
schema_str: str,
|
2026-04-14 10:28:22 +08:00
|
|
|
|
dialect: str = "tsql",
|
|
|
|
|
|
validation_feedback: Optional[str] = None,
|
2026-04-14 18:02:12 +08:00
|
|
|
|
dialog_context: Optional[str] = None,
|
2026-04-16 13:48:44 +08:00
|
|
|
|
sql_stream_callback: Optional[Callable[[str], None]] = None,
|
2026-04-10 16:52:07 +08:00
|
|
|
|
) -> str:
|
|
|
|
|
|
"""
|
|
|
|
|
|
SQL生成(SQL Generator Agent)
|
|
|
|
|
|
|
|
|
|
|
|
Args:
|
|
|
|
|
|
question: 用户问题
|
|
|
|
|
|
schema_str: Schema描述字符串
|
|
|
|
|
|
dialect: SQL方言
|
2026-04-14 10:28:22 +08:00
|
|
|
|
validation_feedback: 非空时附加到用户提示(重试时传入上次校验错误)
|
2026-04-10 16:52:07 +08:00
|
|
|
|
|
2026-04-14 18:02:12 +08:00
|
|
|
|
dialog_context: 前几轮对话摘要;与 ``question`` 一并供指代消解与续问。
|
2026-04-16 13:48:44 +08:00
|
|
|
|
sql_stream_callback: 若提供且 LLM 支持 chat_stream,则在 SQL 主生成/黄金适配时流式回传原始文本分片。
|
2026-04-14 18:02:12 +08:00
|
|
|
|
|
2026-04-10 16:52:07 +08:00
|
|
|
|
Returns:
|
|
|
|
|
|
SQL语句
|
|
|
|
|
|
"""
|
|
|
|
|
|
from config.prompts import SQL_GENERATOR_SYSTEM, SQL_GENERATOR_USER
|
|
|
|
|
|
from utils.sql_parser import normalize_sql_for_dialect
|
|
|
|
|
|
|
2026-04-16 13:48:44 +08:00
|
|
|
|
self._last_fewshot_golden = None
|
2026-04-14 18:02:12 +08:00
|
|
|
|
dc = (dialog_context or "").strip()
|
|
|
|
|
|
fewshot_question = question
|
|
|
|
|
|
if dc:
|
|
|
|
|
|
fewshot_question = f"{dc}\n\n【当前问】{question}"
|
|
|
|
|
|
|
2026-04-16 13:48:44 +08:00
|
|
|
|
golden_reuse = os.getenv("FEWSHOT_GOLDEN_REUSE", "true").lower() in (
|
|
|
|
|
|
"1",
|
|
|
|
|
|
"true",
|
|
|
|
|
|
"yes",
|
|
|
|
|
|
)
|
|
|
|
|
|
only_chroma = os.getenv("FEWSHOT_GOLDEN_ONLY_CHROMA", "true").lower() in (
|
|
|
|
|
|
"1",
|
|
|
|
|
|
"true",
|
|
|
|
|
|
"yes",
|
|
|
|
|
|
)
|
|
|
|
|
|
golden_min = float(os.getenv("FEWSHOT_GOLDEN_MIN_SCORE", "0.88"))
|
|
|
|
|
|
chroma_ok = bool(
|
|
|
|
|
|
self.fewshot_selector and self.fewshot_selector.is_chroma_backend
|
|
|
|
|
|
)
|
|
|
|
|
|
if only_chroma and not chroma_ok:
|
|
|
|
|
|
golden_reuse = False
|
|
|
|
|
|
|
|
|
|
|
|
if (
|
|
|
|
|
|
golden_reuse
|
|
|
|
|
|
and self.fewshot_enabled
|
|
|
|
|
|
and self.fewshot_selector
|
|
|
|
|
|
):
|
|
|
|
|
|
best = self.fewshot_selector.select_best_with_score(
|
|
|
|
|
|
fewshot_question,
|
|
|
|
|
|
min_rating=self.fewshot_min_rating,
|
|
|
|
|
|
)
|
|
|
|
|
|
if (
|
|
|
|
|
|
best
|
|
|
|
|
|
and best[1] >= golden_min
|
|
|
|
|
|
and (best[0].sql or "").strip()
|
|
|
|
|
|
):
|
|
|
|
|
|
ex, sc = best
|
|
|
|
|
|
sql_out = self._generate_sql_golden_adapt(
|
|
|
|
|
|
question=question,
|
|
|
|
|
|
schema_str=schema_str,
|
|
|
|
|
|
dialect=dialect,
|
|
|
|
|
|
golden=ex,
|
|
|
|
|
|
golden_score=sc,
|
|
|
|
|
|
validation_feedback=validation_feedback,
|
|
|
|
|
|
dialog_context=dialog_context,
|
|
|
|
|
|
sql_stream_callback=sql_stream_callback,
|
|
|
|
|
|
)
|
|
|
|
|
|
self._last_fewshot_golden = (ex.qid, sc)
|
|
|
|
|
|
return sql_out
|
|
|
|
|
|
|
2026-04-10 16:52:07 +08:00
|
|
|
|
# Few-shot 增强
|
|
|
|
|
|
if self.fewshot_enabled and self.fewshot_selector:
|
|
|
|
|
|
try:
|
|
|
|
|
|
examples = self.fewshot_selector.select(
|
2026-04-14 18:02:12 +08:00
|
|
|
|
question=fewshot_question,
|
2026-04-10 16:52:07 +08:00
|
|
|
|
top_k=self.fewshot_top_k,
|
|
|
|
|
|
min_rating=self.fewshot_min_rating
|
|
|
|
|
|
)
|
|
|
|
|
|
if examples:
|
|
|
|
|
|
examples_prompt = "\n\n".join([
|
|
|
|
|
|
f"示例 {i+1}:\n问题:{ex.question_zh}\nSQL:\n{ex.sql}"
|
|
|
|
|
|
for i, ex in enumerate(examples)
|
|
|
|
|
|
])
|
|
|
|
|
|
schema_str = f"参考以下相似示例的SQL编写风格:\n\n{examples_prompt}\n\n【当前Schema】\n{schema_str}"
|
2026-04-16 10:53:10 +08:00
|
|
|
|
logger.info(
|
|
|
|
|
|
"已注入 %s 个 few-shot 示例: qid=%s question_zh_preview=%r",
|
|
|
|
|
|
len(examples),
|
|
|
|
|
|
[ex.qid for ex in examples],
|
|
|
|
|
|
[((ex.question_zh or "")[:80] + "…") if len(ex.question_zh or "") > 80 else (ex.question_zh or "") for ex in examples],
|
|
|
|
|
|
)
|
2026-04-10 16:52:07 +08:00
|
|
|
|
except Exception as e:
|
|
|
|
|
|
logger.warning(f"Few-shot检索失败: {e}")
|
|
|
|
|
|
|
|
|
|
|
|
dialect_label = dialect
|
|
|
|
|
|
if dialect == "tsql":
|
|
|
|
|
|
dialect_label = "Microsoft SQL Server (T-SQL)"
|
|
|
|
|
|
|
2026-04-14 18:02:12 +08:00
|
|
|
|
prefix = ""
|
|
|
|
|
|
if dc:
|
|
|
|
|
|
prefix = (
|
|
|
|
|
|
"【对话上文】(用于理解「这/那/同样/上面/刚才」等指代及续问条件;"
|
|
|
|
|
|
"请结合下文「当前用户问题」生成 SQL。)\n"
|
|
|
|
|
|
f"{dc}\n\n"
|
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
user_content = prefix + SQL_GENERATOR_USER.format(
|
2026-04-10 16:52:07 +08:00
|
|
|
|
schema=schema_str,
|
|
|
|
|
|
question=question,
|
|
|
|
|
|
dialect=dialect_label,
|
|
|
|
|
|
)
|
|
|
|
|
|
if dialect == "tsql":
|
|
|
|
|
|
user_content += (
|
|
|
|
|
|
"\n\n【硬性要求】目标库为 SQL Server(T-SQL):禁止使用 MySQL 反引号 `;"
|
|
|
|
|
|
"标识符如需引用请使用方括号,例如 [TableName]、[ColumnName]。"
|
|
|
|
|
|
"字符串连接使用 `+`(与系统提示中的标准版式范例一致)。"
|
|
|
|
|
|
"「今日」「当天」等与日期列比较时,使用 `CAST(GETDATE() AS DATE)`,"
|
|
|
|
|
|
"**禁止** `CURDATE()`、`NOW()`、`CURRENT_DATE`(MySQL)。"
|
|
|
|
|
|
"条件请使用 T-SQL 惯用写法(例如 IS NOT NULL)。"
|
|
|
|
|
|
"排版仍须遵守:关键字大写、SELECT 每列一行缩进、WHERE 续行以 AND 开头、PascalCase 英文别名。"
|
2026-04-17 15:15:01 +08:00
|
|
|
|
"\n在 `WHERE`/`HAVING`/`JOIN ON` 及 `CASE/WHEN` 的**条件**中,**禁止**用含中日韩的 `'`/`N'…'` 字面量与代码型列比对;"
|
|
|
|
|
|
"业务中文须映射为 Schema 注释中的代码或通过维表 JOIN(勿写 `= '过户费'`)。`SELECT` 或 `CASE … THEN/ELSE` 的展示用中文标签允许。"
|
|
|
|
|
|
"若用户问题与对话上文均未要求按时间筛选,不得在 WHERE 中擅自添加日期列条件(与系统提示 1b 一致)。"
|
2026-04-14 10:28:22 +08:00
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
if validation_feedback:
|
|
|
|
|
|
user_content += (
|
|
|
|
|
|
"\n\n【上次校验未通过】请根据下列错误修正 SQL,并输出完整可执行查询;"
|
|
|
|
|
|
"表名、列名必须与「当前Schema」中完全一致,不要臆造字段名。\n"
|
|
|
|
|
|
f"{validation_feedback}"
|
2026-04-10 16:52:07 +08:00
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
messages = [
|
|
|
|
|
|
{"role": "system", "content": SQL_GENERATOR_SYSTEM},
|
|
|
|
|
|
{"role": "user", "content": user_content},
|
|
|
|
|
|
]
|
|
|
|
|
|
|
2026-04-14 10:28:22 +08:00
|
|
|
|
# 与选表一致:生成阶段默认贪心解码,减少同一中文问题多次 SQL 不一致
|
2026-04-16 13:48:44 +08:00
|
|
|
|
sql = self._sql_chat_completion_text(
|
|
|
|
|
|
messages, sql_stream_callback=sql_stream_callback
|
|
|
|
|
|
)
|
2026-04-10 16:52:07 +08:00
|
|
|
|
|
|
|
|
|
|
# 清理可能的markdown代码块
|
|
|
|
|
|
if "```sql" in sql:
|
|
|
|
|
|
sql = sql[sql.find("```sql") + 6:sql.find("```", sql.find("```sql") + 6)].strip()
|
|
|
|
|
|
elif "```" in sql:
|
|
|
|
|
|
sql = sql[sql.find("```") + 3:sql.find("```", sql.find("```") + 3)].strip()
|
|
|
|
|
|
|
|
|
|
|
|
sql = normalize_sql_for_dialect(sql, dialect)
|
|
|
|
|
|
|
2026-04-16 10:53:10 +08:00
|
|
|
|
lim = 12000
|
|
|
|
|
|
body = sql if len(sql) <= lim else sql[:lim] + "\n…(日志已截断)"
|
|
|
|
|
|
logger.info("生成的SQL(chars=%s):\n%s", len(sql), body)
|
2026-04-10 16:52:07 +08:00
|
|
|
|
return sql
|
|
|
|
|
|
|
|
|
|
|
|
def _validate_sql(
|
|
|
|
|
|
self,
|
|
|
|
|
|
sql: str,
|
|
|
|
|
|
schema_str: str,
|
|
|
|
|
|
dialect: str = "tsql",
|
2026-04-14 10:28:22 +08:00
|
|
|
|
question: str = "",
|
2026-04-14 18:02:12 +08:00
|
|
|
|
dialog_context: Optional[str] = None,
|
2026-04-14 10:28:22 +08:00
|
|
|
|
) -> Tuple[bool, List[str], List[str], Optional[int], Optional[str]]:
|
2026-04-10 16:52:07 +08:00
|
|
|
|
"""
|
2026-04-14 10:28:22 +08:00
|
|
|
|
SQL验证(程序验证 + 库探针 + 按探针分支的 LLM)
|
2026-04-10 16:52:07 +08:00
|
|
|
|
|
|
|
|
|
|
Args:
|
|
|
|
|
|
sql: SQL语句
|
|
|
|
|
|
schema_str: Schema描述
|
|
|
|
|
|
dialect: 与生成一致的 SQL 方言(sqlglot 名,默认 tsql)
|
2026-04-14 10:28:22 +08:00
|
|
|
|
question: 用户自然语言(探针为 0 时用于生成补充说明)
|
2026-04-14 18:02:12 +08:00
|
|
|
|
dialog_context: 会话上文;探针 0 时与 question 一并传入说明模型
|
2026-04-10 16:52:07 +08:00
|
|
|
|
|
|
|
|
|
|
Returns:
|
2026-04-14 10:28:22 +08:00
|
|
|
|
(是否通过, 错误列表, 警告列表, 库执行探针状态, 无数据时的用户说明)
|
|
|
|
|
|
探针:``1`` 至少一行数据;``0`` 执行成功但行数为 0;``-1`` 执行失败;``None`` 未配置库或未跑探针
|
2026-04-10 16:52:07 +08:00
|
|
|
|
"""
|
|
|
|
|
|
errors = []
|
|
|
|
|
|
warnings = []
|
2026-04-14 10:28:22 +08:00
|
|
|
|
db_execution_status: Optional[int] = None
|
|
|
|
|
|
empty_feedback: Optional[str] = None
|
2026-04-10 16:52:07 +08:00
|
|
|
|
|
|
|
|
|
|
# === 阶段1:程序验证(确定性规则) ===
|
|
|
|
|
|
from utils.sql_parser import validate_sql_syntax, validate_schema_consistency
|
|
|
|
|
|
|
|
|
|
|
|
# 语法验证
|
|
|
|
|
|
syntax_ok, syntax_errors = validate_sql_syntax(sql, dialect=dialect)
|
|
|
|
|
|
if not syntax_ok:
|
|
|
|
|
|
errors.extend(syntax_errors)
|
|
|
|
|
|
|
|
|
|
|
|
# Schema一致性验证
|
|
|
|
|
|
schema_ok, schema_errors = validate_schema_consistency(
|
|
|
|
|
|
sql, self.schema_manager, dialect=dialect
|
|
|
|
|
|
)
|
|
|
|
|
|
if not schema_ok:
|
|
|
|
|
|
errors.extend(schema_errors)
|
|
|
|
|
|
|
|
|
|
|
|
# 危险操作检查
|
2026-04-14 10:28:22 +08:00
|
|
|
|
from utils.validators import check_dangerous_operations, check_no_cjk_in_sql_string_literals
|
2026-04-10 16:52:07 +08:00
|
|
|
|
danger_ok, danger_errors = check_dangerous_operations(sql)
|
|
|
|
|
|
if not danger_ok:
|
|
|
|
|
|
errors.extend(danger_errors)
|
|
|
|
|
|
|
2026-04-17 15:15:01 +08:00
|
|
|
|
# T-SQL:禁止中文出现在 WHERE/HAVING/ON/CASE 条件等比对语境(展示用 CASE THEN/ELSE 允许)
|
2026-04-14 10:28:22 +08:00
|
|
|
|
if dialect == "tsql":
|
2026-04-17 15:15:01 +08:00
|
|
|
|
cjk_ok, cjk_errors = check_no_cjk_in_sql_string_literals(sql, dialect=dialect)
|
2026-04-14 10:28:22 +08:00
|
|
|
|
if not cjk_ok:
|
|
|
|
|
|
errors.extend(cjk_errors)
|
2026-04-10 16:52:07 +08:00
|
|
|
|
|
2026-04-14 10:28:22 +08:00
|
|
|
|
# === 阶段1.5:数据库试执行(仅程序校验全部通过时;需配置 database_url) ===
|
|
|
|
|
|
if len(errors) == 0:
|
|
|
|
|
|
from db.dbhub_tools import probe_sql_execution_status_ex
|
2026-04-10 16:52:07 +08:00
|
|
|
|
|
2026-04-14 10:28:22 +08:00
|
|
|
|
db_execution_status, db_probe_err = probe_sql_execution_status_ex(sql)
|
|
|
|
|
|
if db_execution_status == -1:
|
|
|
|
|
|
msg = (
|
2026-04-14 18:02:12 +08:00
|
|
|
|
"【库探针结果:-1 执行失败】SQL 在目标库执行报错,"
|
|
|
|
|
|
"系统将依据下列错误**自动重新生成** SQL(请等待重试结果)。"
|
2026-04-14 10:28:22 +08:00
|
|
|
|
)
|
|
|
|
|
|
if db_probe_err:
|
2026-04-14 18:02:12 +08:00
|
|
|
|
msg += f"\n数据库返回:{db_probe_err}"
|
2026-04-14 10:28:22 +08:00
|
|
|
|
errors.append(msg)
|
|
|
|
|
|
elif db_execution_status is None:
|
|
|
|
|
|
warnings.append(
|
|
|
|
|
|
"未配置 database_url,已跳过数据库执行探针"
|
|
|
|
|
|
)
|
2026-04-10 16:52:07 +08:00
|
|
|
|
|
2026-04-14 10:28:22 +08:00
|
|
|
|
# 探针 1:库上至少有一行数据,跳过 Validator LLM,直接将 SQL 视为可交付
|
|
|
|
|
|
# 探针 0:执行成功但行数为 0,跳过 Validator LLM,另调 LLM 生成说明并引导用户补充条件
|
|
|
|
|
|
# 探针 -1:执行失败,跳过 Validator LLM,走重试
|
|
|
|
|
|
# 探针 None:走完整 Validator LLM
|
|
|
|
|
|
skip_validator_llm = db_execution_status in (-1, 0, 1)
|
2026-04-10 16:52:07 +08:00
|
|
|
|
|
2026-04-14 10:28:22 +08:00
|
|
|
|
# === 阶段2:LLM 语义验证(仅未命中库探针 0/1/-1 时) ===
|
|
|
|
|
|
if not skip_validator_llm:
|
|
|
|
|
|
try:
|
|
|
|
|
|
llm_result = self.deepseek.validate_sql(sql=sql, schema=schema_str)
|
|
|
|
|
|
|
|
|
|
|
|
llm_errors = list(llm_result.get("errors", []))
|
|
|
|
|
|
# 程序校验已通过表/列(含别名解析)时,LLM 仍常误报 unknown_*,避免误杀整次生成
|
|
|
|
|
|
if schema_ok:
|
|
|
|
|
|
llm_errors = [
|
|
|
|
|
|
e
|
|
|
|
|
|
for e in llm_errors
|
|
|
|
|
|
if isinstance(e, str)
|
|
|
|
|
|
and not (
|
|
|
|
|
|
e.startswith("unknown_table:")
|
|
|
|
|
|
or e.startswith("unknown_column:")
|
|
|
|
|
|
)
|
|
|
|
|
|
]
|
|
|
|
|
|
|
|
|
|
|
|
if not llm_result.get("valid", True):
|
|
|
|
|
|
errors.extend(llm_errors)
|
|
|
|
|
|
|
|
|
|
|
|
warnings.extend(llm_result.get("warnings", []))
|
|
|
|
|
|
suggestions = llm_result.get("suggestions", [])
|
|
|
|
|
|
if suggestions:
|
|
|
|
|
|
logger.debug(f"优化建议:{suggestions}")
|
|
|
|
|
|
|
|
|
|
|
|
except Exception as e:
|
|
|
|
|
|
logger.warning(f"LLM验证失败(降级为仅程序验证): {e}")
|
|
|
|
|
|
|
|
|
|
|
|
# === 阶段2b:探针 0 时生成用户可读补充说明(仍返回 SQL,由 API/CLI 一并展示) ===
|
|
|
|
|
|
if db_execution_status == 0 and len(errors) == 0:
|
|
|
|
|
|
prefix = (
|
2026-04-14 18:02:12 +08:00
|
|
|
|
"【库探针结果:0 行】该 SQL 已在数据库成功执行,但**返回数据行数为 0**(未查到匹配记录)。"
|
|
|
|
|
|
"下方已附带完整 SQL 与原因分析,请一并阅读。"
|
2026-04-14 10:28:22 +08:00
|
|
|
|
)
|
2026-04-14 18:02:12 +08:00
|
|
|
|
fb_q = question
|
|
|
|
|
|
dc = (dialog_context or "").strip()
|
|
|
|
|
|
if dc:
|
|
|
|
|
|
fb_q = f"{dc}\n\n【当前用户问题】\n{question}"
|
2026-04-14 10:28:22 +08:00
|
|
|
|
try:
|
|
|
|
|
|
llm_fb = self.deepseek.empty_result_user_feedback(
|
2026-04-14 18:02:12 +08:00
|
|
|
|
question=fb_q,
|
2026-04-14 10:28:22 +08:00
|
|
|
|
sql=sql,
|
|
|
|
|
|
schema=schema_str,
|
|
|
|
|
|
)
|
2026-04-14 18:02:12 +08:00
|
|
|
|
empty_feedback = f"{prefix}\n\n【问题分析】\n{llm_fb}"
|
2026-04-14 10:28:22 +08:00
|
|
|
|
except Exception as e:
|
|
|
|
|
|
logger.warning(f"无数据说明生成失败: {e}")
|
|
|
|
|
|
empty_feedback = (
|
2026-04-14 18:02:12 +08:00
|
|
|
|
f"{prefix}\n\n【问题分析】\n"
|
2026-04-14 10:28:22 +08:00
|
|
|
|
"未能自动生成详细分析。请补充时间范围、筛选条件或业务对象后重新提问。"
|
|
|
|
|
|
)
|
2026-04-10 16:52:07 +08:00
|
|
|
|
|
2026-04-14 18:02:12 +08:00
|
|
|
|
follow = (
|
|
|
|
|
|
"\n\n【追问 — 请补充后再次提问以重新生成 SQL】\n"
|
|
|
|
|
|
"1. 请根据上述分析,尽量具体地补充或修正:**时间范围**、**业务对象**(账户/合约/代码等)、"
|
|
|
|
|
|
"**筛选口径** 或 **您认为 SQL 中不合理的条件**。\n"
|
|
|
|
|
|
"2. 补充说明后请**重新发起一次自然语言提问**(无需粘贴 SQL),系统会结合您的新描述**重新生成**查询。"
|
|
|
|
|
|
)
|
|
|
|
|
|
empty_feedback = (empty_feedback or prefix) + follow
|
|
|
|
|
|
|
2026-04-10 16:52:07 +08:00
|
|
|
|
is_valid = len(errors) == 0
|
2026-04-16 10:53:10 +08:00
|
|
|
|
logger.info(
|
|
|
|
|
|
"[validate] 程序+探针+LLM 汇总: valid=%s err_count=%s warn_count=%s "
|
|
|
|
|
|
"db_execution_status=%s sql_chars=%s",
|
|
|
|
|
|
is_valid,
|
|
|
|
|
|
len(errors),
|
|
|
|
|
|
len(warnings),
|
|
|
|
|
|
db_execution_status,
|
|
|
|
|
|
len(sql or ""),
|
|
|
|
|
|
)
|
|
|
|
|
|
if errors:
|
|
|
|
|
|
logger.info("[validate] errors 预览: %s", errors[:5])
|
|
|
|
|
|
if warnings:
|
|
|
|
|
|
logger.info("[validate] warnings: %s", warnings[:5])
|
2026-04-14 10:28:22 +08:00
|
|
|
|
return is_valid, errors, warnings, db_execution_status, empty_feedback
|
2026-04-10 16:52:07 +08:00
|
|
|
|
|
|
|
|
|
|
def generate(
|
|
|
|
|
|
self,
|
|
|
|
|
|
question: str,
|
|
|
|
|
|
dialect: str = "tsql",
|
|
|
|
|
|
top_k_candidates: int = 20,
|
2026-04-14 18:02:12 +08:00
|
|
|
|
include_schema_in_result: bool = False,
|
|
|
|
|
|
dialog_context: Optional[str] = None,
|
2026-04-16 13:48:44 +08:00
|
|
|
|
sql_stream_callback: Optional[Callable[[str], None]] = None,
|
2026-04-10 16:52:07 +08:00
|
|
|
|
) -> GenerationResult:
|
|
|
|
|
|
"""
|
|
|
|
|
|
主生成流程
|
|
|
|
|
|
|
|
|
|
|
|
Args:
|
|
|
|
|
|
question: 用户自然语言问题
|
|
|
|
|
|
dialect: SQL方言(默认 tsql;与 sqlglot 一致)
|
|
|
|
|
|
top_k_candidates: 粗筛候选表数量
|
|
|
|
|
|
include_schema_in_result: 结果中是否包含使用的Schema字符串
|
2026-04-14 18:02:12 +08:00
|
|
|
|
dialog_context: 前几轮对话可读摘要;选表、向量粗筛、SQL 生成与无数据说明会参考
|
2026-04-16 13:48:44 +08:00
|
|
|
|
sql_stream_callback: 可选;SQL 主生成 LLM 输出分片回调(用于 API SSE)
|
2026-04-10 16:52:07 +08:00
|
|
|
|
|
|
|
|
|
|
Returns:
|
|
|
|
|
|
GenerationResult对象
|
|
|
|
|
|
"""
|
2026-04-16 13:48:44 +08:00
|
|
|
|
self._last_fewshot_golden = None
|
2026-04-14 10:28:22 +08:00
|
|
|
|
original_question = (question or "").strip()
|
|
|
|
|
|
translation_meta: Dict = {}
|
|
|
|
|
|
work_question = original_question
|
2026-04-15 11:25:19 +08:00
|
|
|
|
if self.translate_english_to_zh and original_question:
|
2026-04-14 10:28:22 +08:00
|
|
|
|
try:
|
2026-04-15 11:25:19 +08:00
|
|
|
|
zh = self.deepseek.normalize_nl_question_for_text2sql(original_question).strip()
|
2026-04-14 10:28:22 +08:00
|
|
|
|
if zh and len(zh) >= 2:
|
|
|
|
|
|
work_question = zh
|
|
|
|
|
|
translation_meta["question_original"] = original_question
|
|
|
|
|
|
translation_meta["question_zh_normalized"] = zh
|
|
|
|
|
|
logger.info(
|
2026-04-15 11:25:19 +08:00
|
|
|
|
"[GEN] 问句已归一中文:%s",
|
2026-04-14 10:28:22 +08:00
|
|
|
|
zh[:120] + ("…" if len(zh) > 120 else ""),
|
|
|
|
|
|
)
|
|
|
|
|
|
else:
|
2026-04-15 11:25:19 +08:00
|
|
|
|
logger.warning("[GEN] 归一结果为空或过短,使用原文")
|
2026-04-14 10:28:22 +08:00
|
|
|
|
except Exception as e:
|
2026-04-15 11:25:19 +08:00
|
|
|
|
logger.warning("[GEN] 问句归一失败,使用原文: %s", e)
|
2026-04-14 10:28:22 +08:00
|
|
|
|
|
|
|
|
|
|
question = work_question
|
2026-04-14 18:02:12 +08:00
|
|
|
|
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)
|
2026-04-16 10:53:10 +08:00
|
|
|
|
logger.info(
|
|
|
|
|
|
"[GEN] 开始生成SQL: question_chars=%s preview=%r dialog_context_chars=%s",
|
|
|
|
|
|
len(question or ""),
|
|
|
|
|
|
(question or "")[:300] + ("…" if len(question or "") > 300 else ""),
|
|
|
|
|
|
len(dc_raw) if dc_raw else 0,
|
|
|
|
|
|
)
|
2026-04-10 16:52:07 +08:00
|
|
|
|
|
|
|
|
|
|
attempt = 0
|
|
|
|
|
|
last_sql = None
|
|
|
|
|
|
last_errors = []
|
2026-04-14 10:28:22 +08:00
|
|
|
|
last_db_execution_status: Optional[int] = None
|
2026-04-10 16:52:07 +08:00
|
|
|
|
filtered_schema_str = ""
|
|
|
|
|
|
tables_used = []
|
|
|
|
|
|
|
|
|
|
|
|
while attempt < self.max_retry:
|
|
|
|
|
|
logger.info(f" 尝试 #{attempt + 1}")
|
|
|
|
|
|
|
|
|
|
|
|
# === Step 1: Schema筛选(仅首次) ===
|
|
|
|
|
|
if attempt == 0:
|
|
|
|
|
|
# 1.1 粗筛
|
2026-04-14 18:02:12 +08:00
|
|
|
|
candidate_tables = self._coarse_filter(
|
|
|
|
|
|
retrieval_question, top_k=top_k_candidates
|
|
|
|
|
|
)
|
2026-04-10 16:52:07 +08:00
|
|
|
|
|
|
|
|
|
|
# 1.2 LLM精筛
|
2026-04-14 10:28:22 +08:00
|
|
|
|
relevant_tables, reasoning = self._llm_select_tables(
|
2026-04-14 18:02:12 +08:00
|
|
|
|
linker_question,
|
2026-04-14 10:28:22 +08:00
|
|
|
|
candidate_tables,
|
|
|
|
|
|
)
|
2026-04-14 18:02:12 +08:00
|
|
|
|
relevant_tables = self._prioritize_broker_tables(
|
|
|
|
|
|
linker_question, relevant_tables
|
|
|
|
|
|
)
|
2026-04-17 15:15:01 +08:00
|
|
|
|
relevant_tables = self._prioritize_vc_user_accessible_function(
|
|
|
|
|
|
relevant_tables
|
|
|
|
|
|
)
|
2026-04-10 16:52:07 +08:00
|
|
|
|
|
|
|
|
|
|
# 1.3 外键扩展
|
|
|
|
|
|
expanded_tables = self._expand_relations(relevant_tables)
|
|
|
|
|
|
tables_used = expanded_tables
|
|
|
|
|
|
|
|
|
|
|
|
# 1.4 生成Schema字符串
|
|
|
|
|
|
filtered_schema_str = self.schema_manager.to_compact_string(
|
|
|
|
|
|
table_names=expanded_tables,
|
|
|
|
|
|
include_columns=True,
|
|
|
|
|
|
max_columns_per_table=20
|
|
|
|
|
|
)
|
2026-04-17 15:15:01 +08:00
|
|
|
|
if self._VC_USER_ACCESSIBLE_FUNCTION in expanded_tables:
|
|
|
|
|
|
filtered_schema_str += (
|
|
|
|
|
|
"\n\n【选表提示】已包含视图 "
|
|
|
|
|
|
+ self._VC_USER_ACCESSIBLE_FUNCTION
|
|
|
|
|
|
+ "(列含 UserID、FunctionID、Category、Name、DatabaseView)。"
|
|
|
|
|
|
"生成 SQL 时可先通过该视图用 DatabaseView / FunctionID 等定位目标业务视图或功能,"
|
|
|
|
|
|
"再与 Schema 中其余表做 JOIN 或子查询;若问题已明确具体表名,可直接查询该表。"
|
|
|
|
|
|
)
|
2026-04-10 16:52:07 +08:00
|
|
|
|
logger.info(f" 选中表:{relevant_tables},扩展后:{expanded_tables}")
|
|
|
|
|
|
else:
|
2026-04-14 10:28:22 +08:00
|
|
|
|
# 重试:在子 Schema 中并入「上次失败 SQL」实际引用到的表,并对齐程序校验与生成上下文
|
2026-04-10 16:52:07 +08:00
|
|
|
|
logger.info(f" 重试使用之前的Schema({len(tables_used)}张表)")
|
2026-04-14 10:28:22 +08:00
|
|
|
|
if last_sql:
|
|
|
|
|
|
from utils.sql_parser import extract_tables_from_sql
|
|
|
|
|
|
|
|
|
|
|
|
extra = [
|
|
|
|
|
|
t
|
|
|
|
|
|
for t in extract_tables_from_sql(last_sql, dialect=dialect)
|
|
|
|
|
|
if self.schema_manager.get_table(t)
|
|
|
|
|
|
]
|
|
|
|
|
|
merged = list(dict.fromkeys([*(tables_used or []), *extra]))
|
|
|
|
|
|
tables_used = self._expand_relations(merged)
|
|
|
|
|
|
filtered_schema_str = self.schema_manager.to_compact_string(
|
|
|
|
|
|
table_names=tables_used,
|
|
|
|
|
|
include_columns=True,
|
|
|
|
|
|
max_columns_per_table=20,
|
|
|
|
|
|
)
|
|
|
|
|
|
if extra:
|
|
|
|
|
|
logger.info(
|
|
|
|
|
|
" 重试:合并失败SQL中的表 %s,外键扩展后:%s",
|
|
|
|
|
|
extra,
|
|
|
|
|
|
tables_used,
|
|
|
|
|
|
)
|
2026-04-10 16:52:07 +08:00
|
|
|
|
|
|
|
|
|
|
# === Step 2: SQL生成 ===
|
|
|
|
|
|
try:
|
2026-04-14 10:28:22 +08:00
|
|
|
|
feedback: Optional[str] = None
|
|
|
|
|
|
if attempt > 0 and last_errors:
|
|
|
|
|
|
feedback = "\n".join(f"- {e}" for e in last_errors[:20])
|
|
|
|
|
|
sql = self._generate_sql(
|
|
|
|
|
|
question,
|
|
|
|
|
|
filtered_schema_str,
|
|
|
|
|
|
dialect,
|
|
|
|
|
|
validation_feedback=feedback,
|
2026-04-14 18:02:12 +08:00
|
|
|
|
dialog_context=dc_raw or None,
|
2026-04-16 13:48:44 +08:00
|
|
|
|
sql_stream_callback=sql_stream_callback,
|
2026-04-14 10:28:22 +08:00
|
|
|
|
)
|
2026-04-10 16:52:07 +08:00
|
|
|
|
last_sql = sql
|
|
|
|
|
|
except Exception as e:
|
|
|
|
|
|
last_errors = [f"SQL生成失败: {str(e)}"]
|
|
|
|
|
|
attempt += 1
|
|
|
|
|
|
continue
|
|
|
|
|
|
|
|
|
|
|
|
# === Step 3: 验证 ===
|
2026-04-14 10:28:22 +08:00
|
|
|
|
is_valid, errors, warnings, db_probe, empty_feedback = self._validate_sql(
|
|
|
|
|
|
sql,
|
|
|
|
|
|
filtered_schema_str,
|
|
|
|
|
|
dialect=dialect,
|
|
|
|
|
|
question=question,
|
2026-04-14 18:02:12 +08:00
|
|
|
|
dialog_context=dc_raw or None,
|
2026-04-10 16:52:07 +08:00
|
|
|
|
)
|
2026-04-14 10:28:22 +08:00
|
|
|
|
if db_probe is not None:
|
|
|
|
|
|
last_db_execution_status = db_probe
|
2026-04-10 16:52:07 +08:00
|
|
|
|
|
2026-04-14 10:28:22 +08:00
|
|
|
|
if not is_valid:
|
|
|
|
|
|
last_errors = errors
|
|
|
|
|
|
logger.warning(f" [FAIL] 验证失败:{errors}")
|
|
|
|
|
|
attempt += 1
|
|
|
|
|
|
continue
|
2026-04-10 16:52:07 +08:00
|
|
|
|
|
2026-04-14 10:28:22 +08:00
|
|
|
|
logger.info(f"[OK] SQL生成与验证通过({attempt + 1}次尝试)")
|
|
|
|
|
|
|
|
|
|
|
|
meta: Dict = dict(translation_meta)
|
|
|
|
|
|
if db_probe is not None:
|
|
|
|
|
|
meta["db_execution_status"] = db_probe
|
|
|
|
|
|
if empty_feedback:
|
|
|
|
|
|
meta["db_empty_feedback"] = empty_feedback
|
2026-04-14 18:02:12 +08:00
|
|
|
|
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)
|
2026-04-16 13:48:44 +08:00
|
|
|
|
if self._last_fewshot_golden:
|
|
|
|
|
|
gq, gsc = self._last_fewshot_golden
|
|
|
|
|
|
meta["fewshot_golden_reuse"] = True
|
|
|
|
|
|
meta["fewshot_golden_qid"] = gq
|
|
|
|
|
|
meta["fewshot_golden_score"] = gsc
|
2026-04-14 10:28:22 +08:00
|
|
|
|
|
|
|
|
|
|
result = GenerationResult(
|
|
|
|
|
|
sql=sql,
|
|
|
|
|
|
valid=True,
|
|
|
|
|
|
errors=[],
|
|
|
|
|
|
warnings=warnings,
|
|
|
|
|
|
tables_used=tables_used,
|
|
|
|
|
|
attempts=attempt + 1,
|
|
|
|
|
|
reasoning=reasoning if attempt == 0 else None,
|
|
|
|
|
|
metadata=meta,
|
|
|
|
|
|
)
|
|
|
|
|
|
if include_schema_in_result:
|
|
|
|
|
|
result.metadata["schema"] = filtered_schema_str
|
|
|
|
|
|
return result
|
2026-04-10 16:52:07 +08:00
|
|
|
|
|
|
|
|
|
|
# 达到最大重试次数
|
|
|
|
|
|
logger.error(f"[FAIL] 达到最大重试次数({self.max_retry}),生成失败")
|
2026-04-14 10:28:22 +08:00
|
|
|
|
fail_meta: Dict = dict(translation_meta)
|
|
|
|
|
|
if last_db_execution_status is not None:
|
|
|
|
|
|
fail_meta["db_execution_status"] = last_db_execution_status
|
2026-04-14 18:02:12 +08:00
|
|
|
|
if dc_raw:
|
|
|
|
|
|
fail_meta["dialog_context_chars"] = len(dc_raw)
|
2026-04-16 13:48:44 +08:00
|
|
|
|
if self._last_fewshot_golden:
|
|
|
|
|
|
gq, gsc = self._last_fewshot_golden
|
|
|
|
|
|
fail_meta["fewshot_golden_reuse"] = True
|
|
|
|
|
|
fail_meta["fewshot_golden_qid"] = gq
|
|
|
|
|
|
fail_meta["fewshot_golden_score"] = gsc
|
2026-04-10 16:52:07 +08:00
|
|
|
|
return GenerationResult(
|
|
|
|
|
|
sql=last_sql or "",
|
|
|
|
|
|
valid=False,
|
|
|
|
|
|
errors=last_errors,
|
|
|
|
|
|
tables_used=tables_used,
|
|
|
|
|
|
attempts=attempt,
|
2026-04-14 10:28:22 +08:00
|
|
|
|
metadata=fail_meta,
|
2026-04-10 16:52:07 +08:00
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
def build_vector_index(self, force_rebuild: bool = False) -> bool:
|
|
|
|
|
|
"""
|
|
|
|
|
|
构建向量索引(可选,提前构建可加速首次查询)
|
|
|
|
|
|
|
|
|
|
|
|
Args:
|
|
|
|
|
|
force_rebuild: 是否强制重建
|
|
|
|
|
|
|
|
|
|
|
|
Returns:
|
|
|
|
|
|
是否成功构建
|
|
|
|
|
|
"""
|
|
|
|
|
|
indexer = self._get_vector_index()
|
|
|
|
|
|
return indexer.build_index(
|
|
|
|
|
|
self.schema_manager,
|
|
|
|
|
|
force_rebuild=force_rebuild
|
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
def get_statistics(self) -> Dict:
|
|
|
|
|
|
"""获取统计信息"""
|
|
|
|
|
|
schema_stats = self.schema_manager.get_statistics()
|
|
|
|
|
|
|
|
|
|
|
|
indexer = self._get_vector_index()
|
|
|
|
|
|
index_stats = indexer.get_statistics()
|
|
|
|
|
|
|
|
|
|
|
|
return {
|
|
|
|
|
|
"schema": schema_stats,
|
|
|
|
|
|
"vector_index": index_stats,
|
|
|
|
|
|
"max_retry": self.max_retry,
|
|
|
|
|
|
"use_vector_search": self.use_vector_search,
|
|
|
|
|
|
}
|