0.1.1 暂存
This commit is contained in:
@@ -0,0 +1 @@
|
||||
# agents 包初始化
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -0,0 +1,703 @@
|
||||
"""
|
||||
Text2SQL 多智能体编排器
|
||||
协调 Schema Linker、SQL Generator、Validator 三个Agent
|
||||
"""
|
||||
|
||||
import logging
|
||||
import os # 新增
|
||||
from typing import Dict, List, Optional, Tuple
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from schema.manager import SchemaManager
|
||||
from schema.indexer import SchemaIndexer
|
||||
from llm.deepseek_client import DeepSeekClient, DeepSeekConfig
|
||||
from utils.fewshot_selector import FewShotSelector # 新增
|
||||
|
||||
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 多智能体编排器
|
||||
|
||||
工作流程:
|
||||
1. 粗筛候选表
|
||||
2. Schema Linker:LLM 精筛表
|
||||
3. 外键扩展 → 拼 Schema 子集
|
||||
4. SQL Generator:生成 SQL
|
||||
5. Validator:验证 SQL
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
schema_manager: SchemaManager,
|
||||
deepseek_api_key: Optional[str] = None,
|
||||
deepseek_config: Optional[DeepSeekConfig] = None,
|
||||
embedding_model_path: Optional[str] = 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,
|
||||
translate_english_to_zh: bool = True,
|
||||
):
|
||||
"""
|
||||
初始化编排器
|
||||
|
||||
Args:
|
||||
schema_manager: Schema管理器实例
|
||||
deepseek_api_key: DeepSeek API密钥(也可通过环境变量DEEPSEEK_API_KEY)
|
||||
deepseek_config: DeepSeek配置对象(优先于api_key)
|
||||
embedding_model_path: Qwen3-Embedding模型路径
|
||||
vector_db_path: 向量数据库路径
|
||||
max_retry: 最大重试次数(包含首次生成)
|
||||
use_vector_search: 是否使用向量检索粗筛
|
||||
translate_english_to_zh: 无中日韩字符的英文问句是否先译为中文再走检索与生成
|
||||
"""
|
||||
self.schema_manager = schema_manager
|
||||
self.max_retry = max_retry
|
||||
self.use_vector_search = use_vector_search
|
||||
self.translate_english_to_zh = translate_english_to_zh
|
||||
|
||||
# 初始化DeepSeek客户端
|
||||
if deepseek_config:
|
||||
self.deepseek = DeepSeekClient(deepseek_config)
|
||||
else:
|
||||
self.deepseek = DeepSeekClient(
|
||||
DeepSeekConfig(api_key=deepseek_api_key)
|
||||
)
|
||||
|
||||
# 初始化向量索引(延迟加载)
|
||||
self._vector_index: Optional[SchemaIndexer] = None
|
||||
self._vector_db_path = vector_db_path
|
||||
self._embedding_model_path = embedding_model_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:
|
||||
path = fewshot_samples_path or os.getenv(
|
||||
"FEWSHOT_DATA_PATH",
|
||||
"./data/experiences/all_samples.jsonl"
|
||||
)
|
||||
self.fewshot_selector = FewShotSelector(
|
||||
path,
|
||||
embedding_model_path=self._embedding_model_path,
|
||||
)
|
||||
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
|
||||
|
||||
logger.info(
|
||||
f"[OK] Text2SQLOrchestrator初始化完成: "
|
||||
f"max_retry={max_retry}, use_vector_search={use_vector_search}"
|
||||
+ (f", fewshot=on" if self.fewshot_enabled else "")
|
||||
+ (", en→zh=on" if self.translate_english_to_zh else ", en→zh=off")
|
||||
)
|
||||
|
||||
def _get_vector_index(self) -> SchemaIndexer:
|
||||
"""获取或创建向量索引(懒加载)"""
|
||||
if self._vector_index is None:
|
||||
from utils.embedding import get_embedder
|
||||
|
||||
embedder = get_embedder(self._embedding_model_path)
|
||||
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:
|
||||
# 不使用向量检索时,返回所有表
|
||||
return self.schema_manager.list_tables()
|
||||
|
||||
indexer = self._get_vector_index()
|
||||
|
||||
# 确保索引已构建
|
||||
if indexer.count() == 0:
|
||||
logger.info("向量索引为空,正在构建...")
|
||||
indexer.build_index(self.schema_manager, force_rebuild=True)
|
||||
|
||||
# 检索
|
||||
results = indexer.search(
|
||||
query=question,
|
||||
top_k=top_k,
|
||||
score_threshold=0.1 # 降低阈值以提高召回率(原0.2)
|
||||
)
|
||||
|
||||
candidate_tables = [r["table_name"] for r in results]
|
||||
logger.debug(f"粗筛候选表:{candidate_tables[:10]}...(共{len(candidate_tables)}个)")
|
||||
return candidate_tables
|
||||
|
||||
def _llm_select_tables(
|
||||
self,
|
||||
question: str,
|
||||
candidate_tables: List[str],
|
||||
max_tables: int = 5,
|
||||
) -> 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,
|
||||
table_list=table_list_str,
|
||||
)
|
||||
|
||||
relevant_tables = response.get("relevant_tables", [])
|
||||
reasoning = response.get("reasoning", "")
|
||||
|
||||
# 限制数量
|
||||
relevant_tables = relevant_tables[:max_tables]
|
||||
|
||||
logger.info(f"LLM精筛选中表:{relevant_tables}")
|
||||
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]
|
||||
|
||||
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
|
||||
|
||||
def _generate_sql(
|
||||
self,
|
||||
question: str,
|
||||
schema_str: str,
|
||||
dialect: str = "tsql",
|
||||
validation_feedback: Optional[str] = None,
|
||||
) -> str:
|
||||
"""
|
||||
SQL生成(SQL Generator Agent)
|
||||
|
||||
Args:
|
||||
question: 用户问题
|
||||
schema_str: Schema描述字符串
|
||||
dialect: SQL方言
|
||||
validation_feedback: 非空时附加到用户提示(重试时传入上次校验错误)
|
||||
|
||||
Returns:
|
||||
SQL语句
|
||||
"""
|
||||
from config.prompts import SQL_GENERATOR_SYSTEM, SQL_GENERATOR_USER
|
||||
from utils.sql_parser import normalize_sql_for_dialect
|
||||
|
||||
# Few-shot 增强
|
||||
if self.fewshot_enabled and self.fewshot_selector:
|
||||
try:
|
||||
examples = self.fewshot_selector.select(
|
||||
question=question,
|
||||
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}"
|
||||
logger.debug(f"已注入 {len(examples)} 个few-shot示例: {[ex.qid for ex in examples]}")
|
||||
except Exception as e:
|
||||
logger.warning(f"Few-shot检索失败: {e}")
|
||||
|
||||
dialect_label = dialect
|
||||
if dialect == "tsql":
|
||||
dialect_label = "Microsoft SQL Server (T-SQL)"
|
||||
|
||||
user_content = SQL_GENERATOR_USER.format(
|
||||
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 英文别名。"
|
||||
"\n**禁止**在单引号字符串字面量或 `N'…'` 中出现任何中日韩文字;"
|
||||
"业务中文须映射为 Schema 注释中的代码或通过维表 JOIN,勿写 `= '过户费'` 这类比对。"
|
||||
)
|
||||
|
||||
if validation_feedback:
|
||||
user_content += (
|
||||
"\n\n【上次校验未通过】请根据下列错误修正 SQL,并输出完整可执行查询;"
|
||||
"表名、列名必须与「当前Schema」中完全一致,不要臆造字段名。\n"
|
||||
f"{validation_feedback}"
|
||||
)
|
||||
|
||||
messages = [
|
||||
{"role": "system", "content": SQL_GENERATOR_SYSTEM},
|
||||
{"role": "user", "content": user_content},
|
||||
]
|
||||
|
||||
# 与选表一致:生成阶段默认贪心解码,减少同一中文问题多次 SQL 不一致
|
||||
response = self.deepseek.chat(messages, temperature=0.0, top_p=1.0)
|
||||
sql = response.content.strip()
|
||||
|
||||
# 清理可能的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)
|
||||
|
||||
logger.debug(f"生成的SQL:{sql[:200]}...")
|
||||
return sql
|
||||
|
||||
def _validate_sql(
|
||||
self,
|
||||
sql: str,
|
||||
schema_str: str,
|
||||
dialect: str = "tsql",
|
||||
question: str = "",
|
||||
) -> Tuple[bool, List[str], List[str], Optional[int], Optional[str]]:
|
||||
"""
|
||||
SQL验证(程序验证 + 库探针 + 按探针分支的 LLM)
|
||||
|
||||
Args:
|
||||
sql: SQL语句
|
||||
schema_str: Schema描述
|
||||
dialect: 与生成一致的 SQL 方言(sqlglot 名,默认 tsql)
|
||||
question: 用户自然语言(探针为 0 时用于生成补充说明)
|
||||
|
||||
Returns:
|
||||
(是否通过, 错误列表, 警告列表, 库执行探针状态, 无数据时的用户说明)
|
||||
探针:``1`` 至少一行数据;``0`` 执行成功但行数为 0;``-1`` 执行失败;``None`` 未配置库或未跑探针
|
||||
"""
|
||||
errors = []
|
||||
warnings = []
|
||||
db_execution_status: Optional[int] = None
|
||||
empty_feedback: Optional[str] = None
|
||||
|
||||
# === 阶段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)
|
||||
|
||||
# 危险操作检查
|
||||
from utils.validators import check_dangerous_operations, check_no_cjk_in_sql_string_literals
|
||||
danger_ok, danger_errors = check_dangerous_operations(sql)
|
||||
if not danger_ok:
|
||||
errors.extend(danger_errors)
|
||||
|
||||
# T-SQL:禁止中文等业务词出现在字符串字面量(如 FeeNatureID = '过户费')
|
||||
if dialect == "tsql":
|
||||
cjk_ok, cjk_errors = check_no_cjk_in_sql_string_literals(sql)
|
||||
if not cjk_ok:
|
||||
errors.extend(cjk_errors)
|
||||
|
||||
# === 阶段1.5:数据库试执行(仅程序校验全部通过时;需配置 database_url) ===
|
||||
if len(errors) == 0:
|
||||
from db.dbhub_tools import probe_sql_execution_status_ex
|
||||
|
||||
db_execution_status, db_probe_err = probe_sql_execution_status_ex(sql)
|
||||
if db_execution_status == -1:
|
||||
msg = (
|
||||
"数据库执行验证失败:SQL 在目标库执行报错(探针状态 -1),"
|
||||
"将据此重新生成 SQL。"
|
||||
)
|
||||
if db_probe_err:
|
||||
msg += f" 数据库返回:{db_probe_err}"
|
||||
errors.append(msg)
|
||||
elif db_execution_status is None:
|
||||
warnings.append(
|
||||
"未配置 database_url,已跳过数据库执行探针"
|
||||
)
|
||||
|
||||
# 探针 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)
|
||||
|
||||
# === 阶段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 = (
|
||||
"该 SQL 已在数据库成功执行,但返回的数据行数为 0(未查到匹配记录)。"
|
||||
"请将下方 SQL 与说明一并核对;若不符合预期,请补充或调整条件后再次提问。"
|
||||
)
|
||||
try:
|
||||
llm_fb = self.deepseek.empty_result_user_feedback(
|
||||
question=question,
|
||||
sql=sql,
|
||||
schema=schema_str,
|
||||
)
|
||||
empty_feedback = f"{prefix}\n\n【分析与建议】\n{llm_fb}"
|
||||
except Exception as e:
|
||||
logger.warning(f"无数据说明生成失败: {e}")
|
||||
empty_feedback = (
|
||||
f"{prefix}\n\n【分析与建议】\n"
|
||||
"未能自动生成详细分析。请补充时间范围、筛选条件或业务对象后重新提问。"
|
||||
)
|
||||
|
||||
is_valid = len(errors) == 0
|
||||
return is_valid, errors, warnings, db_execution_status, empty_feedback
|
||||
|
||||
def generate(
|
||||
self,
|
||||
question: str,
|
||||
dialect: str = "tsql",
|
||||
top_k_candidates: int = 20,
|
||||
include_schema_in_result: bool = False
|
||||
) -> GenerationResult:
|
||||
"""
|
||||
主生成流程
|
||||
|
||||
Args:
|
||||
question: 用户自然语言问题
|
||||
dialect: SQL方言(默认 tsql;与 sqlglot 一致)
|
||||
top_k_candidates: 粗筛候选表数量
|
||||
include_schema_in_result: 结果中是否包含使用的Schema字符串
|
||||
|
||||
Returns:
|
||||
GenerationResult对象
|
||||
"""
|
||||
from utils.question_locale import looks_like_english_only
|
||||
|
||||
original_question = (question or "").strip()
|
||||
translation_meta: Dict = {}
|
||||
work_question = original_question
|
||||
if self.translate_english_to_zh and looks_like_english_only(original_question):
|
||||
try:
|
||||
zh = self.deepseek.translate_nl_question_to_zh(original_question).strip()
|
||||
if zh and len(zh) >= 2:
|
||||
work_question = zh
|
||||
translation_meta["question_original"] = original_question
|
||||
translation_meta["question_zh_normalized"] = zh
|
||||
logger.info(
|
||||
"[GEN] 英文已译为中文:%s",
|
||||
zh[:120] + ("…" if len(zh) > 120 else ""),
|
||||
)
|
||||
else:
|
||||
logger.warning("[GEN] 英译中结果为空或过短,使用原文")
|
||||
except Exception as e:
|
||||
logger.warning("[GEN] 英译中失败,使用原文: %s", e)
|
||||
|
||||
question = work_question
|
||||
logger.info(f"[GEN] 开始生成SQL:{question[:50]}...")
|
||||
|
||||
attempt = 0
|
||||
last_sql = None
|
||||
last_errors = []
|
||||
last_db_execution_status: Optional[int] = None
|
||||
filtered_schema_str = ""
|
||||
tables_used = []
|
||||
|
||||
while attempt < self.max_retry:
|
||||
logger.info(f" 尝试 #{attempt + 1}")
|
||||
|
||||
# === Step 1: Schema筛选(仅首次) ===
|
||||
if attempt == 0:
|
||||
# 1.1 粗筛
|
||||
candidate_tables = self._coarse_filter(question, top_k=top_k_candidates)
|
||||
|
||||
# 1.2 LLM精筛
|
||||
relevant_tables, reasoning = self._llm_select_tables(
|
||||
question,
|
||||
candidate_tables,
|
||||
)
|
||||
relevant_tables = self._prioritize_broker_tables(question, relevant_tables)
|
||||
|
||||
# 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
|
||||
)
|
||||
logger.info(f" 选中表:{relevant_tables},扩展后:{expanded_tables}")
|
||||
else:
|
||||
# 重试:在子 Schema 中并入「上次失败 SQL」实际引用到的表,并对齐程序校验与生成上下文
|
||||
logger.info(f" 重试使用之前的Schema({len(tables_used)}张表)")
|
||||
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,
|
||||
)
|
||||
|
||||
# === Step 2: SQL生成 ===
|
||||
try:
|
||||
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,
|
||||
)
|
||||
last_sql = sql
|
||||
except Exception as e:
|
||||
last_errors = [f"SQL生成失败: {str(e)}"]
|
||||
attempt += 1
|
||||
continue
|
||||
|
||||
# === Step 3: 验证 ===
|
||||
is_valid, errors, warnings, db_probe, empty_feedback = self._validate_sql(
|
||||
sql,
|
||||
filtered_schema_str,
|
||||
dialect=dialect,
|
||||
question=question,
|
||||
)
|
||||
if db_probe is not None:
|
||||
last_db_execution_status = db_probe
|
||||
|
||||
if not is_valid:
|
||||
last_errors = errors
|
||||
logger.warning(f" [FAIL] 验证失败:{errors}")
|
||||
attempt += 1
|
||||
continue
|
||||
|
||||
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
|
||||
|
||||
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
|
||||
|
||||
# 达到最大重试次数
|
||||
logger.error(f"[FAIL] 达到最大重试次数({self.max_retry}),生成失败")
|
||||
fail_meta: Dict = dict(translation_meta)
|
||||
if last_db_execution_status is not None:
|
||||
fail_meta["db_execution_status"] = last_db_execution_status
|
||||
return GenerationResult(
|
||||
sql=last_sql or "",
|
||||
valid=False,
|
||||
errors=last_errors,
|
||||
tables_used=tables_used,
|
||||
attempts=attempt,
|
||||
metadata=fail_meta,
|
||||
)
|
||||
|
||||
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,
|
||||
}
|
||||
@@ -0,0 +1,163 @@
|
||||
"""
|
||||
Schema Linker Agent - 表筛选专家
|
||||
从大量数据表中识别与用户问题相关的表
|
||||
"""
|
||||
|
||||
import logging
|
||||
from typing import Dict, List, Tuple, Optional
|
||||
import json
|
||||
|
||||
from camel.agents import ChatAgent
|
||||
from camel.models import ChatModel
|
||||
from config.prompts import SCHEMA_LINKER_SYSTEM
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class SchemaLinkerAgent:
|
||||
"""
|
||||
Schema Linker Agent
|
||||
|
||||
职责:
|
||||
- 分析用户问题中的实体和意图
|
||||
- 从候选表中筛选真正相关的表
|
||||
- 提供选择理由
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model: ChatModel,
|
||||
system_message: Optional[str] = None
|
||||
):
|
||||
"""
|
||||
初始化Agent
|
||||
|
||||
Args:
|
||||
model: CAMEL AI模型实例
|
||||
system_message: 系统提示词(默认使用SCHEMA_LINKER_SYSTEM)
|
||||
"""
|
||||
self.system_message = system_message or SCHEMA_LINKER_SYSTEM
|
||||
self.agent = ChatAgent(
|
||||
system_message=self.system_message,
|
||||
model=model,
|
||||
)
|
||||
logger.info("[OK] SchemaLinkerAgent初始化完成")
|
||||
|
||||
def select_tables(
|
||||
self,
|
||||
question: str,
|
||||
candidate_tables: List[str],
|
||||
table_metadata: Optional[Dict[str, str]] = None,
|
||||
max_tables: int = 5
|
||||
) -> Tuple[List[str], str]:
|
||||
"""
|
||||
选择相关表
|
||||
|
||||
Args:
|
||||
question: 用户问题
|
||||
candidate_tables: 候选表列表(粗筛结果)
|
||||
table_metadata: 表元数据 {表名: 注释}
|
||||
max_tables: 最多返回表数量
|
||||
|
||||
Returns:
|
||||
(相关表列表, 推理理由)
|
||||
"""
|
||||
# 构造候选表信息字符串
|
||||
if table_metadata:
|
||||
table_list = "\n".join([
|
||||
f"- {tbl}: {table_metadata.get(tbl, '无描述')}"
|
||||
for tbl in candidate_tables
|
||||
])
|
||||
else:
|
||||
table_list = "\n".join([f"- {tbl}" for tbl in candidate_tables])
|
||||
|
||||
from config.prompts import SCHEMA_LINKER_USER
|
||||
prompt = SCHEMA_LINKER_USER.format(
|
||||
question=question,
|
||||
table_list=table_list
|
||||
)
|
||||
|
||||
try:
|
||||
response = self.agent.step(prompt)
|
||||
content = response.msg.content.strip()
|
||||
|
||||
# 解析JSON响应
|
||||
result = self._parse_json_response(content)
|
||||
|
||||
relevant_tables = result.get("relevant_tables", [])
|
||||
reasoning = result.get("reasoning", "")
|
||||
|
||||
# 限制数量
|
||||
relevant_tables = relevant_tables[:max_tables]
|
||||
|
||||
logger.info(
|
||||
f"SchemaLinker选中 {len(relevant_tables)} 张表: {relevant_tables}"
|
||||
)
|
||||
return relevant_tables, reasoning
|
||||
|
||||
except json.JSONDecodeError as e:
|
||||
logger.error(f"JSON解析失败: {e}, 原始内容: {content[:200]}")
|
||||
# 降级:返回前max_tables个候选表
|
||||
return candidate_tables[:max_tables], "JSON解析失败,使用粗筛结果"
|
||||
except Exception as e:
|
||||
logger.error(f"Agent调用失败: {e}")
|
||||
return candidate_tables[:max_tables], f"Agent错误: {str(e)}"
|
||||
|
||||
def _parse_json_response(self, content: str) -> Dict:
|
||||
"""
|
||||
解析Agent的JSON响应
|
||||
|
||||
处理可能的markdown代码块包裹
|
||||
"""
|
||||
import json
|
||||
|
||||
# 尝试提取```json```块
|
||||
if "```json" in content:
|
||||
start = content.find("```json") + 7
|
||||
end = content.find("```", start)
|
||||
if end != -1:
|
||||
content = content[start:end].strip()
|
||||
elif "```" in content:
|
||||
start = content.find("```") + 3
|
||||
end = content.find("```", start)
|
||||
if end != -1:
|
||||
content = content[start:end].strip()
|
||||
|
||||
return json.loads(content)
|
||||
|
||||
def expand_by_relations(
|
||||
self,
|
||||
selected_tables: List[str],
|
||||
schema_manager,
|
||||
depth: int = 1
|
||||
) -> List[str]:
|
||||
"""
|
||||
通过外键关系扩展表(备选方案,也可在Orchestrator中完成)
|
||||
|
||||
Args:
|
||||
selected_tables: 已选中的表
|
||||
schema_manager: Schema管理器
|
||||
depth: 递归深度
|
||||
|
||||
Returns:
|
||||
扩展后的表列表
|
||||
"""
|
||||
result = set(selected_tables)
|
||||
|
||||
for tbl_name in selected_tables:
|
||||
table = 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)
|
||||
|
||||
# 引用当前表的表(反向外键)
|
||||
for other in 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)
|
||||
|
||||
return list(result)
|
||||
@@ -0,0 +1,206 @@
|
||||
"""
|
||||
SQL Generator Agent - SQL生成专家
|
||||
根据Schema和问题生成高质量SQL
|
||||
"""
|
||||
|
||||
import logging
|
||||
import json
|
||||
import re
|
||||
from typing import Dict, Optional
|
||||
|
||||
from camel.agents import ChatAgent
|
||||
from camel.models import ChatModel
|
||||
from config.prompts import SQL_GENERATOR_SYSTEM
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class SQLGeneratorAgent:
|
||||
"""
|
||||
SQL Generator Agent
|
||||
|
||||
职责:
|
||||
- 理解用户问题和Schema结构
|
||||
- 生成准确的SQL语句
|
||||
- 处理复杂的JOIN、聚合、子查询
|
||||
- 遵循金融/证券业务特殊规则
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model: ChatModel,
|
||||
system_message: Optional[str] = None,
|
||||
dialect: str = "tsql"
|
||||
):
|
||||
"""
|
||||
初始化Agent
|
||||
|
||||
Args:
|
||||
model: CAMEL AI模型实例
|
||||
system_message: 系统提示词
|
||||
dialect: SQL方言
|
||||
"""
|
||||
self.dialect = dialect
|
||||
self.system_message = system_message or SQL_GENERATOR_SYSTEM
|
||||
self.agent = ChatAgent(
|
||||
system_message=self.system_message,
|
||||
model=model,
|
||||
)
|
||||
logger.info(f"[OK] SQLGeneratorAgent初始化完成 (dialect={dialect})")
|
||||
|
||||
def generate(
|
||||
self,
|
||||
question: str,
|
||||
schema_str: str,
|
||||
examples: Optional[List[Dict]] = None,
|
||||
dialect: Optional[str] = None
|
||||
) -> str:
|
||||
"""
|
||||
生成SQL
|
||||
|
||||
Args:
|
||||
question: 用户问题
|
||||
schema_str: Schema描述字符串
|
||||
examples: Few-shot示例列表
|
||||
dialect: 覆盖默认dialect
|
||||
|
||||
Returns:
|
||||
SQL语句字符串
|
||||
"""
|
||||
from config.prompts import SQL_GENERATOR_USER
|
||||
|
||||
dialect = dialect or self.dialect
|
||||
prompt = SQL_GENERATOR_USER.format(
|
||||
schema=schema_str,
|
||||
question=question,
|
||||
dialect=dialect
|
||||
)
|
||||
|
||||
# 添加few-shot示例(如果有)
|
||||
if examples:
|
||||
prompt = self._inject_examples(prompt, examples)
|
||||
|
||||
try:
|
||||
response = self.agent.step(prompt)
|
||||
sql = self._extract_sql(response.msg.content)
|
||||
|
||||
logger.debug(f"生成的SQL: {sql[:200]}...")
|
||||
return sql
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"SQL生成失败: {e}")
|
||||
raise
|
||||
|
||||
def _extract_sql(self, content: str) -> str:
|
||||
"""
|
||||
从Agent响应中提取SQL语句
|
||||
|
||||
处理:
|
||||
- Markdown代码块 (```sql ... ```)
|
||||
- 纯SQL文本
|
||||
- JSON格式 {"sql": "..."}
|
||||
"""
|
||||
content = content.strip()
|
||||
|
||||
# 尝试提取```sql```块
|
||||
if "```sql" in content:
|
||||
start = content.find("```sql") + 6
|
||||
end = content.find("```", start)
|
||||
if end != -1:
|
||||
return content[start:end].strip()
|
||||
|
||||
# 尝试提取通用代码块```
|
||||
if "```" in content:
|
||||
start = content.find("```") + 3
|
||||
end = content.find("```", start)
|
||||
if end != -1:
|
||||
return content[start:end].strip()
|
||||
|
||||
# 尝试解析JSON
|
||||
try:
|
||||
data = json.loads(content)
|
||||
if "sql" in data:
|
||||
return data["sql"]
|
||||
except json.JSONDecodeError:
|
||||
pass
|
||||
|
||||
# 返回原始内容(假设是纯SQL)
|
||||
return content
|
||||
|
||||
def _inject_examples(
|
||||
self,
|
||||
prompt: str,
|
||||
examples: List[Dict]
|
||||
) -> str:
|
||||
"""
|
||||
注入few-shot示例到提示词
|
||||
|
||||
Args:
|
||||
prompt: 原始提示词
|
||||
examples: 示例列表,每项为 {"question": "...", "schema": "...", "sql": "..."}
|
||||
|
||||
Returns:
|
||||
增强后的提示词
|
||||
"""
|
||||
examples_text = []
|
||||
for ex in examples[:3]: # 最多3个示例
|
||||
examples_text.append(
|
||||
f"示例:\n问题:{ex['question']}\n"
|
||||
f"Schema: {ex['schema'][:200]}...\n"
|
||||
f"SQL: {ex['sql']}"
|
||||
)
|
||||
|
||||
examples_block = "\n\n".join(examples_text)
|
||||
|
||||
# 插入到提示词末尾(要求之前)
|
||||
return f"{prompt}\n\n参考示例:\n{examples_block}\n\n请生成SQL:"
|
||||
|
||||
def generate_with_reasoning(
|
||||
self,
|
||||
question: str,
|
||||
schema_str: str,
|
||||
dialect: Optional[str] = None
|
||||
) -> Tuple[str, str]:
|
||||
"""
|
||||
生成SQL并返回解释
|
||||
|
||||
Args:
|
||||
question: 用户问题
|
||||
schema_str: Schema描述
|
||||
dialect: SQL方言
|
||||
|
||||
Returns:
|
||||
(SQL语句, 解释)
|
||||
"""
|
||||
from config.prompts import SQL_GENERATOR_USER
|
||||
|
||||
dialect = dialect or self.dialect
|
||||
prompt = f"""
|
||||
{sql_generator_user.format(schema=schema_str, question=question, dialect=dialect)}
|
||||
|
||||
请同时输出SQL和简要解释(JSON格式):
|
||||
{{
|
||||
"sql": "SELECT ...",
|
||||
"explanation": "SQL逻辑说明"
|
||||
}}
|
||||
"""
|
||||
|
||||
try:
|
||||
response = self.agent.step(prompt)
|
||||
content = response.msg.content.strip()
|
||||
|
||||
# 尝试解析JSON
|
||||
try:
|
||||
if "```json" in content:
|
||||
start = content.find("```json") + 7
|
||||
end = content.find("```", start)
|
||||
content = content[start:end].strip()
|
||||
data = json.loads(content)
|
||||
return data.get("sql", ""), data.get("explanation", "")
|
||||
except (json.JSONDecodeError, ValueError):
|
||||
# 降级:提取SQL,解释为空
|
||||
return self._extract_sql(content), ""
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"生成失败: {e}")
|
||||
raise
|
||||
@@ -0,0 +1,299 @@
|
||||
"""
|
||||
Validator Agent - SQL审核员
|
||||
验证SQL的正确性和安全性
|
||||
"""
|
||||
|
||||
import logging
|
||||
import json
|
||||
from typing import Dict, List, Tuple, Optional
|
||||
|
||||
from camel.agents import ChatAgent
|
||||
from camel.models import ChatModel
|
||||
from config.prompts import VALIDATOR_SYSTEM
|
||||
|
||||
from utils.validators import full_validation_pipeline
|
||||
from utils.sql_parser import validate_sql_syntax, validate_schema_consistency
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class ValidatorAgent:
|
||||
"""
|
||||
Validator Agent
|
||||
|
||||
职责:
|
||||
- 语法正确性验证
|
||||
- Schema一致性检查
|
||||
- 安全性检查(禁止DML/DDL)
|
||||
- 性能问题识别
|
||||
- 提供修正建议
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model: ChatModel,
|
||||
system_message: Optional[str] = None,
|
||||
schema_manager = None
|
||||
):
|
||||
"""
|
||||
初始化Agent
|
||||
|
||||
Args:
|
||||
model: CAMEL AI模型实例
|
||||
system_message: 系统提示词
|
||||
schema_manager: Schema管理器(程序验证用)
|
||||
"""
|
||||
self.schema_manager = schema_manager
|
||||
self.system_message = system_message or VALIDATOR_SYSTEM
|
||||
self.agent = ChatAgent(
|
||||
system_message=self.system_message,
|
||||
model=model,
|
||||
)
|
||||
logger.info("[OK] ValidatorAgent初始化完成")
|
||||
|
||||
def validate(
|
||||
self,
|
||||
sql: str,
|
||||
schema_str: str,
|
||||
dialect: str = "tsql",
|
||||
check_dangerous: bool = True
|
||||
) -> Dict:
|
||||
"""
|
||||
完整验证流程(程序 + LLM双重验证)
|
||||
|
||||
Args:
|
||||
sql: SQL语句
|
||||
schema_str: Schema描述
|
||||
dialect: SQL方言
|
||||
check_dangerous: 是否检查危险操作
|
||||
|
||||
Returns:
|
||||
验证结果字典 {
|
||||
"valid": bool,
|
||||
"errors": [...],
|
||||
"warnings": [...],
|
||||
"suggestions": [...]
|
||||
}
|
||||
"""
|
||||
result = {
|
||||
"valid": True,
|
||||
"errors": [],
|
||||
"warnings": [],
|
||||
"suggestions": []
|
||||
}
|
||||
|
||||
# === 阶段1:程序验证(快速、确定性) ===
|
||||
program_result = self._program_validation(
|
||||
sql, dialect, check_dangerous
|
||||
)
|
||||
|
||||
result["errors"].extend(program_result.get("errors", []))
|
||||
result["warnings"].extend(program_result.get("warnings", []))
|
||||
result["suggestions"].extend(program_result.get("suggestions", []))
|
||||
|
||||
# 如果程序验证已发现致命错误,跳过LLM验证
|
||||
if program_result.get("fatal", False):
|
||||
result["valid"] = False
|
||||
return result
|
||||
|
||||
# === 阶段2:LLM语义验证 ===
|
||||
try:
|
||||
llm_result = self._llm_validation(sql, schema_str)
|
||||
result["errors"].extend(llm_result.get("errors", []))
|
||||
result["warnings"].extend(llm_result.get("warnings", []))
|
||||
result["suggestions"].extend(llm_result.get("suggestions", []))
|
||||
except Exception as e:
|
||||
logger.warning(f"LLM验证失败,使用程序验证结果: {e}")
|
||||
result["warnings"].append(f"LLM验证异常: {str(e)}")
|
||||
|
||||
result["valid"] = len(result["errors"]) == 0
|
||||
return result
|
||||
|
||||
def _program_validation(
|
||||
self,
|
||||
sql: str,
|
||||
dialect: str,
|
||||
check_dangerous: bool
|
||||
) -> Dict:
|
||||
"""
|
||||
程序验证(规则引擎)
|
||||
|
||||
Returns:
|
||||
{"errors": [], "warnings": [], "suggestions": [], "fatal": bool}
|
||||
"""
|
||||
errors = []
|
||||
warnings = []
|
||||
suggestions = []
|
||||
fatal = False
|
||||
|
||||
# 1. 语法检查
|
||||
syntax_ok, syntax_errors = validate_sql_syntax(sql, dialect)
|
||||
if not syntax_ok:
|
||||
errors.extend(syntax_errors)
|
||||
fatal = True # 语法错误无法修复,直接失败
|
||||
return {
|
||||
"errors": errors, "warnings": warnings,
|
||||
"suggestions": suggestions, "fatal": fatal
|
||||
}
|
||||
|
||||
# 2. Schema一致性(如果有schema_manager)
|
||||
if self.schema_manager:
|
||||
schema_ok, schema_errors = validate_schema_consistency(
|
||||
sql, self.schema_manager, dialect
|
||||
)
|
||||
if not schema_ok:
|
||||
errors.extend(schema_errors)
|
||||
# Schema错误通常也是fatal的
|
||||
fatal = True
|
||||
|
||||
# 3. 危险操作检查
|
||||
if check_dangerous:
|
||||
from utils.validators import check_dangerous_operations
|
||||
safe, dangers = check_dangerous_operations(sql)
|
||||
if not safe:
|
||||
errors.append(f"包含危险操作: {', '.join(dangers)}")
|
||||
fatal = True
|
||||
|
||||
return {
|
||||
"errors": errors,
|
||||
"warnings": warnings,
|
||||
"suggestions": suggestions,
|
||||
"fatal": fatal
|
||||
}
|
||||
|
||||
def _llm_validation(self, sql: str, schema_str: str) -> Dict:
|
||||
"""
|
||||
LLM语义验证
|
||||
|
||||
Args:
|
||||
sql: SQL语句
|
||||
schema_str: Schema描述
|
||||
|
||||
Returns:
|
||||
验证结果
|
||||
"""
|
||||
from config.prompts import VALIDATOR_USER
|
||||
|
||||
prompt = VALIDATOR_USER.format(sql=sql, schema=schema_str)
|
||||
|
||||
try:
|
||||
response = self.agent.step(prompt)
|
||||
content = response.msg.content.strip()
|
||||
|
||||
# 解析JSON响应
|
||||
result = self._parse_json_response(content)
|
||||
|
||||
# 标准化字段
|
||||
return {
|
||||
"errors": result.get("errors", []),
|
||||
"warnings": result.get("warnings", []),
|
||||
"suggestions": result.get("suggestions", []),
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"LLM验证异常: {e}")
|
||||
return {
|
||||
"errors": [f"LLM验证失败: {str(e)}"],
|
||||
"warnings": [],
|
||||
"suggestions": []
|
||||
}
|
||||
|
||||
def _parse_json_response(self, content: str) -> Dict:
|
||||
"""解析JSON响应"""
|
||||
try:
|
||||
# 提取```json```块
|
||||
if "```json" in content:
|
||||
start = content.find("```json") + 7
|
||||
end = content.find("```", start)
|
||||
if end != -1:
|
||||
content = content[start:end].strip()
|
||||
|
||||
return json.loads(content)
|
||||
except json.JSONDecodeError as e:
|
||||
logger.warning(f"JSON解析失败: {e}, content={content[:200]}")
|
||||
return {"errors": ["验证结果解析失败"], "warnings": [], "suggestions": []}
|
||||
|
||||
def quick_check(self, sql: str) -> Tuple[bool, List[str]]:
|
||||
"""
|
||||
快速检查(仅语法和危险操作)
|
||||
|
||||
Returns:
|
||||
(是否通过, 错误列表)
|
||||
"""
|
||||
errors = []
|
||||
|
||||
# 语法
|
||||
syntax_ok, syntax_errors = validate_sql_syntax(sql)
|
||||
if not syntax_ok:
|
||||
errors.extend(syntax_errors)
|
||||
return False, errors
|
||||
|
||||
# 危险操作
|
||||
from utils.validators import check_dangerous_operations
|
||||
safe, dangers = check_dangerous_operations(sql)
|
||||
if not safe:
|
||||
errors.append(f"危险操作: {', '.join(dangers)}")
|
||||
return False, errors
|
||||
|
||||
return True, []
|
||||
|
||||
def suggest_fixes(
|
||||
self,
|
||||
sql: str,
|
||||
errors: List[str],
|
||||
schema_str: str
|
||||
) -> List[str]:
|
||||
"""
|
||||
根据错误建议修复方案
|
||||
|
||||
Args:
|
||||
sql: 原始SQL
|
||||
errors: 错误列表
|
||||
schema_str: Schema描述
|
||||
|
||||
Returns:
|
||||
修复建议列表
|
||||
"""
|
||||
suggestions = []
|
||||
|
||||
# 常见错误模式匹配
|
||||
for error in errors:
|
||||
error_lower = error.lower()
|
||||
|
||||
if "field not exist" in error_lower or "字段不存在" in error_lower:
|
||||
suggestions.append(
|
||||
"建议:检查字段名拼写,或使用schema_manager.get_table(table).column_names查看可用字段"
|
||||
)
|
||||
|
||||
if "table not exist" in error_lower or "表不存在" in error_lower:
|
||||
suggestions.append(
|
||||
"建议:检查表名拼写,或使用schema_manager.list_tables()查看所有表"
|
||||
)
|
||||
|
||||
if "missing join condition" in error_lower or "缺少on条件" in error_lower:
|
||||
suggestions.append(
|
||||
"建议:为每个JOIN添加明确的ON条件,基于外键关系"
|
||||
)
|
||||
|
||||
if "group by" in error_lower:
|
||||
suggestions.append(
|
||||
"建议:SELECT中的所有非聚合字段都必须出现在GROUP BY子句中"
|
||||
)
|
||||
|
||||
# LLM补充建议
|
||||
if len(suggestions) < len(errors):
|
||||
try:
|
||||
prompt = f"""
|
||||
SQL: {sql}
|
||||
错误: {errors}
|
||||
Schema: {schema_str[:1000]}
|
||||
|
||||
请给出2-3条具体的修复建议(简洁明了):
|
||||
"""
|
||||
response = self.agent.step(prompt)
|
||||
llm_suggestions = response.msg.content.strip().split('\n')
|
||||
suggestions.extend([s for s in llm_suggestions if s.strip()])
|
||||
except Exception as e:
|
||||
logger.warning(f"获取LLM建议失败: {e}")
|
||||
|
||||
return suggestions[:5] # 最多5条
|
||||
Reference in New Issue
Block a user