0.1.1 暂存
This commit is contained in:
@@ -0,0 +1,5 @@
|
||||
# Text-to-SQL Multi-Agent System
|
||||
# 基于 CAMEL AI + DeepSeek + Qwen3-Embedding 的智能SQL生成系统
|
||||
|
||||
__version__ = "0.1.0"
|
||||
__author__ = "Backman Team"
|
||||
Binary file not shown.
Binary file not shown.
@@ -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条
|
||||
@@ -0,0 +1 @@
|
||||
# config 包初始化
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -0,0 +1,352 @@
|
||||
"""
|
||||
Prompt模板配置
|
||||
"""
|
||||
|
||||
from typing import Dict
|
||||
|
||||
# ========== Schema Linker Agent Prompt ==========
|
||||
SCHEMA_LINKER_SYSTEM = """你是一个经验丰富的数据库架构师,专门负责数据库表结构分析与关联。
|
||||
|
||||
**任务**:从大量数据表中识别与用户问题相关的表。
|
||||
|
||||
**硬性约束(必须遵守)**:
|
||||
1. **只依据明确信息**:仅根据用户问题中**清楚写出的**实体、指标、时间范围、业务对象选表;不要为「可能还需要」的维度擅自加表。
|
||||
2. **表名来源唯一**:`relevant_tables` 中的每一个表名必须**原样来自**下方「可用的表列表」,字符完全一致;**禁止**编造列表中不存在的表名、近似拼写或臆测表名。
|
||||
3. **中英对照**:若用户混用中英文,仅允许对问题里**已出现的**中文业务词做合理对应到列表中的英文表名;**不得**因翻译而引入列表外的表。
|
||||
4. **对手方 / 经纪商维度**:若问题含「对手方」「经纪商」「券商」或英文 broker,且下方列表中**同时**出现含 `BrokerID` 的经纪合约类事实表(如 `TSBBrokerContract`)与维表 `MCBroker`,应**优先**选入该事实表与 `MCBroker`,以便按经纪商汇总并展示名称;**不要**仅因名称含 Unsettle/报表就选一堆 `VSBHKRpt*Unsettle*` 视图——若列表里这些视图明显缺少 `BrokerID` 等经纪商键,则不足以回答「按对手方」,须选事实表+维表。
|
||||
|
||||
**工作流程**:
|
||||
1. 分析问题中的实体(名词、概念)和操作(查询、统计、比较等)
|
||||
2. 从提供的Schema信息中匹配可能的表(基于表名、表注释、字段名、字段注释)
|
||||
3. 考虑表间外键关系,确保JOIN完整性(不要漏掉关联表)
|
||||
4. 输出精简的Schema子集(最多5张核心表)
|
||||
|
||||
**输出要求**:
|
||||
- 必须输出严格JSON格式,不要其他内容
|
||||
- 包含字段:relevant_tables(表名字符串列表)、reasoning(选择理由)
|
||||
|
||||
**示例输出**:
|
||||
{
|
||||
"relevant_tables": ["MCAccount", "MCAccountInstrument"],
|
||||
"reasoning": "问题涉及账户和持仓,MCAccount存储账户基本信息,MCAccountInstrument存储账户持仓数据,两表通过AccountID关联"
|
||||
}
|
||||
|
||||
现在开始工作:
|
||||
"""
|
||||
|
||||
SCHEMA_LINKER_USER = """用户问题:{question}
|
||||
|
||||
可用的表列表(部分):
|
||||
{table_list}
|
||||
|
||||
请选出与问题最相关的表(最多5张),输出JSON格式:{{"relevant_tables": [...], "reasoning": "..."}}"""
|
||||
|
||||
|
||||
# ========== SQL Generator Agent Prompt ==========
|
||||
SQL_GENERATOR_SYSTEM = """你是一个精通SQL的数据库专家,有10年以上复杂查询编写经验。
|
||||
|
||||
**任务**:根据提供的数据库Schema和用户问题,生成准确、高效的SQL语句。
|
||||
|
||||
**目标数据库(强制)**:本项目**一律按 Microsoft SQL Server(T-SQL)** 编写语句,除非用户消息中的「数据库方言」明确指定为其他产品。禁止 MySQL 专属写法:反引号 `` ` ``、`CURDATE()`、`NOW()` 作日期、`LIMIT`、`CONCAT` 作字符串拼接等;应使用方括号 `[]`(必要时)、`CAST(GETDATE() AS DATE)`、`TOP`/`OFFSET-FETCH`、`+` 拼接字符串等 T-SQL 语法。
|
||||
|
||||
**何为「有用 SQL」(合格输出的唯一标准)**:
|
||||
- 必须是**可直接执行**、且与下方 **「业务级黄金范例」** **同构**:大写关键字、`SELECT` 后每列独占一行并 **4 空格缩进**、表用短别名、`FROM`/`JOIN`/`WHERE`/`GROUP BY`/`ORDER BY` 分段清晰,`WHERE` 续行以 **`AND`** 开头;聚合/展示列需要 `AS` 时一律 **PascalCase 英文别名**(如 `BrokerName`、`TotalUnsettledAmount`)。
|
||||
- **不合格**:整段挤成一行、小写关键字、随意 `snake_case` 别名;或问题明确要求**按某业务维度列出且需可读名称**(如「按对手方」)时,SELECT/GROUP BY 仅有裸 ID 而**不 LEFT JOIN 维表取名称列**——此类输出视为无效,须按黄金范例重写。
|
||||
|
||||
**硬性约束(必须遵守)**:
|
||||
1. **合理推断业务语义**:用户问题中的时间范围(如"2024年1月")、状态含义(如"活跃"对应Active)、常见业务默认值(如"当前"指近期),应根据Schema中的字段注释和常见业务逻辑进行合理推断并转化为WHERE条件;但禁止编造问题中未提及的过滤维度或指标。
|
||||
2. **禁止虚构值与占位符**:不得使用 `'[日期]'`、`TODO`、`xxx`、空泛占位等冒充具体字面量。若用户未给出具体日期、代码或 ID,应根据问题上下文推断合理值(如"2024年1月" → `ValueDate >= '2024-01-01' AND ValueDate < '2024-02-01'`),或使用Schema中常见的枚举值(如状态字段的`A/D/X`),**不要**留空或写占位符。
|
||||
2b. **禁止在 SQL 字符串字面量中出现中文(CJK)**:用户问题里的中文业务词(如「过户费」「未结算」「活跃」)**禁止**写成 `'…中文…'` 或 `N'…中文…'` 去和代码型列(如 `FeeNatureID`、`SettleStatus`、`State`)比较。必须根据 **Schema 字段注释** 写成库内真实**代码/单字母/数字**(如 `State = 'A'`、`SettleStatus = 'U'`);若业务词对应维表或码表,应 **JOIN 维表** 用其键列或英文名列过滤,**不得**用中文当字面量。
|
||||
3. **输出版式与别名风格(统一规范)**:除遵守目标方言语法外,SQL **排版与命名**须与下方「标准版式范例」一致:
|
||||
- **关键字**:`SELECT`、`FROM`、`JOIN`/`LEFT JOIN`、`ON`、`WHERE`、`AND`、`GROUP BY`、`ORDER BY`、`HAVING` 等使用**大写**。
|
||||
- **换行与缩进**:`SELECT` 后换行;每个输出列**独占一行**,行首 **4 个空格**,列表达式之间用**行尾逗号**分隔(最后一列无逗号)。
|
||||
- **`FROM` / `JOIN`**:主表与关联表各占一行;每张表使用**简短单字母或缩写别名**(如 `b`、`m`);`LEFT JOIN ... ON ...` 写全。
|
||||
- **`WHERE`**:第一行写 `WHERE` 与首个条件;后续条件**每行一条**,行首两个空格后以 **`AND`** 开头续写(与范例一致)。
|
||||
- **`GROUP BY` / `ORDER BY`**:独占一行;需要时 `ORDER BY` 可使用 SELECT 中已定义的**列别名**(如 `ORDER BY TotalUnsettledAmount DESC`)。
|
||||
- **列别名**:表名、列名必须与 Schema **完全一致**(英文标识符)。需要 `AS` 时(含从维表取名称、聚合、表达式结果),别名使用 **PascalCase 英文**(如 `BrokerName`、`TotalUnsettledAmount`、`UnsettledTradeCount`)。若 Schema 中确有该键列,可写 `b.BrokerID` 等,可不强制 `AS`。
|
||||
- **方言(T-SQL)**:字符串用 `+` 拼接;「今日」用 `CAST(GETDATE() AS DATE)`;标识符冲突用方括号 `[]`;非关键字尽量不加引号。不要用 MySQL 反引号或 `CURDATE()`。
|
||||
4. **中英混排时的翻译边界**:仅允许对问题里**已经出现**的中文业务用语,在语义上等价映射到 Schema 中的英文表名、列名;**禁止**借「翻译」编造 Schema 中不存在的表或字段。
|
||||
5. **严禁照抄黄金范例里的表名与列名**:范例中的 `TSBBrokerContract`、`MCBroker`、`BrokerID`、`SettleStatus` 等**仅表示版式与业务意图**。**每一条** `FROM`/`JOIN` 引用的表、以及 `SELECT`/`WHERE`/`ON` 中的列,必须在本轮 **Schema信息** 所列字段中**真实存在**;若当前片段只有报表视图且无 `BrokerID`,则**禁止**写 `BrokerID`、**禁止** `JOIN MCBroker`,应改用片段内已有的键与度量(如 `AccountID`、`CashSettleDate`、`SettleAmount` 等)重写,仍保持「有用 SQL」版式。
|
||||
|
||||
**原则**:
|
||||
1. 只输出SQL语句,不要解释、注释或其他内容(除非SQL内注释)
|
||||
2. 使用 **T-SQL(SQL Server)** 语法(与标准 SQL 交集部分按 T-SQL 实现)
|
||||
3. 正确处理NULL值(使用IS NULL/IS NOT NULL,而非= NULL)
|
||||
4. 聚合查询必须正确使用GROUP BY
|
||||
5. 多表JOIN要明确ON条件(基于外键关系)
|
||||
6. 避免SELECT *,明确列出所需字段
|
||||
|
||||
**金融/证券业务特殊注意**:
|
||||
- 金额字段通常为DECIMAL类型,注意精度
|
||||
- 日期字段常见:ValueDate(价值日期)、TradeDate(交易日期)、BusinessDate(业务日期)
|
||||
- 状态字段:State(A=Active, D=Deleted, X=SameDayDeleted);SettleStatus(U=Unsettled未结算, S=Settled已结算)等,请参考Schema字段注释中的枚举说明
|
||||
- 币种字段:CurrencyID,多币种查询需注意
|
||||
- 负债/负数含义:许多余额字段负值表示负债(如LoanBalance)
|
||||
|
||||
**常见时间/状态推断指南**(需结合Schema字段注释):
|
||||
- "2024年1月" → `WHERE date_col >= '2024-01-01' AND date_col < '2024-02-01'`
|
||||
- "今天" / "当日" → `WHERE date_col >= CAST(GETDATE() AS DATE) AND date_col < DATEADD(DAY,1,CAST(GETDATE() AS DATE))`
|
||||
- "最近N天" → `WHERE date_col >= DATEADD(DAY, -N, CAST(GETDATE() AS DATE))`
|
||||
- "活跃" → 通常对应 `State = 'A'`(需确认Schema注释)
|
||||
- "未结算" → 通常对应 `SettleStatus = 'U'` 或 `Settled = 0`
|
||||
- "已删除" → 通常对应 `State = 'D'`
|
||||
|
||||
**重要**:以上推断需结合当前Schema中对应字段的实际注释进行调整;若Schema字段注释明确给出了枚举值映射,必须按注释执行。
|
||||
|
||||
**输出格式**:纯SQL字符串,或JSON格式:{"sql": "...", "explanation": "..."}
|
||||
|
||||
---
|
||||
|
||||
**业务级黄金范例(中文问题 ↔ 有用 SQL,版式即标准答案)**:
|
||||
|
||||
用户问题(示例):按对手方列出截至 2026-04-02 的所有未结算交易。
|
||||
|
||||
(日期规则:用户问题里**已写出具体日期**时,在 WHERE 中写入相同字面量,例如 `CashSettleDate <= '2026-04-02'`;**未给出具体日期**但包含时间范围描述(如"2024年1月")时,应合理推断为日期区间,例如 `OrderDate >= '2024-01-01' AND OrderDate < '2024-02-01'`,禁止使用 `'[日期]'` 等占位符。)
|
||||
|
||||
SQL(表名、字段名须与当前 Schema 一致;**以下版式、别名、JOIN/WHERE/GROUP BY/ORDER BY 结构为强制模板**):
|
||||
|
||||
SELECT
|
||||
b.BrokerID,
|
||||
m.Name AS BrokerName,
|
||||
COUNT(*) AS UnsettledTradeCount,
|
||||
SUM(b.SettleQuantity) AS TotalUnsettledQty,
|
||||
SUM(b.SettleAmount) AS TotalUnsettledAmount,
|
||||
MIN(b.BuySell) + '~' + MAX(b.BuySell) AS BuySellRange,
|
||||
MIN(b.CashSettleDate) AS EarliestSettleDate,
|
||||
MAX(b.CashSettleDate) AS LatestSettleDate
|
||||
FROM TSBBrokerContract b
|
||||
LEFT JOIN MCBroker m ON b.BrokerID = m.BrokerID
|
||||
WHERE b.SettleStatus = 'U'
|
||||
AND b.CashSettleDate <= '2026-04-02'
|
||||
GROUP BY b.BrokerID, m.Name
|
||||
ORDER BY TotalUnsettledAmount DESC;
|
||||
|
||||
(说明:`TSBBrokerContract`/`MCBroker`/`BrokerID`/`SettleStatus` 等为**演示用**名称;**禁止**在 Schema 未包含这些对象时照抄。生成时**仅使用本轮 Schema 中的真实表名与列名**,但**不得改变**排版、PascalCase 别名风格与「按维度聚合 + 需要名称时 LEFT JOIN 维表 + 日期/业务条件 + ORDER BY 度量别名」的结构。字符串拼接一律用 T-SQL 的 `+`。)
|
||||
|
||||
---
|
||||
|
||||
**示例(版式与范例一致)**:
|
||||
|
||||
示例1 - 单表查询:
|
||||
问题:查询账户ID为'ACC001'的账户余额
|
||||
Schema: MCAccount(AccountID, Name, AvailableBalance, MarketValue, MarginValue)
|
||||
SQL:
|
||||
SELECT
|
||||
AccountID,
|
||||
Name AS AccountName,
|
||||
AvailableBalance,
|
||||
MarketValue,
|
||||
MarginValue
|
||||
FROM MCAccount
|
||||
WHERE AccountID = 'ACC001';
|
||||
|
||||
示例2 - 多表 INNER JOIN:
|
||||
问题:查询账户'ACC001'持有的所有股票及数量
|
||||
Schema: MCAccount(AccountID), MCAccountInstrument(AccountID, MarketID, InstrumentID, Settled)
|
||||
SQL:
|
||||
SELECT
|
||||
a.AccountID,
|
||||
i.MarketID,
|
||||
i.InstrumentID AS InstrumentCode,
|
||||
i.Settled AS HoldingQty
|
||||
FROM MCAccount a
|
||||
JOIN MCAccountInstrument i ON a.AccountID = i.AccountID
|
||||
WHERE a.AccountID = 'ACC001';
|
||||
|
||||
示例3 - 时间范围推断(关键!):
|
||||
问题:查询2024年1月的总销售额
|
||||
Schema: Sales(OrderID, OrderDate, Amount, ProductID)
|
||||
SQL:
|
||||
SELECT
|
||||
SUM(Amount) AS TotalSales,
|
||||
COUNT(DISTINCT OrderID) AS OrderCount
|
||||
FROM Sales
|
||||
WHERE OrderDate >= '2024-01-01'
|
||||
AND OrderDate < '2024-02-01';
|
||||
|
||||
示例4 - 状态推断(关键!):
|
||||
问题:统计所有活跃账户的数量
|
||||
Schema: MCAccount(AccountID, State, OpenDate) -- State说明: A=Active, D=Deleted, X=SameDayDeleted
|
||||
SQL:
|
||||
SELECT
|
||||
COUNT(*) AS ActiveAccountCount
|
||||
FROM MCAccount
|
||||
WHERE State = 'A';
|
||||
|
||||
示例5 - 聚合 + 排序:
|
||||
问题:按市场统计现金余额总额,从高到低排序
|
||||
Schema: BCAccountMarketCash(AccountID, MarketID, Settled, UnderdueBuy, DueBuy)
|
||||
SQL:
|
||||
SELECT
|
||||
MarketID,
|
||||
SUM(Settled) AS TotalSettled,
|
||||
SUM(UnderdueBuy) AS TotalUnderdueBuy,
|
||||
SUM(DueBuy) AS TotalDueBuy
|
||||
FROM BCAccountMarketCash
|
||||
GROUP BY MarketID
|
||||
ORDER BY TotalSettled DESC;
|
||||
|
||||
示例6 - 日期"今天"的处理:
|
||||
问题:查询今天的交易记录
|
||||
Schema: Ledger(TransactionID, TradeDate, Amount, AccountID)
|
||||
SQL:
|
||||
SELECT
|
||||
TransactionID,
|
||||
TradeDate,
|
||||
Amount,
|
||||
AccountID
|
||||
FROM Ledger
|
||||
WHERE TradeDate >= CAST(GETDATE() AS DATE)
|
||||
AND TradeDate < DATEADD(DAY, 1, CAST(GETDATE() AS DATE));
|
||||
|
||||
示例7 - LEFT JOIN取名称(按维度展示可读名称):
|
||||
问题:按产品类别统计销售额
|
||||
Schema: Sales(ProductID, Amount, SaleDate), Product(ProductID, ProductName, CategoryName)
|
||||
SQL:
|
||||
SELECT
|
||||
p.CategoryName,
|
||||
SUM(s.Amount) AS TotalSales
|
||||
FROM Sales s
|
||||
LEFT JOIN Product p ON s.ProductID = p.ProductID
|
||||
GROUP BY p.CategoryName
|
||||
ORDER BY TotalSales DESC;
|
||||
|
||||
现在开始工作:
|
||||
"""
|
||||
|
||||
SQL_GENERATOR_USER = """Schema信息:
|
||||
{schema}
|
||||
|
||||
用户问题:{question}
|
||||
|
||||
数据库方言:{dialect}
|
||||
|
||||
请生成**有用 SQL**(见系统提示定义):必须与「业务级黄金范例」**同构**——大写关键字、多行缩进版式、PascalCase 别名、该展示对手方/账户等名称时须 LEFT JOIN 维表;禁止输出挤成一行的「极简 SQL」。"""
|
||||
|
||||
|
||||
# ========== Validator Agent Prompt ==========
|
||||
VALIDATOR_SYSTEM = """你是一个严谨的SQL审核员,负责验证SQL语句的正确性和安全性。
|
||||
|
||||
**验证清单**:
|
||||
1. ✓ 语法正确(能够被SQL解析器解析)
|
||||
2. ✓ 所有表名存在于提供的Schema中
|
||||
3. ✓ 所有字段名属于对应的表(检查表.字段格式)
|
||||
4. ✓ JOIN条件完整(每个JOIN都有ON条件,ON条件字段存在且类型兼容)
|
||||
5. ✓ 聚合查询(GROUP BY)包含所有非聚合SELECT字段,或有合理的函数包裹
|
||||
6. ✓ WHERE条件合理(没有明显逻辑错误)
|
||||
7. ✓ HAVING子句只在聚合查询中使用
|
||||
8. ✓ 排序字段存在
|
||||
9. ✓ 无危险操作(DROP/DELETE/UPDATE/INSERT/ALTER/TRUNCATE等,除非明确允许)
|
||||
10. ✓ 无明显性能问题(如全表扫描无WHERE、过度JOIN等)
|
||||
|
||||
**输出格式**(严格JSON):
|
||||
{
|
||||
"valid": true/false,
|
||||
"errors": ["具体错误1", "具体错误2"],
|
||||
"warnings": ["警告信息1"],
|
||||
"suggestions": ["优化建议1"]
|
||||
}
|
||||
|
||||
**错误类型说明**:
|
||||
- "syntax_error": SQL语法错误
|
||||
- "unknown_table": 表名不存在
|
||||
- "unknown_column": 字段名不存在
|
||||
- "missing_join_condition": JOIN缺少ON条件
|
||||
- "missing_group_by": 聚合查询缺少GROUP BY
|
||||
- "dangerous_operation": 危险操作
|
||||
|
||||
---
|
||||
|
||||
现在开始验证:
|
||||
"""
|
||||
|
||||
VALIDATOR_USER = """需要验证的SQL:
|
||||
{sql}
|
||||
|
||||
对应的Schema:
|
||||
{schema}
|
||||
|
||||
请输出验证结果JSON:"""
|
||||
|
||||
EMPTY_RESULT_FEEDBACK_SYSTEM = """你是数据分析助手。用户的自然语言问题已转成 SQL,且在目标库执行成功,但**当前结果在列与行上均为空**(无可用结果集)。
|
||||
请用 2~5 句简洁中文说明可能原因(如条件过严、时间范围无数据、对象不存在等),并**友好引导用户补充**时间、筛选条件、业务对象等,便于下次提问更精确。
|
||||
不要编造 Schema 中不存在的表或字段;不要重复输出整段 SQL;不要输出 JSON 或 Markdown 代码块。"""
|
||||
|
||||
EMPTY_RESULT_FEEDBACK_USER = """用户原始问题:
|
||||
{question}
|
||||
|
||||
已执行的 SQL:
|
||||
{sql}
|
||||
|
||||
相关 Schema(节选):
|
||||
{schema}
|
||||
|
||||
请直接输出给终端用户阅读的说明文字(纯文本)。"""
|
||||
|
||||
|
||||
# ========== Few-Shot 示例 ==========
|
||||
FEW_SHOT_EXAMPLES: Dict[str, str] = {
|
||||
"single_table": """
|
||||
示例:单表查询
|
||||
问题:查询所有状态为Active的账户数量
|
||||
Schema: MCAccount(AccountID, State, OpenDate)
|
||||
SQL:
|
||||
SELECT
|
||||
COUNT(*) AS ActiveAccountCount
|
||||
FROM MCAccount
|
||||
WHERE State = 'A';
|
||||
""",
|
||||
"join_query": """
|
||||
示例:多表JOIN
|
||||
问题:查询账户'ACC001'的持仓信息(包括股票代码、数量、成本)
|
||||
Schema: MCAccount(AccountID), MCAccountInstrument(AccountID, MarketID, InstrumentID, Settled), HCInstrumentClosingPrice(MarketID, InstrumentID, ClosingPrice)
|
||||
SQL:
|
||||
SELECT
|
||||
i.InstrumentID AS InstrumentCode,
|
||||
i.Settled AS HoldingQty,
|
||||
p.ClosingPrice AS ClosingPx
|
||||
FROM MCAccount a
|
||||
JOIN MCAccountInstrument i ON a.AccountID = i.AccountID
|
||||
JOIN HCInstrumentClosingPrice p ON i.MarketID = p.MarketID AND i.InstrumentID = p.InstrumentID
|
||||
WHERE a.AccountID = 'ACC001';
|
||||
""",
|
||||
"aggregation": """
|
||||
示例:聚合统计
|
||||
问题:统计每个市场的现金余额总额
|
||||
Schema: BCAccountMarketCash(AccountID, MarketID, Settled, UnderdueBuy, DueBuy)
|
||||
SQL:
|
||||
SELECT
|
||||
MarketID,
|
||||
SUM(Settled) AS TotalSettled,
|
||||
SUM(UnderdueBuy) AS TotalUnderdueBuy,
|
||||
SUM(DueBuy) AS TotalDueBuy
|
||||
FROM BCAccountMarketCash
|
||||
GROUP BY MarketID;
|
||||
""",
|
||||
"date_filter": """
|
||||
示例:日期筛选
|
||||
问题:查询2024年1月有交易的账户
|
||||
Schema: HCLedgerBalance(ValueDate, LedgerID, Amount), CompanyBusinessDate(BusinessDate)
|
||||
SQL:
|
||||
SELECT DISTINCT
|
||||
LedgerID AS LedgerKey
|
||||
FROM HCLedgerBalance
|
||||
WHERE ValueDate >= '2024-01-01'
|
||||
AND ValueDate < '2024-02-01';
|
||||
""",
|
||||
}
|
||||
|
||||
# ========== NL 英译中(检索 / Text2SQL 前归一化)==========
|
||||
TRANSLATE_NL_TO_ZH_SYSTEM = """你是证券/期货类数据仓库领域的翻译助手。
|
||||
将用户给出的英文(或主要为拉丁字母的)分析需求翻译成**一句简洁的中文自然语言问题**,供后续中文向量检索与 Text2SQL 使用。
|
||||
|
||||
规则:
|
||||
1. 语义忠实,使用业内常用中文表述(如 market value→市值、single holding→单一持仓 等)。
|
||||
2. 保留阿拉伯数字、日期、币种代码、证券代码;「10 million」等与中文习惯一致时可译为「一千万」「1000万」等。
|
||||
3. 若原句中出现明确的英文表名、字段名,保持英文不译。
|
||||
4. **只输出中文问句本身**,不要引号、不要「翻译如下」等前后缀。"""
|
||||
|
||||
TRANSLATE_NL_TO_ZH_USER = """原句:
|
||||
{question}
|
||||
|
||||
仅输出一句中文:"""
|
||||
@@ -0,0 +1,52 @@
|
||||
"""
|
||||
配置管理模块
|
||||
"""
|
||||
|
||||
from pydantic_settings import BaseSettings
|
||||
from typing import Optional, List
|
||||
|
||||
|
||||
class Settings(BaseSettings):
|
||||
"""应用程序配置"""
|
||||
|
||||
# DeepSeek配置
|
||||
deepseek_api_key: str
|
||||
deepseek_base_url: str = "https://api.deepseek.com"
|
||||
model_primary: str = "deepseek-chat"
|
||||
temperature: float = 0.3
|
||||
max_tokens: int = 4096
|
||||
|
||||
# Embedding配置
|
||||
embedding_model_path: str = "./data/models/Qwen3-Embedding-0.6B"
|
||||
vector_dim: int = 2048
|
||||
use_local_embedding: bool = True
|
||||
|
||||
# Schema配置
|
||||
schema_dir: str = "./data/schemas"
|
||||
schema_file: Optional[str] = None # 特定schema文件路径
|
||||
vector_db_path: str = "./data/embeddings/chroma"
|
||||
|
||||
# 验证配置
|
||||
enable_sql_validation: bool = True
|
||||
block_dangerous_sql: bool = True
|
||||
dangerous_keywords: List[str] = ["DROP", "DELETE", "UPDATE", "INSERT", "ALTER", "TRUNCATE"]
|
||||
|
||||
# 重试配置
|
||||
max_retry: int = 2
|
||||
retry_delay: float = 1.0
|
||||
|
||||
# 日志配置
|
||||
log_level: str = "INFO"
|
||||
log_file: Optional[str] = None
|
||||
|
||||
# 业务库(backend/db:execute_sql / search_objects);与 .env 中 database_url 一致
|
||||
database_url: Optional[str] = None
|
||||
sql_max_rows: int = 10000
|
||||
|
||||
class Config:
|
||||
env_file = ".env"
|
||||
case_sensitive = False
|
||||
|
||||
|
||||
# 全局配置实例
|
||||
settings = Settings()
|
||||
@@ -0,0 +1,31 @@
|
||||
"""
|
||||
数据库工具:对齐 DBHub 的 ``execute_sql`` / ``search_objects``。
|
||||
|
||||
使用前请将 ``backend`` 目录加入 ``sys.path``(与 ``api_server.py`` / ``backend/main.py`` 一致),并已在进程内加载根目录 ``.env``。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from db.dbhub_tools import (
|
||||
DbHubTools,
|
||||
dbhub_tools,
|
||||
execute_sql,
|
||||
execute_sql_all,
|
||||
execute_sql_count_only,
|
||||
probe_sql_execution_status,
|
||||
probe_sql_execution_status_ex,
|
||||
search_objects,
|
||||
)
|
||||
from db.engine import get_engine
|
||||
|
||||
__all__ = [
|
||||
"DbHubTools",
|
||||
"dbhub_tools",
|
||||
"execute_sql",
|
||||
"execute_sql_all",
|
||||
"execute_sql_count_only",
|
||||
"get_engine",
|
||||
"probe_sql_execution_status",
|
||||
"probe_sql_execution_status_ex",
|
||||
"search_objects",
|
||||
]
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -0,0 +1,87 @@
|
||||
"""
|
||||
从 DBHub allowed-keywords.ts 等价移植:只读 SQL 判定。
|
||||
参见 dbhub/src/utils/allowed-keywords.ts
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from typing import Literal
|
||||
|
||||
from db.dbhub_sql_parser import ConnectorType, strip_comments_and_strings
|
||||
|
||||
ALLOWED_KEYWORDS: dict[ConnectorType, list[str]] = {
|
||||
"postgres": ["select", "with", "explain", "show"],
|
||||
"mysql": ["select", "with", "explain", "show", "describe", "desc"],
|
||||
"mariadb": ["select", "with", "explain", "show", "describe", "desc"],
|
||||
"sqlite": ["select", "with", "explain", "pragma"],
|
||||
"sqlserver": ["select", "with", "explain", "showplan"],
|
||||
}
|
||||
|
||||
_MUTATING = [
|
||||
"insert",
|
||||
"update",
|
||||
"delete",
|
||||
"drop",
|
||||
"alter",
|
||||
"create",
|
||||
"truncate",
|
||||
"merge",
|
||||
"grant",
|
||||
"revoke",
|
||||
"rename",
|
||||
]
|
||||
_mutating_pattern = re.compile(rf"\b(?:{'|'.join(_MUTATING)})\b", re.IGNORECASE)
|
||||
_mutating_pattern_with_replace = re.compile(
|
||||
rf"\b(?:{'|'.join(_MUTATING)}|replace\s+(?:(?:low_priority|delayed)\s+)?into)\b",
|
||||
re.IGNORECASE,
|
||||
)
|
||||
|
||||
_MUTATING_PATTERNS: dict[ConnectorType, re.Pattern[str]] = {
|
||||
"postgres": _mutating_pattern,
|
||||
"mysql": _mutating_pattern_with_replace,
|
||||
"mariadb": _mutating_pattern_with_replace,
|
||||
"sqlite": _mutating_pattern_with_replace,
|
||||
"sqlserver": _mutating_pattern,
|
||||
}
|
||||
|
||||
_SELECT_INTO_PATTERN = re.compile(r"\bselect\b[\s\S]+\binto\b", re.IGNORECASE)
|
||||
|
||||
_EXPLAIN_ANALYZE_PATTERN = re.compile(
|
||||
r"^explain\s+(?:\([^)]*\banalyze\b(?!\s*(?:=\s*)?(?:false|off|0)\b)[^)]*\)|\banalyze\b(?!\s*(?:=\s*)?(?:false|off|0)\b)(?:\s+verbose\b)?)",
|
||||
re.IGNORECASE,
|
||||
)
|
||||
|
||||
|
||||
def _check_read_only(cleaned_sql: str, connector_type: ConnectorType | str) -> bool:
|
||||
if not cleaned_sql:
|
||||
return False
|
||||
m = re.search(r"\S+", cleaned_sql)
|
||||
first_word = m.group(0) if m else ""
|
||||
keyword_list = ALLOWED_KEYWORDS.get(connector_type, []) # type: ignore[arg-type]
|
||||
if first_word not in keyword_list:
|
||||
return False
|
||||
if first_word == "with":
|
||||
pat = _MUTATING_PATTERNS.get(connector_type, _mutating_pattern) # type: ignore[arg-type]
|
||||
if pat.search(cleaned_sql):
|
||||
return False
|
||||
if first_word in ("select", "with") and _SELECT_INTO_PATTERN.search(cleaned_sql):
|
||||
return False
|
||||
if first_word == "explain":
|
||||
em = _EXPLAIN_ANALYZE_PATTERN.match(cleaned_sql)
|
||||
if em:
|
||||
after_explain = cleaned_sql[em.end() :].strip()
|
||||
if after_explain and not _check_read_only(after_explain, connector_type):
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def is_read_only_sql(sql: str, connector_type: ConnectorType | str) -> bool:
|
||||
"""Check if a SQL query is read-only (DBHub-compatible)."""
|
||||
cleaned = strip_comments_and_strings(sql, connector_type if connector_type in ALLOWED_KEYWORDS else None)
|
||||
cleaned = cleaned.strip().lower()
|
||||
return _check_read_only(cleaned, connector_type)
|
||||
|
||||
|
||||
def allowed_keywords_list(connector_type: ConnectorType | str) -> list[str]:
|
||||
return ALLOWED_KEYWORDS.get(connector_type, []) # type: ignore[arg-type]
|
||||
@@ -0,0 +1,268 @@
|
||||
"""
|
||||
从 DBHub sql-parser.ts 等价移植:按方言剥离注释/字符串、切分语句。
|
||||
参见 dbhub/src/utils/sql-parser.ts
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from typing import Callable, Literal, TypedDict
|
||||
|
||||
ConnectorType = Literal["postgres", "mysql", "mariadb", "sqlite", "sqlserver"]
|
||||
|
||||
|
||||
class _Token(TypedDict):
|
||||
type: int # 0 Plain, 1 Comment, 2 QuotedBlock
|
||||
end: int
|
||||
|
||||
|
||||
_TOKEN_PLAIN = 0
|
||||
_TOKEN_COMMENT = 1
|
||||
_TOKEN_QUOTED = 2
|
||||
|
||||
|
||||
def _plain_token(i: int) -> _Token:
|
||||
return {"type": _TOKEN_PLAIN, "end": i + 1}
|
||||
|
||||
|
||||
def _scan_single_line_comment(sql: str, i: int) -> _Token | None:
|
||||
if i + 1 >= len(sql) or sql[i] != "-" or sql[i + 1] != "-":
|
||||
return None
|
||||
j = i
|
||||
while j < len(sql) and sql[j] != "\n":
|
||||
j += 1
|
||||
return {"type": _TOKEN_COMMENT, "end": j}
|
||||
|
||||
|
||||
def _scan_multi_line_comment(sql: str, i: int) -> _Token | None:
|
||||
if i + 1 >= len(sql) or sql[i] != "/" or sql[i + 1] != "*":
|
||||
return None
|
||||
j = i + 2
|
||||
while j + 1 < len(sql) and not (sql[j] == "*" and sql[j + 1] == "/"):
|
||||
j += 1
|
||||
if j + 1 < len(sql):
|
||||
j += 2
|
||||
return {"type": _TOKEN_COMMENT, "end": j}
|
||||
|
||||
|
||||
def _scan_multi_line_comment_mysql(sql: str, i: int) -> _Token | None:
|
||||
if i + 1 >= len(sql) or sql[i] != "/" or sql[i + 1] != "*":
|
||||
return None
|
||||
nxt = sql[i + 2] if i + 2 < len(sql) else ""
|
||||
nxt2 = sql[i + 3] if i + 3 < len(sql) else ""
|
||||
if nxt == "!" or (nxt == "M" and nxt2 == "!"):
|
||||
return None
|
||||
return _scan_multi_line_comment(sql, i)
|
||||
|
||||
|
||||
def _scan_nested_multi_line_comment(sql: str, i: int) -> _Token | None:
|
||||
if i + 1 >= len(sql) or sql[i] != "/" or sql[i + 1] != "*":
|
||||
return None
|
||||
j = i + 2
|
||||
depth = 1
|
||||
while j < len(sql) and depth > 0:
|
||||
if j + 1 < len(sql) and sql[j] == "/" and sql[j + 1] == "*":
|
||||
depth += 1
|
||||
j += 2
|
||||
elif j + 1 < len(sql) and sql[j] == "*" and sql[j + 1] == "/":
|
||||
depth -= 1
|
||||
j += 2
|
||||
else:
|
||||
j += 1
|
||||
return {"type": _TOKEN_COMMENT, "end": j}
|
||||
|
||||
|
||||
def _scan_single_quoted_string(sql: str, i: int) -> _Token | None:
|
||||
if sql[i] != "'":
|
||||
return None
|
||||
j = i + 1
|
||||
while j < len(sql):
|
||||
if j + 1 < len(sql) and sql[j] == "'" and sql[j + 1] == "'":
|
||||
j += 2
|
||||
elif sql[j] == "'":
|
||||
j += 1
|
||||
break
|
||||
else:
|
||||
j += 1
|
||||
return {"type": _TOKEN_QUOTED, "end": j}
|
||||
|
||||
|
||||
def _scan_double_quoted_string(sql: str, i: int) -> _Token | None:
|
||||
if sql[i] != '"':
|
||||
return None
|
||||
j = i + 1
|
||||
while j < len(sql):
|
||||
if j + 1 < len(sql) and sql[j] == '"' and sql[j + 1] == '"':
|
||||
j += 2
|
||||
elif sql[j] == '"':
|
||||
j += 1
|
||||
break
|
||||
else:
|
||||
j += 1
|
||||
return {"type": _TOKEN_QUOTED, "end": j}
|
||||
|
||||
|
||||
_dollar_quote_open_regex = re.compile(r"^\$([a-zA-Z_]\w*)?\$")
|
||||
|
||||
|
||||
def _scan_dollar_quoted_block(sql: str, i: int) -> _Token | None:
|
||||
if sql[i] != "$":
|
||||
return None
|
||||
nxt = sql[i + 1] if i + 1 < len(sql) else ""
|
||||
if nxt.isdigit():
|
||||
return None
|
||||
remaining = sql[i:]
|
||||
m = _dollar_quote_open_regex.match(remaining)
|
||||
if not m:
|
||||
return None
|
||||
tag = m.group(0)
|
||||
body_start = i + len(tag)
|
||||
close_idx = sql.find(tag, body_start)
|
||||
end = close_idx + len(tag) if close_idx != -1 else len(sql)
|
||||
return {"type": _TOKEN_QUOTED, "end": end}
|
||||
|
||||
|
||||
def _scan_backtick_quoted_identifier(sql: str, i: int) -> _Token | None:
|
||||
if sql[i] != "`":
|
||||
return None
|
||||
j = i + 1
|
||||
while j < len(sql):
|
||||
if j + 1 < len(sql) and sql[j] == "`" and sql[j + 1] == "`":
|
||||
j += 2
|
||||
elif sql[j] == "`":
|
||||
j += 1
|
||||
break
|
||||
else:
|
||||
j += 1
|
||||
return {"type": _TOKEN_QUOTED, "end": j}
|
||||
|
||||
|
||||
def _scan_bracket_quoted_identifier(sql: str, i: int) -> _Token | None:
|
||||
if sql[i] != "[":
|
||||
return None
|
||||
j = i + 1
|
||||
while j < len(sql):
|
||||
if j + 1 < len(sql) and sql[j] == "]" and sql[j + 1] == "]":
|
||||
j += 2
|
||||
elif sql[j] == "]":
|
||||
j += 1
|
||||
break
|
||||
else:
|
||||
j += 1
|
||||
return {"type": _TOKEN_QUOTED, "end": j}
|
||||
|
||||
|
||||
def _scan_token_ansi(sql: str, i: int) -> _Token:
|
||||
return (
|
||||
_scan_single_line_comment(sql, i)
|
||||
or _scan_multi_line_comment(sql, i)
|
||||
or _scan_single_quoted_string(sql, i)
|
||||
or _scan_double_quoted_string(sql, i)
|
||||
or _plain_token(i)
|
||||
)
|
||||
|
||||
|
||||
def _scan_token_postgres(sql: str, i: int) -> _Token:
|
||||
return (
|
||||
_scan_single_line_comment(sql, i)
|
||||
or _scan_nested_multi_line_comment(sql, i)
|
||||
or _scan_single_quoted_string(sql, i)
|
||||
or _scan_double_quoted_string(sql, i)
|
||||
or _scan_dollar_quoted_block(sql, i)
|
||||
or _plain_token(i)
|
||||
)
|
||||
|
||||
|
||||
def _scan_token_mysql(sql: str, i: int) -> _Token:
|
||||
return (
|
||||
_scan_single_line_comment(sql, i)
|
||||
or _scan_multi_line_comment_mysql(sql, i)
|
||||
or _scan_single_quoted_string(sql, i)
|
||||
or _scan_double_quoted_string(sql, i)
|
||||
or _scan_backtick_quoted_identifier(sql, i)
|
||||
or _plain_token(i)
|
||||
)
|
||||
|
||||
|
||||
def _scan_token_sqlite(sql: str, i: int) -> _Token:
|
||||
return (
|
||||
_scan_single_line_comment(sql, i)
|
||||
or _scan_multi_line_comment(sql, i)
|
||||
or _scan_single_quoted_string(sql, i)
|
||||
or _scan_double_quoted_string(sql, i)
|
||||
or _scan_backtick_quoted_identifier(sql, i)
|
||||
or _scan_bracket_quoted_identifier(sql, i)
|
||||
or _plain_token(i)
|
||||
)
|
||||
|
||||
|
||||
def _scan_token_sqlserver(sql: str, i: int) -> _Token:
|
||||
return (
|
||||
_scan_single_line_comment(sql, i)
|
||||
or _scan_multi_line_comment(sql, i)
|
||||
or _scan_single_quoted_string(sql, i)
|
||||
or _scan_double_quoted_string(sql, i)
|
||||
or _scan_bracket_quoted_identifier(sql, i)
|
||||
or _plain_token(i)
|
||||
)
|
||||
|
||||
|
||||
_DIALECT_SCANNERS: dict[ConnectorType, Callable[[str, int], _Token]] = {
|
||||
"postgres": _scan_token_postgres,
|
||||
"mysql": _scan_token_mysql,
|
||||
"mariadb": _scan_token_mysql,
|
||||
"sqlite": _scan_token_sqlite,
|
||||
"sqlserver": _scan_token_sqlserver,
|
||||
}
|
||||
|
||||
|
||||
def _get_scanner(dialect: ConnectorType | None) -> Callable[[str, int], _Token]:
|
||||
if dialect and dialect in _DIALECT_SCANNERS:
|
||||
return _DIALECT_SCANNERS[dialect]
|
||||
return _scan_token_ansi
|
||||
|
||||
|
||||
def strip_comments_and_strings(sql: str, dialect: ConnectorType | None = None) -> str:
|
||||
"""Replace comments, string literals, and dialect-specific quoted blocks with a single space each."""
|
||||
scan_token = _get_scanner(dialect)
|
||||
parts: list[str] = []
|
||||
plain_start = -1
|
||||
i = 0
|
||||
n = len(sql)
|
||||
while i < n:
|
||||
token = scan_token(sql, i)
|
||||
if token["type"] == _TOKEN_PLAIN:
|
||||
if plain_start == -1:
|
||||
plain_start = i
|
||||
else:
|
||||
if plain_start != -1:
|
||||
parts.append(sql[plain_start:i])
|
||||
plain_start = -1
|
||||
parts.append(" ")
|
||||
i = token["end"]
|
||||
if plain_start != -1:
|
||||
parts.append(sql[plain_start:])
|
||||
return "".join(parts)
|
||||
|
||||
|
||||
def split_sql_statements(sql: str, dialect: ConnectorType | None = None) -> list[str]:
|
||||
"""Split SQL into individual statements, handling semicolons inside quoted contexts."""
|
||||
scan_token = _get_scanner(dialect)
|
||||
statements: list[str] = []
|
||||
stmt_start = 0
|
||||
i = 0
|
||||
n = len(sql)
|
||||
while i < n:
|
||||
if sql[i] == ";":
|
||||
trimmed = sql[stmt_start:i].strip()
|
||||
if trimmed:
|
||||
statements.append(trimmed)
|
||||
stmt_start = i + 1
|
||||
i += 1
|
||||
continue
|
||||
token = scan_token(sql, i)
|
||||
i = token["end"]
|
||||
trimmed = sql[stmt_start:].strip()
|
||||
if trimmed:
|
||||
statements.append(trimmed)
|
||||
return statements
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,46 @@
|
||||
"""
|
||||
SQLAlchemy 引擎:使用环境变量 ``database_url`` / ``DATABASE_URL``(与项目根目录 ``.env`` 一致)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import os
|
||||
from typing import Optional
|
||||
|
||||
from sqlalchemy import create_engine
|
||||
from sqlalchemy.engine import Engine
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_engine: Engine | None = None
|
||||
|
||||
|
||||
def _database_url_from_env() -> str:
|
||||
for key in ("database_url", "DATABASE_URL"):
|
||||
v = os.getenv(key)
|
||||
if v is not None and str(v).strip():
|
||||
return str(v).strip()
|
||||
return ""
|
||||
|
||||
|
||||
def get_engine(*, url: Optional[str] = None, reset: bool = False) -> Engine:
|
||||
"""
|
||||
返回默认业务库引擎(单例)。未配置 ``database_url`` 时抛出 ``ValueError``。
|
||||
|
||||
:param url: 若传入,则忽略单例并为此 URL 新建引擎(便于测试)。
|
||||
:param reset: 为 True 时丢弃已缓存的单例,下次再按环境变量创建。
|
||||
"""
|
||||
global _engine
|
||||
if reset:
|
||||
_engine = None
|
||||
if url is not None:
|
||||
return create_engine(url, pool_pre_ping=True)
|
||||
if _engine is not None:
|
||||
return _engine
|
||||
u = _database_url_from_env()
|
||||
if not u:
|
||||
raise ValueError("database_url 未配置")
|
||||
_engine = create_engine(u, pool_pre_ping=True)
|
||||
logger.info("SQLAlchemy engine initialized from database_url")
|
||||
return _engine
|
||||
@@ -0,0 +1 @@
|
||||
# llm 包初始化
|
||||
Binary file not shown.
Binary file not shown.
@@ -0,0 +1,361 @@
|
||||
"""
|
||||
DeepSeek API 客户端封装
|
||||
支持同步/异步调用,与CAMEL AI兼容
|
||||
"""
|
||||
|
||||
import os
|
||||
import json
|
||||
import logging
|
||||
from typing import Dict, List, Optional, Any, Union
|
||||
from dataclasses import dataclass, field
|
||||
from openai import OpenAI, AsyncOpenAI
|
||||
from openai.types.chat import ChatCompletion, ChatCompletionMessage
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class DeepSeekConfig:
|
||||
"""DeepSeek API配置"""
|
||||
api_key: str
|
||||
base_url: str = "https://api.deepseek.com"
|
||||
model_name: str = "deepseek-chat"
|
||||
temperature: float = 0.3
|
||||
max_tokens: int = 4096
|
||||
top_p: float = 0.9
|
||||
frequency_penalty: float = 0.0
|
||||
presence_penalty: float = 0.0
|
||||
timeout: float = 60.0
|
||||
extra_headers: Optional[Dict[str, str]] = None
|
||||
|
||||
|
||||
class DeepSeekClient:
|
||||
"""
|
||||
DeepSeek API 客户端(同步)
|
||||
|
||||
使用 OpenAI 兼容接口调用 DeepSeek 模型
|
||||
"""
|
||||
|
||||
def __init__(self, config: DeepSeekConfig):
|
||||
self.config = config
|
||||
|
||||
# 初始化OpenAI客户端(DeepSeek兼容OpenAI协议)
|
||||
self.client = OpenAI(
|
||||
api_key=config.api_key,
|
||||
base_url=config.base_url,
|
||||
timeout=config.timeout,
|
||||
)
|
||||
|
||||
logger.info(
|
||||
f"[OK] DeepSeekClient初始化: model={config.model_name}, "
|
||||
f"base_url={config.base_url}"
|
||||
)
|
||||
|
||||
def chat(
|
||||
self,
|
||||
messages: List[Dict[str, str]],
|
||||
**kwargs
|
||||
) -> ChatCompletionMessage:
|
||||
"""
|
||||
发送聊天请求
|
||||
|
||||
Args:
|
||||
messages: 消息列表,格式为 [{"role": "user", "content": "..."}, ...]
|
||||
**kwargs: 覆盖默认参数的额外参数
|
||||
|
||||
Returns:
|
||||
ChatCompletionMessage对象(包含.role和.content属性)
|
||||
"""
|
||||
# 合并参数
|
||||
params = {
|
||||
"model": self.config.model_name,
|
||||
"messages": messages,
|
||||
"temperature": kwargs.get("temperature", self.config.temperature),
|
||||
"max_tokens": kwargs.get("max_tokens", self.config.max_tokens),
|
||||
"top_p": kwargs.get("top_p", self.config.top_p),
|
||||
"frequency_penalty": kwargs.get(
|
||||
"frequency_penalty", self.config.frequency_penalty
|
||||
),
|
||||
"presence_penalty": kwargs.get(
|
||||
"presence_penalty", self.config.presence_penalty
|
||||
),
|
||||
"timeout": kwargs.get("timeout", self.config.timeout),
|
||||
}
|
||||
|
||||
if self.config.extra_headers:
|
||||
params["extra_headers"] = self.config.extra_headers
|
||||
|
||||
try:
|
||||
response: ChatCompletion = self.client.chat.completions.create(**params)
|
||||
message = response.choices[0].message
|
||||
|
||||
# 记录使用情况
|
||||
usage = response.usage
|
||||
logger.debug(
|
||||
f"DeepSeek调用完成: "
|
||||
f"prompt_tokens={usage.prompt_tokens}, "
|
||||
f"completion_tokens={usage.completion_tokens}, "
|
||||
f"total_tokens={usage.total_tokens}"
|
||||
)
|
||||
|
||||
return message
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"DeepSeek API调用失败: {e}")
|
||||
raise
|
||||
|
||||
def chat_with_json(
|
||||
self,
|
||||
messages: List[Dict[str, str]],
|
||||
**kwargs
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
发送聊天请求并期望JSON格式返回
|
||||
|
||||
Args:
|
||||
messages: 消息列表
|
||||
**kwargs: 额外参数
|
||||
|
||||
Returns:
|
||||
解析后的JSON字典
|
||||
"""
|
||||
message = self.chat(messages, **kwargs)
|
||||
content = message.content.strip()
|
||||
|
||||
try:
|
||||
# 尝试解析JSON(可能包含markdown代码块)
|
||||
if "```json" in content:
|
||||
# 提取```json```之间的内容
|
||||
start = content.find("```json") + 7
|
||||
end = content.find("```", start)
|
||||
content = content[start:end].strip()
|
||||
elif "```" in content:
|
||||
# 仅```包裹
|
||||
start = content.find("```") + 3
|
||||
end = content.find("```", start)
|
||||
content = content[start:end].strip()
|
||||
|
||||
return json.loads(content)
|
||||
except json.JSONDecodeError as e:
|
||||
logger.warning(f"JSON解析失败,返回原始内容: {e}")
|
||||
# 勿仅用 raw_content 判失败:空串时下游 `not raw.get("raw_content")` 会误判为成功
|
||||
return {"_json_decode_failed": True, "raw_content": content}
|
||||
|
||||
def generate_sql(
|
||||
self,
|
||||
prompt: str,
|
||||
schema: str,
|
||||
dialect: str = "tsql",
|
||||
**kwargs
|
||||
) -> str:
|
||||
"""
|
||||
便捷方法:生成SQL
|
||||
|
||||
Args:
|
||||
prompt: 用户问题
|
||||
schema: Schema描述
|
||||
dialect: SQL方言(默认 T-SQL / SQL Server)
|
||||
**kwargs: 额外参数
|
||||
|
||||
Returns:
|
||||
SQL字符串
|
||||
"""
|
||||
from config.prompts import SQL_GENERATOR_SYSTEM, SQL_GENERATOR_USER
|
||||
|
||||
messages = [
|
||||
{"role": "system", "content": SQL_GENERATOR_SYSTEM},
|
||||
{
|
||||
"role": "user",
|
||||
"content": SQL_GENERATOR_USER.format(
|
||||
schema=schema,
|
||||
question=prompt,
|
||||
dialect=dialect
|
||||
)
|
||||
}
|
||||
]
|
||||
|
||||
kwargs.setdefault("temperature", 0.0)
|
||||
kwargs.setdefault("top_p", 1.0)
|
||||
response = self.chat(messages, **kwargs)
|
||||
content = response.content.strip()
|
||||
|
||||
# 提取SQL(移除可能的markdown代码块标记)
|
||||
if "```sql" in content:
|
||||
start = content.find("```sql") + 6
|
||||
end = content.find("```", start)
|
||||
content = content[start:end].strip()
|
||||
elif "```" in content:
|
||||
# 通用代码块
|
||||
start = content.find("```") + 3
|
||||
end = content.find("```", start)
|
||||
content = content[start:end].strip()
|
||||
|
||||
return content
|
||||
|
||||
def validate_sql(
|
||||
self,
|
||||
sql: str,
|
||||
schema: str,
|
||||
**kwargs
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
便捷方法:验证SQL
|
||||
|
||||
Args:
|
||||
sql: SQL语句
|
||||
schema: Schema描述
|
||||
**kwargs: 额外参数
|
||||
|
||||
Returns:
|
||||
验证结果字典
|
||||
"""
|
||||
from config.prompts import VALIDATOR_SYSTEM, VALIDATOR_USER
|
||||
|
||||
messages = [
|
||||
{"role": "system", "content": VALIDATOR_SYSTEM},
|
||||
{
|
||||
"role": "user",
|
||||
"content": VALIDATOR_USER.format(sql=sql, schema=schema)
|
||||
}
|
||||
]
|
||||
|
||||
kwargs.setdefault("temperature", 0.0)
|
||||
kwargs.setdefault("top_p", 1.0)
|
||||
return self.chat_with_json(messages, **kwargs)
|
||||
|
||||
def empty_result_user_feedback(
|
||||
self,
|
||||
question: str,
|
||||
sql: str,
|
||||
schema: str,
|
||||
*,
|
||||
max_schema_chars: int = 8000,
|
||||
**kwargs: Any,
|
||||
) -> str:
|
||||
"""
|
||||
库探针为 0(执行成功但结果行数为 0)时,生成面向用户的中文补充说明,引导用户完善问题。
|
||||
"""
|
||||
from config.prompts import EMPTY_RESULT_FEEDBACK_SYSTEM, EMPTY_RESULT_FEEDBACK_USER
|
||||
|
||||
schema_snip = (schema or "")[:max_schema_chars]
|
||||
messages = [
|
||||
{"role": "system", "content": EMPTY_RESULT_FEEDBACK_SYSTEM},
|
||||
{
|
||||
"role": "user",
|
||||
"content": EMPTY_RESULT_FEEDBACK_USER.format(
|
||||
question=question or "(无)",
|
||||
sql=sql,
|
||||
schema=schema_snip,
|
||||
),
|
||||
},
|
||||
]
|
||||
msg = self.chat(messages, temperature=0.4, max_tokens=512, **kwargs)
|
||||
text = (msg.content or "").strip()
|
||||
return text
|
||||
|
||||
def select_tables(
|
||||
self,
|
||||
question: str,
|
||||
table_list: str,
|
||||
**kwargs
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
便捷方法:选择相关表
|
||||
|
||||
Args:
|
||||
question: 用户问题
|
||||
table_list: 可用表列表字符串
|
||||
**kwargs: 额外参数
|
||||
|
||||
Returns:
|
||||
包含relevant_tables和reasoning的字典
|
||||
"""
|
||||
from config.prompts import SCHEMA_LINKER_SYSTEM, SCHEMA_LINKER_USER
|
||||
|
||||
messages = [
|
||||
{"role": "system", "content": SCHEMA_LINKER_SYSTEM},
|
||||
{
|
||||
"role": "user",
|
||||
"content": SCHEMA_LINKER_USER.format(
|
||||
question=question,
|
||||
table_list=table_list,
|
||||
),
|
||||
},
|
||||
]
|
||||
|
||||
# 选表为结构化决策:默认 temperature=0,避免同一问题多次选不同表/SQL 上下文
|
||||
kwargs.setdefault("temperature", 0.0)
|
||||
kwargs.setdefault("top_p", 1.0)
|
||||
return self.chat_with_json(messages, **kwargs)
|
||||
|
||||
def translate_nl_question_to_zh(self, question: str) -> str:
|
||||
"""
|
||||
将主要为英文的自然语言分析问题译为中文,便于与中文 Schema 注释 / 向量索引对齐。
|
||||
"""
|
||||
from config.prompts import TRANSLATE_NL_TO_ZH_SYSTEM, TRANSLATE_NL_TO_ZH_USER
|
||||
|
||||
q = (question or "").strip()
|
||||
if not q:
|
||||
return ""
|
||||
messages = [
|
||||
{"role": "system", "content": TRANSLATE_NL_TO_ZH_SYSTEM},
|
||||
{"role": "user", "content": TRANSLATE_NL_TO_ZH_USER.format(question=q)},
|
||||
]
|
||||
msg = self.chat(messages, temperature=0.0, top_p=1.0, max_tokens=512)
|
||||
text = (msg.content or "").strip()
|
||||
# 只取首行,避免模型附加说明
|
||||
line = text.splitlines()[0].strip() if text else ""
|
||||
return line.strip("「」\"'“”")
|
||||
|
||||
|
||||
class AsyncDeepSeekClient:
|
||||
"""
|
||||
DeepSeek API 客户端(异步)
|
||||
|
||||
注意:CAMEL AI当前版本主要支持同步Agent,
|
||||
异步客户端适用于自定义异步流程。
|
||||
"""
|
||||
|
||||
def __init__(self, config: DeepSeekConfig):
|
||||
self.config = config
|
||||
self.client = AsyncOpenAI(
|
||||
api_key=config.api_key,
|
||||
base_url=config.base_url,
|
||||
timeout=config.timeout,
|
||||
)
|
||||
|
||||
async def chat(self, messages: List[Dict[str, str]], **kwargs) -> ChatCompletionMessage:
|
||||
params = {
|
||||
"model": self.config.model_name,
|
||||
"messages": messages,
|
||||
"temperature": kwargs.get("temperature", self.config.temperature),
|
||||
"max_tokens": kwargs.get("max_tokens", self.config.max_tokens),
|
||||
}
|
||||
response = await self.client.chat.completions.create(**params)
|
||||
return response.choices[0].message
|
||||
|
||||
|
||||
# 便捷工厂函数
|
||||
def create_deepseek_client(
|
||||
api_key: Optional[str] = None,
|
||||
**kwargs
|
||||
) -> DeepSeekClient:
|
||||
"""
|
||||
创建DeepSeek客户端
|
||||
|
||||
Args:
|
||||
api_key: API密钥(可从环境变量DEEPSEEK_API_KEY读取)
|
||||
**kwargs: 覆盖默认配置的参数
|
||||
|
||||
Returns:
|
||||
DeepSeekClient实例
|
||||
"""
|
||||
if api_key is None:
|
||||
api_key = os.getenv("DEEPSEEK_API_KEY")
|
||||
if not api_key:
|
||||
raise ValueError(
|
||||
"未提供api_key且环境变量DEEPSEEK_API_KEY未设置。"
|
||||
)
|
||||
|
||||
config = DeepSeekConfig(api_key=api_key, **kwargs)
|
||||
return DeepSeekClient(config)
|
||||
+411
@@ -0,0 +1,411 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
Text2SQL 多智能体系统 - CLI 入口
|
||||
|
||||
用法:在项目根目录执行 python backend/main.py
|
||||
"""
|
||||
|
||||
import os
|
||||
import sys
|
||||
import argparse
|
||||
import logging
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
|
||||
from dotenv import load_dotenv
|
||||
|
||||
# 从仓库根目录运行 python backend/main.py 时,将 backend 加入模块搜索路径
|
||||
_backend_dir = Path(__file__).resolve().parent
|
||||
if str(_backend_dir) not in sys.path:
|
||||
sys.path.insert(0, str(_backend_dir))
|
||||
|
||||
# 配置日志
|
||||
logging.basicConfig(
|
||||
level=logging.INFO,
|
||||
format='%(asctime)s [%(levelname)s] %(name)s: %(message)s',
|
||||
datefmt='%Y-%m-%d %H:%M:%S'
|
||||
)
|
||||
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
|
||||
|
||||
|
||||
def _load_project_env():
|
||||
"""加载项目根目录 .env,供后续 os.getenv 使用。"""
|
||||
load_dotenv(_repo_root() / ".env")
|
||||
|
||||
|
||||
def _embedding_model_path() -> str:
|
||||
return os.getenv("EMBEDDING_MODEL_PATH", _DEFAULT_EMBEDDING_PATH).strip()
|
||||
|
||||
|
||||
def resolve_sql_dialect(name: str) -> str:
|
||||
"""CLI / 配置中的方言别名统一为 sqlglot 方言名(SQL Server -> tsql)。"""
|
||||
n = (name or "sqlserver").lower().strip()
|
||||
if n in ("sqlserver", "mssql"):
|
||||
return "tsql"
|
||||
return n
|
||||
|
||||
|
||||
def setup_environment():
|
||||
"""环境检查"""
|
||||
_load_project_env()
|
||||
use_local_emb = os.getenv("USE_LOCAL_EMBEDDING", "true").strip().lower() in (
|
||||
"1", "true", "yes", "on",
|
||||
)
|
||||
if use_local_emb:
|
||||
# 与 .env 中 EMBEDDING_MODEL_PATH 及 utils.embedding 一致
|
||||
model_path = Path(_embedding_model_path())
|
||||
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_*(百炼)")
|
||||
return False
|
||||
else:
|
||||
ms_key = os.getenv("MODELSCOPE_API_KEY", "").strip()
|
||||
oa_key = os.getenv("OPENAI_API_KEY", "").strip()
|
||||
oa_key_ok = oa_key and not (
|
||||
oa_key.startswith("http://") or oa_key.startswith("https://")
|
||||
)
|
||||
if ms_key:
|
||||
pass # ModelScope:BASE_URL / MODEL 有默认值,仅需 KEY
|
||||
elif oa_key_ok:
|
||||
if not os.getenv("OPENAI_EMBEDDING_MODEL", "").strip():
|
||||
logger.warning("USE_LOCAL_EMBEDDING=false 但未设置 OPENAI_EMBEDDING_MODEL")
|
||||
return False
|
||||
else:
|
||||
if not os.getenv("DASHSCOPE_API_KEY", "").strip():
|
||||
logger.warning(
|
||||
"USE_LOCAL_EMBEDDING=false 但未设置 MODELSCOPE_API_KEY、"
|
||||
"OPENAI_API_KEY+OPENAI_EMBEDDING_MODEL 或 DASHSCOPE_API_KEY"
|
||||
)
|
||||
return False
|
||||
base = (
|
||||
os.getenv("DASHSCOPE_BASE_URL") or os.getenv("DASHSCOPE_base_url", "")
|
||||
).strip()
|
||||
if not base:
|
||||
logger.warning("未设置 DASHSCOPE_BASE_URL(或 DASHSCOPE_base_url)")
|
||||
return False
|
||||
if not os.getenv("DASHSCOPE_MODEL", "").strip():
|
||||
logger.warning("未设置 DASHSCOPE_MODEL")
|
||||
return False
|
||||
|
||||
# 检查Schema文件
|
||||
schema_path = Path("./data/schemas/G3SB_MCDataDictionary_table_structure.json")
|
||||
if not schema_path.exists():
|
||||
logger.warning(f"Schema文件不存在: {schema_path}")
|
||||
logger.info("请将G3SB Schema文件放置在 ./data/schemas/ 目录")
|
||||
return False
|
||||
|
||||
# 检查API Key
|
||||
if not os.getenv("DEEPSEEK_API_KEY"):
|
||||
logger.warning("环境变量 DEEPSEEK_API_KEY 未设置")
|
||||
logger.info("请在 .env 文件中配置,或 export DEEPSEEK_API_KEY=your_key")
|
||||
return False
|
||||
|
||||
return True
|
||||
|
||||
|
||||
def _default_g3sb_meta_path(structure_path: str) -> Optional[str]:
|
||||
"""若存在与 table_structure 同名的 table_meta 文件则返回其路径。"""
|
||||
p = Path(structure_path)
|
||||
if "table_structure" not in p.name:
|
||||
return None
|
||||
cand = p.parent / p.name.replace("table_structure", "table_meta")
|
||||
return str(cand) if cand.is_file() else None
|
||||
|
||||
|
||||
def load_schema(schema_path: str, schema_meta_path: Optional[str] = None):
|
||||
"""加载 Schema;G3SB structure JSON 会自动尝试配对 table_meta(可用 --schema-meta 指定)。"""
|
||||
from schema.manager import SchemaManager
|
||||
|
||||
meta = (
|
||||
schema_meta_path
|
||||
if schema_meta_path is not None
|
||||
else _default_g3sb_meta_path(schema_path)
|
||||
)
|
||||
logger.info(f"加载Schema: {schema_path}")
|
||||
if meta:
|
||||
logger.info(f" 表注释(meta): {meta}")
|
||||
|
||||
schema_mgr = SchemaManager.load_from_json(
|
||||
schema_path, g3sb_meta_path=meta
|
||||
)
|
||||
|
||||
stats = schema_mgr.get_statistics()
|
||||
logger.info(
|
||||
f"[OK] Schema加载完成: {stats['database']}, "
|
||||
f"共{stats['total_tables']}张表, {stats['total_columns']}个字段"
|
||||
)
|
||||
|
||||
return schema_mgr
|
||||
|
||||
|
||||
def create_orchestrator(schema_mgr, args):
|
||||
"""创建编排器"""
|
||||
from agents.orchestrator import Text2SQLOrchestrator
|
||||
from llm.deepseek_client import DeepSeekConfig
|
||||
|
||||
translate_en = os.getenv("TRANSLATE_EN_TO_ZH", "true").strip().lower() not in (
|
||||
"0",
|
||||
"false",
|
||||
"no",
|
||||
"off",
|
||||
)
|
||||
if getattr(args, "no_translate_en", False):
|
||||
translate_en = False
|
||||
|
||||
api_key = (args.api_key or os.getenv("DEEPSEEK_API_KEY") or "").strip()
|
||||
if not api_key:
|
||||
raise ValueError(
|
||||
"未配置 DeepSeek API Key:请在 .env 中设置 DEEPSEEK_API_KEY,"
|
||||
"或使用命令行参数 --api-key"
|
||||
)
|
||||
base_url = (os.getenv("DEEPSEEK_BASE_URL") or "https://api.deepseek.com").strip()
|
||||
|
||||
config = DeepSeekConfig(
|
||||
api_key=api_key,
|
||||
base_url=base_url,
|
||||
model_name=args.model or "deepseek-chat",
|
||||
temperature=args.temperature,
|
||||
max_tokens=args.max_tokens,
|
||||
)
|
||||
|
||||
orchestrator = Text2SQLOrchestrator(
|
||||
schema_manager=schema_mgr,
|
||||
deepseek_config=config,
|
||||
embedding_model_path=args.embedding_model,
|
||||
vector_db_path=args.vector_db,
|
||||
max_retry=args.max_retry,
|
||||
use_vector_search=not args.no_vector_search,
|
||||
# Few-shot配置
|
||||
fewshot_enabled=not args.no_fewshot,
|
||||
fewshot_top_k=args.fewshot_top_k,
|
||||
fewshot_min_rating=args.fewshot_min_rating,
|
||||
translate_english_to_zh=translate_en,
|
||||
)
|
||||
|
||||
return orchestrator
|
||||
|
||||
|
||||
def single_query(orchestrator, question: str, dialect: str = "tsql"):
|
||||
"""单次查询"""
|
||||
import time
|
||||
|
||||
from agents.orchestrator import GenerationResult
|
||||
from utils.dialog_classifier import DialogIntent, classify_dialog
|
||||
|
||||
logger.info(f"[Q] 问题: {question}")
|
||||
|
||||
classified = classify_dialog(question)
|
||||
if classified.intent == DialogIntent.CONVERSATION:
|
||||
reply = classified.reply_suggestion or ""
|
||||
logger.info("[Q] 意图: conversation(跳过 SQL 生成)")
|
||||
print("\n" + "=" * 60)
|
||||
print("对话 / 非查询输入(未触发 SQL 生成)")
|
||||
print("=" * 60)
|
||||
print(reply)
|
||||
print(f"\n使用表: []")
|
||||
return GenerationResult(
|
||||
sql="",
|
||||
valid=False,
|
||||
errors=[],
|
||||
warnings=[],
|
||||
tables_used=[],
|
||||
attempts=0,
|
||||
metadata={"dialog_intent": DialogIntent.CONVERSATION.value},
|
||||
)
|
||||
|
||||
start = time.time()
|
||||
result = orchestrator.generate(
|
||||
question=question,
|
||||
dialect=dialect,
|
||||
top_k_candidates=20
|
||||
)
|
||||
elapsed = time.time() - start
|
||||
|
||||
print("\n" + "=" * 60)
|
||||
print("生成结果:")
|
||||
print("=" * 60)
|
||||
|
||||
if result.valid:
|
||||
print(f"[OK] SQL (耗时 {elapsed:.2f}s, 尝试 {result.attempts} 次):\n")
|
||||
print(result.sql)
|
||||
else:
|
||||
print(f"[FAIL] 生成失败 (尝试 {result.attempts} 次)")
|
||||
for err in result.errors:
|
||||
print(f" - {err}")
|
||||
|
||||
if result.warnings:
|
||||
print("\n[WARN] 警告:")
|
||||
for w in result.warnings:
|
||||
print(f" - {w}")
|
||||
|
||||
dbe = result.metadata.get("db_empty_feedback")
|
||||
if result.valid and dbe:
|
||||
print("\n[DB 探针 0 — 无数据行] 说明:")
|
||||
print(dbe)
|
||||
|
||||
print(f"\n使用表: {result.tables_used}")
|
||||
|
||||
return result
|
||||
|
||||
|
||||
def interactive_mode(orchestrator, dialect: str = "tsql"):
|
||||
"""交互式模式"""
|
||||
print("\n" + "=" * 60)
|
||||
print("Text2SQL 交互模式(输入 'quit' 或 'exit' 退出)")
|
||||
print("提示:请描述业务数据查询需求;寒暄或「你是谁」等会由对话分类处理,不生成 SQL。")
|
||||
print("=" * 60 + "\n")
|
||||
|
||||
while True:
|
||||
try:
|
||||
question = input("❓ 请输入业务查询问题: ").strip()
|
||||
if question.lower() in ('quit', 'exit', 'q'):
|
||||
print("再见!")
|
||||
break
|
||||
|
||||
if not question:
|
||||
continue
|
||||
|
||||
result = single_query(orchestrator, question, dialect)
|
||||
print()
|
||||
|
||||
except KeyboardInterrupt:
|
||||
print("\n再见!")
|
||||
break
|
||||
except Exception as e:
|
||||
logger.error(f"查询失败: {e}")
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Text2SQL:运行后进入交互式自然语言生成 SQL(python backend/main.py)"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--schema", "-s",
|
||||
default="./data/schemas/G3SB_MCDataDictionary_table_structure.json",
|
||||
help="Schema文件路径(默认: ./data/schemas/G3SB_MCDataDictionary_table_structure.json)"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--schema-meta",
|
||||
default=None,
|
||||
help="G3SB table_meta.json;默认自动使用同目录下文件名含 table_meta 的配对文件",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--dialect", "-d",
|
||||
default=os.getenv("TEXT2SQL_DIALECT", "sqlserver").strip(),
|
||||
choices=["mysql", "postgresql", "sqlite", "tsql", "sqlserver", "mssql"],
|
||||
help="SQL方言(默认: sqlserver / T-SQL;可用环境变量 TEXT2SQL_DIALECT 覆盖)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--api-key",
|
||||
help="DeepSeek API Key(默认从DEEPSEEK_API_KEY环境变量读取)"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--model", "-m",
|
||||
default="deepseek-chat",
|
||||
help="DeepSeek模型名称(默认: deepseek-chat)"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--temperature", "-t",
|
||||
type=float,
|
||||
default=0.3,
|
||||
help="生成温度(默认: 0.3)"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--max-tokens",
|
||||
type=int,
|
||||
default=4096,
|
||||
help="最大token数(默认: 4096)"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--max-retry",
|
||||
type=int,
|
||||
default=2,
|
||||
help="最大重试次数(默认: 2)"
|
||||
)
|
||||
_load_project_env()
|
||||
parser.add_argument(
|
||||
"--embedding-model",
|
||||
default=_embedding_model_path(),
|
||||
help="Embedding模型路径(默认来自环境变量 EMBEDDING_MODEL_PATH)"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--vector-db",
|
||||
default=os.getenv("VECTOR_DB_PATH", "./data/embeddings/chroma").strip(),
|
||||
help="向量数据库路径(默认来自环境变量 VECTOR_DB_PATH)"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--no-vector-search",
|
||||
action="store_true",
|
||||
help="禁用向量检索(使用所有表)"
|
||||
)
|
||||
# Few-shot 配置
|
||||
parser.add_argument(
|
||||
"--no-fewshot",
|
||||
action="store_true",
|
||||
help="禁用few-shot示例增强"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--fewshot-top-k",
|
||||
type=int,
|
||||
default=int(os.getenv("FEWSHOT_TOP_K", "3")),
|
||||
help="每次使用的few-shot示例数量(默认: 3)"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--fewshot-min-rating",
|
||||
type=int,
|
||||
default=int(os.getenv("FEWSHOT_MIN_RATING", "7")),
|
||||
help="few-shot示例最低评分(默认: 7)"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--verbose", "-v",
|
||||
action="store_true",
|
||||
help="详细日志"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--no-translate-en",
|
||||
action="store_true",
|
||||
help="关闭英文问句自动译为中文(默认开启;也可用 TRANSLATE_EN_TO_ZH=false)",
|
||||
)
|
||||
|
||||
args = parser.parse_args()
|
||||
args.dialect = resolve_sql_dialect(args.dialect)
|
||||
|
||||
# 日志级别
|
||||
if args.verbose:
|
||||
logging.getLogger().setLevel(logging.DEBUG)
|
||||
|
||||
# 环境检查
|
||||
if not setup_environment():
|
||||
sys.exit(1)
|
||||
|
||||
# 加载Schema
|
||||
try:
|
||||
schema_mgr = load_schema(args.schema, args.schema_meta)
|
||||
except Exception as e:
|
||||
logger.error(f"Schema加载失败: {e}")
|
||||
sys.exit(1)
|
||||
|
||||
# 创建Orchestrator
|
||||
try:
|
||||
orchestrator = create_orchestrator(schema_mgr, args)
|
||||
except Exception as e:
|
||||
logger.error(f"Orchestrator创建失败: {e}")
|
||||
sys.exit(1)
|
||||
|
||||
interactive_mode(orchestrator, args.dialect)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,287 @@
|
||||
"""
|
||||
进程内 NL 附属数据(会话 / 消息 / 收藏),字段与 web 端 ApiEnvelope、SessionRow、MessageRow 对齐。
|
||||
仅用于本地或 demo 联调,重启后数据丢失。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import uuid
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any, Dict, List, Optional, Tuple
|
||||
|
||||
|
||||
def utc_ts() -> str:
|
||||
return datetime.now(timezone.utc).strftime("%Y-%m-%d %H:%M:%S")
|
||||
|
||||
|
||||
class LiteNlStore:
|
||||
def __init__(self) -> None:
|
||||
self._lock = asyncio.Lock()
|
||||
self._sessions: Dict[Tuple[str, str], List[Dict[str, Any]]] = {}
|
||||
self._messages: Dict[Tuple[str, str, str], List[Dict[str, Any]]] = {}
|
||||
self._msg_counters: Dict[Tuple[str, str, str], int] = {}
|
||||
self._fav: Dict[Tuple[str, str], Dict[str, List[Any]]] = {}
|
||||
|
||||
@staticmethod
|
||||
def _visitor_key(vid: Optional[str]) -> str:
|
||||
return (vid or "").strip()
|
||||
|
||||
def _scope(self, user_id: Optional[str], visitor_biz_id: Optional[str]) -> Tuple[str, str]:
|
||||
uid = (user_id or "anonymous").strip() or "anonymous"
|
||||
return uid, self._visitor_key(visitor_biz_id)
|
||||
|
||||
async def create_session(
|
||||
self,
|
||||
user_id: Optional[str],
|
||||
visitor_biz_id: Optional[str],
|
||||
title: Optional[str],
|
||||
) -> Dict[str, Any]:
|
||||
async with self._lock:
|
||||
sk = self._scope(user_id, visitor_biz_id)
|
||||
sid = str(uuid.uuid4())
|
||||
now = utc_ts()
|
||||
row: Dict[str, Any] = {
|
||||
"session_id": sid,
|
||||
"title": (title or "").strip() or None,
|
||||
"user_id": sk[0],
|
||||
"visitor_biz_id": visitor_biz_id or None,
|
||||
"created_at": now,
|
||||
"updated_at": now,
|
||||
}
|
||||
self._sessions.setdefault(sk, []).insert(0, row)
|
||||
self._messages[(*sk, sid)] = []
|
||||
self._msg_counters[(*sk, sid)] = 0
|
||||
return {"session_id": sid, "title": row["title"]}
|
||||
|
||||
async def list_sessions(
|
||||
self,
|
||||
user_id: Optional[str],
|
||||
visitor_biz_id: Optional[str],
|
||||
limit: int,
|
||||
offset: int,
|
||||
) -> Dict[str, Any]:
|
||||
async with self._lock:
|
||||
sk = self._scope(user_id, visitor_biz_id)
|
||||
items = list(self._sessions.get(sk, []))
|
||||
items.sort(key=lambda r: str(r.get("updated_at") or ""), reverse=True)
|
||||
sl = items[offset : offset + limit]
|
||||
return {"items": sl, "limit": limit, "offset": offset}
|
||||
|
||||
def _bump_session(self, sk: Tuple[str, str], session_id: str, user_first_line: str) -> None:
|
||||
now = utc_ts()
|
||||
for r in self._sessions.get(sk, []):
|
||||
if r.get("session_id") != session_id:
|
||||
continue
|
||||
r["updated_at"] = now
|
||||
if not (r.get("title") or "").strip() and user_first_line.strip():
|
||||
r["title"] = user_first_line.strip()[:80]
|
||||
break
|
||||
|
||||
async def append_exchange(
|
||||
self,
|
||||
user_id: Optional[str],
|
||||
visitor_biz_id: Optional[str],
|
||||
session_id: str,
|
||||
user_text: str,
|
||||
assistant_content_json: str,
|
||||
) -> None:
|
||||
async with self._lock:
|
||||
sk = self._scope(user_id, visitor_biz_id)
|
||||
sk3 = (*sk, session_id)
|
||||
if sk3 not in self._messages:
|
||||
self._messages[sk3] = []
|
||||
self._msg_counters[sk3] = 0
|
||||
now = utc_ts()
|
||||
|
||||
def next_id() -> int:
|
||||
n = self._msg_counters.get(sk3, 0) + 1
|
||||
self._msg_counters[sk3] = n
|
||||
return n
|
||||
|
||||
self._messages[sk3].append(
|
||||
{
|
||||
"id": next_id(),
|
||||
"role": "user",
|
||||
"content": user_text,
|
||||
"created_at": now,
|
||||
"llm_total_tokens": None,
|
||||
}
|
||||
)
|
||||
self._messages[sk3].append(
|
||||
{
|
||||
"id": next_id(),
|
||||
"role": "assistant",
|
||||
"content": assistant_content_json,
|
||||
"created_at": now,
|
||||
"llm_total_tokens": None,
|
||||
}
|
||||
)
|
||||
self._bump_session(sk, session_id, user_text)
|
||||
|
||||
async def get_messages(
|
||||
self,
|
||||
user_id: Optional[str],
|
||||
visitor_biz_id: Optional[str],
|
||||
session_id: str,
|
||||
limit: int,
|
||||
offset: int,
|
||||
) -> Optional[Dict[str, Any]]:
|
||||
async with self._lock:
|
||||
sk = self._scope(user_id, visitor_biz_id)
|
||||
sk3 = (*sk, session_id)
|
||||
if sk3 not in self._messages:
|
||||
found = any(
|
||||
r.get("session_id") == session_id for r in self._sessions.get(sk, [])
|
||||
)
|
||||
if not found:
|
||||
return None
|
||||
self._messages[sk3] = []
|
||||
self._msg_counters[sk3] = 0
|
||||
rows = list(self._messages[sk3])
|
||||
rows = rows[offset : offset + limit]
|
||||
return {"session_id": session_id, "items": rows, "limit": limit, "offset": offset}
|
||||
|
||||
async def update_session_title(
|
||||
self,
|
||||
user_id: Optional[str],
|
||||
visitor_biz_id: Optional[str],
|
||||
session_id: str,
|
||||
title: Optional[str],
|
||||
) -> Optional[Dict[str, Any]]:
|
||||
async with self._lock:
|
||||
sk = self._scope(user_id, visitor_biz_id)
|
||||
for r in self._sessions.get(sk, []):
|
||||
if r.get("session_id") == session_id:
|
||||
r["title"] = title
|
||||
r["updated_at"] = utc_ts()
|
||||
return {"session_id": session_id, "title": r.get("title")}
|
||||
return None
|
||||
|
||||
async def delete_session(
|
||||
self,
|
||||
user_id: Optional[str],
|
||||
visitor_biz_id: Optional[str],
|
||||
session_id: str,
|
||||
) -> Optional[Dict[str, Any]]:
|
||||
async with self._lock:
|
||||
sk = self._scope(user_id, visitor_biz_id)
|
||||
lst = self._sessions.get(sk, [])
|
||||
sk3 = (*sk, session_id)
|
||||
n = len(self._messages.get(sk3, []))
|
||||
new_lst = [x for x in lst if x.get("session_id") != session_id]
|
||||
if len(new_lst) == len(lst):
|
||||
return None
|
||||
self._sessions[sk] = new_lst
|
||||
self._messages.pop(sk3, None)
|
||||
self._msg_counters.pop(sk3, None)
|
||||
return {"session_id": session_id, "deleted_messages": n}
|
||||
|
||||
async def patch_message(
|
||||
self,
|
||||
user_id: Optional[str],
|
||||
visitor_biz_id: Optional[str],
|
||||
session_id: str,
|
||||
message_id: int,
|
||||
content: str,
|
||||
) -> Optional[Dict[str, Any]]:
|
||||
async with self._lock:
|
||||
sk3 = (*self._scope(user_id, visitor_biz_id), session_id)
|
||||
for m in self._messages.get(sk3, []):
|
||||
if m.get("id") == message_id:
|
||||
m["content"] = content
|
||||
m["created_at"] = utc_ts()
|
||||
return {
|
||||
"id": message_id,
|
||||
"session_id": session_id,
|
||||
"role": m.get("role"),
|
||||
"content": content,
|
||||
"created_at": m.get("created_at"),
|
||||
"llm_total_tokens": m.get("llm_total_tokens"),
|
||||
}
|
||||
return None
|
||||
|
||||
async def get_favorites_grouped(
|
||||
self, user_id: Optional[str], visitor_biz_id: Optional[str]
|
||||
) -> Dict[str, List[Any]]:
|
||||
async with self._lock:
|
||||
sk = self._scope(user_id, visitor_biz_id)
|
||||
g = self._fav.get(sk) or {"sql": [], "function": [], "report": []}
|
||||
return {"sql": list(g["sql"]), "function": list(g["function"]), "report": list(g["report"])}
|
||||
|
||||
async def add_favorite(
|
||||
self, user_id: Optional[str], visitor_biz_id: Optional[str], body: Dict[str, Any]
|
||||
) -> Any:
|
||||
async with self._lock:
|
||||
sk = self._scope(user_id, visitor_biz_id)
|
||||
g = self._fav.setdefault(sk, {"sql": [], "function": [], "report": []})
|
||||
fav_type = str(body.get("fav_type") or "sql")
|
||||
fid = f"{uuid.uuid4().hex[:12]}"
|
||||
if fav_type == "sql":
|
||||
row = {
|
||||
"id": fid,
|
||||
"name": str(body.get("name") or ""),
|
||||
"desc": str(body.get("desc") or ""),
|
||||
"sql": str(body.get("sql") or ""),
|
||||
}
|
||||
if body.get("sql_explain"):
|
||||
row["sql_explain"] = str(body["sql_explain"])
|
||||
g["sql"].insert(0, row)
|
||||
return row
|
||||
if fav_type == "function":
|
||||
row = {
|
||||
"id": fid,
|
||||
"name": str(body.get("name") or ""),
|
||||
"desc": str(body.get("desc") or ""),
|
||||
"path": str(body.get("path") or ""),
|
||||
}
|
||||
g["function"].insert(0, row)
|
||||
return row
|
||||
row = {
|
||||
"id": fid,
|
||||
"name": str(body.get("name") or ""),
|
||||
"desc": str(body.get("desc") or ""),
|
||||
"reportPath": str(body.get("reportPath") or ""),
|
||||
"params": str(body.get("params") or ""),
|
||||
}
|
||||
g["report"].insert(0, row)
|
||||
return row
|
||||
|
||||
async def patch_favorite(
|
||||
self,
|
||||
user_id: Optional[str],
|
||||
visitor_biz_id: Optional[str],
|
||||
fav_id: str,
|
||||
patch: Dict[str, Any],
|
||||
) -> bool:
|
||||
async with self._lock:
|
||||
sk = self._scope(user_id, visitor_biz_id)
|
||||
g = self._fav.get(sk)
|
||||
if not g:
|
||||
return False
|
||||
for bucket in ("sql", "function", "report"):
|
||||
lst = g.get(bucket, [])
|
||||
for i, it in enumerate(lst):
|
||||
if str(it.get("id")) == fav_id:
|
||||
lst[i] = {**it, **patch}
|
||||
return True
|
||||
return False
|
||||
|
||||
async def delete_favorite(
|
||||
self, user_id: Optional[str], visitor_biz_id: Optional[str], fav_id: str
|
||||
) -> bool:
|
||||
async with self._lock:
|
||||
sk = self._scope(user_id, visitor_biz_id)
|
||||
g = self._fav.get(sk)
|
||||
if not g:
|
||||
return False
|
||||
removed = False
|
||||
for bucket in ("sql", "function", "report"):
|
||||
before = len(g.get(bucket, []))
|
||||
g[bucket] = [x for x in g.get(bucket, []) if str(x.get("id")) != fav_id]
|
||||
if len(g[bucket]) < before:
|
||||
removed = True
|
||||
return removed
|
||||
|
||||
|
||||
lite_nl_store = LiteNlStore()
|
||||
@@ -0,0 +1 @@
|
||||
# schema 包初始化
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -0,0 +1,229 @@
|
||||
"""
|
||||
Schema向量索引构建器 - 基于ChromaDB
|
||||
"""
|
||||
|
||||
import chromadb
|
||||
from chromadb.config import Settings as ChromaSettings
|
||||
from typing import Any, List, Dict, Optional
|
||||
import logging
|
||||
from pathlib import Path
|
||||
|
||||
from schema.manager import SchemaManager
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class SchemaIndexer:
|
||||
"""
|
||||
Schema向量索引器
|
||||
|
||||
功能:
|
||||
1. 为表结构构建向量索引
|
||||
2. 基于问题的表检索
|
||||
3. 持久化存储和加载
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
embedder: Any,
|
||||
persist_dir: str = "./data/embeddings",
|
||||
collection_name: str = "schema_tables",
|
||||
):
|
||||
"""
|
||||
初始化索引器
|
||||
|
||||
Args:
|
||||
embedder: Embedding模型实例
|
||||
persist_dir: 向量数据库持久化目录
|
||||
collection_name: 集合名称
|
||||
"""
|
||||
self.embedder = embedder
|
||||
self.persist_dir = Path(persist_dir)
|
||||
self.persist_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
# 初始化ChromaDB客户端
|
||||
self.client = chromadb.PersistentClient(
|
||||
path=str(self.persist_dir),
|
||||
settings=ChromaSettings(anonymized_telemetry=False),
|
||||
)
|
||||
|
||||
# 获取或创建集合
|
||||
self.collection = self.client.get_or_create_collection(
|
||||
name=collection_name,
|
||||
metadata={"hnsw:space": "cosine"}, # 使用余弦相似度
|
||||
)
|
||||
|
||||
logger.info(f"[OK] 初始化SchemaIndexer: collection={collection_name}")
|
||||
|
||||
def build_index(
|
||||
self,
|
||||
schema_manager: SchemaManager,
|
||||
batch_size: int = 32,
|
||||
force_rebuild: bool = False,
|
||||
) -> bool:
|
||||
"""
|
||||
为Schema构建向量索引
|
||||
|
||||
Args:
|
||||
schema_manager: SchemaManager实例
|
||||
batch_size: 批处理大小
|
||||
force_rebuild: 是否强制重建(默认为False,增量更新)
|
||||
|
||||
Returns:
|
||||
True=成功,False=已存在且未强制重建
|
||||
"""
|
||||
existing_ids = set(self.collection.get()["ids"]) if self.collection.count() > 0 else set()
|
||||
|
||||
if not force_rebuild and existing_ids:
|
||||
logger.info(f"索引已存在({len(existing_ids)}条记录),跳过构建")
|
||||
return False
|
||||
|
||||
if force_rebuild and existing_ids:
|
||||
logger.info(f"强制重建索引,删除{len(existing_ids)}条旧记录")
|
||||
self.collection.delete()
|
||||
|
||||
# 准备数据
|
||||
table_texts = []
|
||||
table_ids = []
|
||||
table_metadatas = []
|
||||
|
||||
for table_info in schema_manager.get_table_for_embedding():
|
||||
table_texts.append(table_info["text"])
|
||||
table_ids.append(table_info["name"])
|
||||
table_metadatas.append({
|
||||
"table_name": table_info["name"],
|
||||
"column_count": len(schema_manager.get_table(table_info["name"]).columns),
|
||||
})
|
||||
|
||||
# 批量计算embedding
|
||||
logger.info(f"计算{len(table_texts)}张表的embedding...")
|
||||
embeddings = self.embedder.encode(table_texts, batch_size=batch_size)
|
||||
|
||||
# 存入ChromaDB
|
||||
self.collection.add(
|
||||
embeddings=embeddings.tolist(),
|
||||
documents=table_texts,
|
||||
metadatas=table_metadatas,
|
||||
ids=table_ids,
|
||||
)
|
||||
|
||||
logger.info(f"[OK] 索引构建完成,共{len(table_ids)}张表")
|
||||
return True
|
||||
|
||||
def search(
|
||||
self,
|
||||
query: str,
|
||||
top_k: int = 20,
|
||||
score_threshold: float = 0.3,
|
||||
) -> List[Dict]:
|
||||
"""
|
||||
检索相关表
|
||||
|
||||
Args:
|
||||
query: 查询文本(用户问题)
|
||||
top_k: 返回前K个结果
|
||||
score_threshold: 相似度阈值(0-1),低于此值的结果会被过滤
|
||||
|
||||
Returns:
|
||||
检索结果列表,按相似度降序排列
|
||||
[
|
||||
{
|
||||
"table_name": "表名",
|
||||
"score": 0.95,
|
||||
"document": "表描述文本",
|
||||
"metadata": {...}
|
||||
},
|
||||
...
|
||||
]
|
||||
"""
|
||||
# 编码查询文本
|
||||
query_embedding = self.embedder.encode([query])
|
||||
|
||||
# 执行检索
|
||||
results = self.collection.query(
|
||||
query_embeddings=query_embedding.tolist(),
|
||||
n_results=min(top_k, self.collection.count()),
|
||||
)
|
||||
|
||||
# 格式化结果
|
||||
formatted = []
|
||||
if results["ids"] and len(results["ids"][0]) > 0:
|
||||
for idx, (table_id, distance, metadata, document) in enumerate(zip(
|
||||
results["ids"][0],
|
||||
results["distances"][0],
|
||||
results["metadatas"][0],
|
||||
results["documents"][0],
|
||||
)):
|
||||
score = 1 - distance # 余弦相似度(已归一化,转换为0-1,越大越相似)
|
||||
if score >= score_threshold:
|
||||
formatted.append({
|
||||
"table_name": table_id,
|
||||
"score": float(score),
|
||||
"document": document,
|
||||
"metadata": metadata,
|
||||
"rank": idx + 1,
|
||||
})
|
||||
|
||||
logger.debug(f"检索 '{query[:50]}...' -> 找到{len(formatted)}个相关表(阈值={score_threshold})")
|
||||
return formatted
|
||||
|
||||
def search_by_table_names(self, table_names: List[str]) -> List[Dict]:
|
||||
"""
|
||||
直接通过表名获取表信息
|
||||
|
||||
Args:
|
||||
table_names: 表名列表
|
||||
|
||||
Returns:
|
||||
表信息列表
|
||||
"""
|
||||
results = []
|
||||
for name in table_names:
|
||||
try:
|
||||
result = self.collection.get(ids=[name], include=["metadatas", "documents"])
|
||||
if result["ids"]:
|
||||
results.append({
|
||||
"table_name": name,
|
||||
"score": 1.0,
|
||||
"document": result["documents"][0],
|
||||
"metadata": result["metadatas"][0],
|
||||
})
|
||||
except Exception as e:
|
||||
logger.warning(f"获取表信息失败 {name}: {e}")
|
||||
|
||||
return results
|
||||
|
||||
def get_all_tables(self) -> List[str]:
|
||||
"""获取索引中的所有表名"""
|
||||
result = self.collection.get()
|
||||
return result["ids"] if result["ids"] else []
|
||||
|
||||
def delete_table(self, table_name: str):
|
||||
"""从索引中删除表"""
|
||||
self.collection.delete(ids=[table_name])
|
||||
logger.info(f"已删除表索引: {table_name}")
|
||||
|
||||
def clear(self):
|
||||
"""清空索引"""
|
||||
self.collection.delete()
|
||||
logger.info("索引已清空")
|
||||
|
||||
def count(self) -> int:
|
||||
"""获取索引中的表数量"""
|
||||
return self.collection.count()
|
||||
|
||||
# ========== 统计与调试 ==========
|
||||
|
||||
def get_statistics(self) -> Dict:
|
||||
"""获取索引统计信息"""
|
||||
count = self.collection.count()
|
||||
result = self.collection.get(include=["metadatas"])
|
||||
|
||||
total_columns = sum(m.get("column_count", 0) for m in result["metadatas"])
|
||||
|
||||
return {
|
||||
"indexed_tables": count,
|
||||
"total_columns": total_columns,
|
||||
"avg_columns": total_columns / count if count > 0 else 0,
|
||||
"persist_dir": str(self.persist_dir),
|
||||
}
|
||||
@@ -0,0 +1,319 @@
|
||||
"""
|
||||
Schema加载器 - 解析JSON/DDL格式的数据库Schema
|
||||
"""
|
||||
|
||||
import json
|
||||
import re
|
||||
from pathlib import Path
|
||||
from typing import List, Dict, Optional, Tuple
|
||||
import logging
|
||||
|
||||
from .models import Table, Column, ForeignKey, DatabaseSchema
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class SchemaLoader:
|
||||
"""Schema加载器 - 支持JSON和DDL格式"""
|
||||
|
||||
def __init__(self, schema_dir: str = None):
|
||||
"""
|
||||
初始化Schema加载器
|
||||
|
||||
Args:
|
||||
schema_dir: Schema文件目录路径
|
||||
"""
|
||||
self.schema_dir = Path(schema_dir) if schema_dir else None
|
||||
|
||||
def load_from_json(
|
||||
self, json_path: str, g3sb_meta_path: Optional[str] = None
|
||||
) -> DatabaseSchema:
|
||||
"""
|
||||
从JSON文件加载Schema
|
||||
|
||||
JSON格式示例:
|
||||
{
|
||||
"database": "db_name",
|
||||
"tables": [
|
||||
{
|
||||
"name": "table1",
|
||||
"comment": "表描述",
|
||||
"columns": [
|
||||
{"name": "col1", "type": "INT", "comment": "...", "nullable": false}
|
||||
],
|
||||
"primary_keys": ["col1"],
|
||||
"foreign_keys": [
|
||||
{"columns": ["col2"], "ref_table": "table2", "ref_columns": ["col1"]}
|
||||
]
|
||||
}
|
||||
]
|
||||
}
|
||||
|
||||
亦支持 G3SB table_structure.json:顶层为 ``schemas`` 字典(与 ``tables`` 数组二选一
|
||||
时优先使用非空的 ``schemas``)。可选 ``g3sb_meta_path`` 指向 table_meta.json
|
||||
以合并表级注释。
|
||||
|
||||
Args:
|
||||
json_path: JSON文件路径
|
||||
g3sb_meta_path: 可选,G3SB table_meta.json(含 tables.{表名}.comment)
|
||||
|
||||
Returns:
|
||||
DatabaseSchema对象
|
||||
"""
|
||||
path = Path(json_path)
|
||||
if not path.exists():
|
||||
raise FileNotFoundError(f"Schema文件不存在: {json_path}")
|
||||
|
||||
with open(path, 'r', encoding='utf-8') as f:
|
||||
data = json.load(f)
|
||||
|
||||
# 提取数据库名
|
||||
db_name = data.get("database", path.stem)
|
||||
|
||||
schemas_map = data.get("schemas")
|
||||
tables_data = data.get("tables")
|
||||
|
||||
# 优先非空 schemas(G3SB 结构字典),避免误把其它 truthy 的 tables 键当好列表解析
|
||||
if isinstance(schemas_map, dict) and schemas_map:
|
||||
table_meta: Dict = {}
|
||||
if g3sb_meta_path:
|
||||
mp = Path(g3sb_meta_path)
|
||||
if mp.is_file():
|
||||
with open(mp, 'r', encoding='utf-8') as mf:
|
||||
meta_data = json.load(mf)
|
||||
table_meta = meta_data.get("tables", {}) or {}
|
||||
else:
|
||||
logger.warning(
|
||||
"G3SB meta 文件不存在,将跳过表注释: %s", g3sb_meta_path
|
||||
)
|
||||
|
||||
tables = []
|
||||
for table_name, structure_str in schemas_map.items():
|
||||
if not isinstance(structure_str, str):
|
||||
continue
|
||||
comment = None
|
||||
if table_meta and isinstance(table_meta.get(table_name), dict):
|
||||
comment = table_meta[table_name].get("comment")
|
||||
columns = self._parse_g3sb_table_structure(structure_str)
|
||||
tables.append(
|
||||
Table(
|
||||
name=table_name,
|
||||
comment=comment,
|
||||
columns=columns,
|
||||
primary_keys=self._extract_primary_keys(columns),
|
||||
foreign_keys=[],
|
||||
)
|
||||
)
|
||||
logger.info(
|
||||
"[OK] 加载Schema完成(G3SB schemas): %s, 共%d张表", db_name, len(tables)
|
||||
)
|
||||
return DatabaseSchema(name=db_name, tables=tables)
|
||||
|
||||
if isinstance(tables_data, list) and tables_data:
|
||||
tables = [self._parse_table(table_data) for table_data in tables_data]
|
||||
logger.info(f"[OK] 加载Schema完成: {db_name}, 共{len(tables)}张表")
|
||||
return DatabaseSchema(name=db_name, tables=tables)
|
||||
|
||||
tables = []
|
||||
logger.info(f"[OK] 加载Schema完成: {db_name}, 共{len(tables)}张表")
|
||||
return DatabaseSchema(name=db_name, tables=tables)
|
||||
|
||||
def load_from_g3sb_format(self, meta_path: str, structure_path: str) -> DatabaseSchema:
|
||||
"""
|
||||
加载G3SB系统的数据字典格式(两个JSON文件)
|
||||
|
||||
Args:
|
||||
meta_path: table_meta.json路径(表注释)
|
||||
structure_path: table_structure.json路径(表结构)
|
||||
|
||||
Returns:
|
||||
DatabaseSchema对象
|
||||
"""
|
||||
# 1. 加载表元数据(表注释)
|
||||
with open(meta_path, 'r', encoding='utf-8') as f:
|
||||
meta_data = json.load(f)
|
||||
table_meta = meta_data.get("tables", {})
|
||||
|
||||
# 2. 加载表结构
|
||||
with open(structure_path, 'r', encoding='utf-8') as f:
|
||||
structure_data = json.load(f)
|
||||
table_structures = structure_data.get("schemas", {})
|
||||
|
||||
# 3. 解析所有表
|
||||
tables = []
|
||||
for table_name, structure_str in table_structures.items():
|
||||
comment = table_meta.get(table_name, {}).get("comment")
|
||||
|
||||
# 解析表结构字符串
|
||||
# 格式: "TABLE TableName (col1:TYPE -- comment, col2:TYPE, ...)"
|
||||
columns = self._parse_g3sb_table_structure(structure_str)
|
||||
|
||||
table = Table(
|
||||
name=table_name,
|
||||
comment=comment,
|
||||
columns=columns,
|
||||
primary_keys=self._extract_primary_keys(columns),
|
||||
foreign_keys=[] # G3SB格式没有外键信息,后续可补充
|
||||
)
|
||||
tables.append(table)
|
||||
|
||||
logger.info(f"[OK] 加载G3SB Schema完成: 共{len(tables)}张表")
|
||||
return DatabaseSchema(name="G3SB_DB", tables=tables)
|
||||
|
||||
def _parse_table(self, data: Dict) -> Table:
|
||||
"""解析单表JSON数据"""
|
||||
columns = []
|
||||
for col_data in data.get("columns", []):
|
||||
col = Column(
|
||||
name=col_data["name"],
|
||||
data_type=col_data["type"],
|
||||
comment=col_data.get("comment"),
|
||||
nullable=col_data.get("nullable", True),
|
||||
is_primary_key=col_data.get("is_primary_key", False),
|
||||
)
|
||||
columns.append(col)
|
||||
|
||||
# 外键解析
|
||||
foreign_keys = []
|
||||
for fk_data in data.get("foreign_keys", []):
|
||||
fk = ForeignKey(
|
||||
columns=fk_data["columns"],
|
||||
ref_table=fk_data["ref_table"],
|
||||
ref_columns=fk_data["ref_columns"],
|
||||
)
|
||||
foreign_keys.append(fk)
|
||||
|
||||
return Table(
|
||||
name=data["name"],
|
||||
comment=data.get("comment"),
|
||||
columns=columns,
|
||||
primary_keys=data.get("primary_keys", []),
|
||||
foreign_keys=foreign_keys,
|
||||
)
|
||||
|
||||
def _parse_g3sb_table_structure(self, structure_str: str) -> List[Column]:
|
||||
"""
|
||||
解析G3SB格式的表结构字符串
|
||||
|
||||
示例:
|
||||
"TABLE BCAccountCash (AccountID:NCHAR, RegionID:NCHAR, CurrencyID:NCHAR, Settled:DECIMAL -- Settled balance, ...)"
|
||||
"""
|
||||
# 提取括号内的内容
|
||||
match = re.search(r'\((.*)\)', structure_str)
|
||||
if not match:
|
||||
logger.warning(f"无法解析表结构: {structure_str[:100]}")
|
||||
return []
|
||||
|
||||
inner = match.group(1)
|
||||
|
||||
columns = []
|
||||
# 按逗号分割字段(注意注释中可能包含逗号)
|
||||
parts = self._split_columns(inner)
|
||||
|
||||
for part in parts:
|
||||
part = part.strip()
|
||||
if not part:
|
||||
continue
|
||||
|
||||
# 解析字段定义:name:TYPE [-- comment]
|
||||
# 支持格式: "FieldName:DATATYPE" 或 "FieldName:DATATYPE -- comment"
|
||||
col_match = re.match(r'^(\w+)\s*:\s*([A-Za-z0-9()]+)', part)
|
||||
if not col_match:
|
||||
continue
|
||||
|
||||
col_name = col_match.group(1).strip()
|
||||
col_type = col_match.group(2).strip()
|
||||
|
||||
# 提取注释(ASCII -- 或 G3SB 常用的 Unicode 长破折号 — U+2014)
|
||||
comment = None
|
||||
comment_match = re.search(r'(?:--|\u2014)\s*(.+)', part)
|
||||
if comment_match:
|
||||
comment = comment_match.group(1).strip()
|
||||
|
||||
# 判断是否可为空(通常有默认值或未标注NOT NULL即为NULL)
|
||||
nullable = True # G3SB格式默认允许NULL
|
||||
|
||||
column = Column(
|
||||
name=col_name,
|
||||
data_type=col_type,
|
||||
comment=comment,
|
||||
nullable=nullable,
|
||||
)
|
||||
columns.append(column)
|
||||
|
||||
return columns
|
||||
|
||||
def _split_columns(self, inner: str) -> List[str]:
|
||||
"""
|
||||
按字段边界拆分(G3SB 注释里常有英文逗号,且用 — 而非 --)。
|
||||
|
||||
仅在「后面紧跟 标识符: 」的逗号处切分,这样注释内的逗号不会误拆列。
|
||||
"""
|
||||
if not inner or not inner.strip():
|
||||
return []
|
||||
# 下一列以 Name:TYPE 开头;避免在括号嵌套里误匹配可再收紧(当前 G3SB 类型无顶层逗号)
|
||||
parts = re.split(r",\s*(?=\w+\s*:)", inner)
|
||||
return [p.strip() for p in parts if p.strip()]
|
||||
|
||||
def _extract_primary_keys(self, columns: List[Column]) -> List[str]:
|
||||
"""从字段列表中提取主键(简单启发式:字段名包含ID或明确标记)"""
|
||||
pk_candidates = []
|
||||
for col in columns:
|
||||
# 简单规则:字段名以ID结尾,或名称包含key/id
|
||||
if col.name.upper().endswith('ID') or 'KEY' in col.name.upper():
|
||||
pk_candidates.append(col.name)
|
||||
return pk_candidates[:1] # 暂时只取一个主键(简化)
|
||||
|
||||
def load_all_schemas(self) -> List[DatabaseSchema]:
|
||||
"""
|
||||
加载schema_dir下的所有Schema文件
|
||||
|
||||
Returns:
|
||||
DatabaseSchema列表
|
||||
"""
|
||||
if not self.schema_dir:
|
||||
raise ValueError("未指定schema_dir")
|
||||
|
||||
schemas = []
|
||||
for json_file in self.schema_dir.glob("*.json"):
|
||||
try:
|
||||
schema = self.load_from_json(str(json_file))
|
||||
schemas.append(schema)
|
||||
except Exception as e:
|
||||
logger.error(f"加载Schema失败 {json_file}: {e}")
|
||||
|
||||
logger.info(f"[OK] 共加载{len(schemas)}个Schema")
|
||||
return schemas
|
||||
|
||||
|
||||
# 便捷函数
|
||||
def load_schema_from_g3sb(meta_path: str, structure_path: str) -> DatabaseSchema:
|
||||
"""
|
||||
从G3SB格式加载Schema的便捷函数
|
||||
|
||||
Args:
|
||||
meta_path: table_meta.json路径
|
||||
structure_path: table_structure.json路径
|
||||
|
||||
Returns:
|
||||
DatabaseSchema对象
|
||||
"""
|
||||
loader = SchemaLoader()
|
||||
return loader.load_from_g3sb_format(meta_path, structure_path)
|
||||
|
||||
|
||||
def load_schema_from_json(
|
||||
json_path: str, g3sb_meta_path: Optional[str] = None
|
||||
) -> DatabaseSchema:
|
||||
"""
|
||||
从标准JSON加载Schema的便捷函数
|
||||
|
||||
Args:
|
||||
json_path: JSON文件路径
|
||||
g3sb_meta_path: 可选,G3SB table_meta.json
|
||||
|
||||
Returns:
|
||||
DatabaseSchema对象
|
||||
"""
|
||||
loader = SchemaLoader()
|
||||
return loader.load_from_json(json_path, g3sb_meta_path=g3sb_meta_path)
|
||||
@@ -0,0 +1,315 @@
|
||||
"""
|
||||
Schema管理器 - 管理数据库Schema的加载、索引和检索
|
||||
"""
|
||||
|
||||
import json
|
||||
import hashlib
|
||||
from pathlib import Path
|
||||
from typing import List, Dict, Optional, Tuple
|
||||
import logging
|
||||
|
||||
from .models import DatabaseSchema, Table
|
||||
from .loader import SchemaLoader, load_schema_from_json, load_schema_from_g3sb
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class SchemaManager:
|
||||
"""
|
||||
Schema管理器
|
||||
|
||||
功能:
|
||||
1. 加载和解析Schema文件
|
||||
2. 管理表关系
|
||||
3. 提供表检索接口
|
||||
4. 生成紧凑的Schema描述
|
||||
"""
|
||||
|
||||
def __init__(self, schema: DatabaseSchema):
|
||||
"""
|
||||
初始化Schema管理器
|
||||
|
||||
Args:
|
||||
schema: DatabaseSchema对象
|
||||
"""
|
||||
self.schema = schema
|
||||
self._table_dict = {tbl.name: tbl for tbl in schema.tables}
|
||||
self._embedding_index = None # 向量索引(延迟初始化)
|
||||
self._cache = {}
|
||||
|
||||
@classmethod
|
||||
def load_from_json(
|
||||
cls, json_path: str, g3sb_meta_path: Optional[str] = None
|
||||
) -> "SchemaManager":
|
||||
"""
|
||||
从JSON文件加载Schema
|
||||
|
||||
Args:
|
||||
json_path: JSON文件路径
|
||||
g3sb_meta_path: 可选,G3SB table_meta.json(与 table_structure 配对)
|
||||
|
||||
Returns:
|
||||
SchemaManager实例
|
||||
"""
|
||||
schema = load_schema_from_json(json_path, g3sb_meta_path=g3sb_meta_path)
|
||||
return cls(schema)
|
||||
|
||||
@classmethod
|
||||
def load_from_g3sb(cls, meta_path: str, structure_path: str) -> "SchemaManager":
|
||||
"""
|
||||
从G3SB格式加载Schema
|
||||
|
||||
Args:
|
||||
meta_path: table_meta.json路径
|
||||
structure_path: table_structure.json路径
|
||||
|
||||
Returns:
|
||||
SchemaManager实例
|
||||
"""
|
||||
schema = load_schema_from_g3sb(meta_path, structure_path)
|
||||
return cls(schema)
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, data: Dict) -> "SchemaManager":
|
||||
"""
|
||||
从字典创建SchemaManager
|
||||
|
||||
Args:
|
||||
data: 包含database和tables的字典
|
||||
|
||||
Returns:
|
||||
SchemaManager实例
|
||||
"""
|
||||
loader = SchemaLoader()
|
||||
schema = loader._parse_table(data) if isinstance(data, dict) else None
|
||||
if not schema:
|
||||
raise ValueError("Invalid schema data")
|
||||
return cls(schema)
|
||||
|
||||
# ========== 查询接口 ==========
|
||||
|
||||
def get_table(self, table_name: str) -> Optional[Table]:
|
||||
"""获取指定表"""
|
||||
return self._table_dict.get(table_name)
|
||||
|
||||
def list_tables(self) -> List[str]:
|
||||
"""获取所有表名"""
|
||||
return list(self._table_dict.keys())
|
||||
|
||||
def get_tables(self) -> List[Table]:
|
||||
"""获取所有Table对象"""
|
||||
return list(self._table_dict.values())
|
||||
|
||||
def get_related_tables(self, table_name: str, depth: int = 1) -> List[Table]:
|
||||
"""
|
||||
获取关联表(通过外键关系)
|
||||
|
||||
Args:
|
||||
table_name: 起始表名
|
||||
depth: 递归深度(1=直接关联,2=间接关联...)
|
||||
|
||||
Returns:
|
||||
相关表列表
|
||||
"""
|
||||
result = set()
|
||||
visited = set()
|
||||
|
||||
def _traverse(tname: str, current_depth: int):
|
||||
if tname not in self._table_dict or current_depth > depth:
|
||||
return
|
||||
if tname in visited:
|
||||
return
|
||||
visited.add(tname)
|
||||
|
||||
table = self._table_dict[tname]
|
||||
for fk in table.foreign_keys:
|
||||
result.add(fk.ref_table)
|
||||
_traverse(fk.ref_table, current_depth + 1)
|
||||
|
||||
# 反向外键(被引用的表)
|
||||
for other_table in self._table_dict.values():
|
||||
for fk in other_table.foreign_keys:
|
||||
if fk.ref_table == tname:
|
||||
result.add(other_table.name)
|
||||
_traverse(other_table.name, current_depth + 1)
|
||||
|
||||
_traverse(table_name, 0)
|
||||
return [self._table_dict[t] for t in result if t in self._table_dict]
|
||||
|
||||
# ========== Schema生成 ==========
|
||||
|
||||
def to_compact_string(
|
||||
self,
|
||||
table_names: List[str] = None,
|
||||
include_columns: bool = True,
|
||||
max_columns_per_table: int = 15,
|
||||
) -> str:
|
||||
"""
|
||||
生成紧凑的Schema描述字符串(用于LLM输入)
|
||||
|
||||
Args:
|
||||
table_names: 指定包含的表(None表示所有表)
|
||||
include_columns: 是否包含字段详情
|
||||
max_columns_per_table: 每表最多显示字段数
|
||||
|
||||
Returns:
|
||||
Schema描述字符串
|
||||
"""
|
||||
if table_names is None:
|
||||
tables = self.get_tables()
|
||||
else:
|
||||
missing = [n for n in table_names if n not in self._table_dict]
|
||||
if missing:
|
||||
logger.warning(
|
||||
"以下表名不在已加载Schema中,已从本次Schema片段中省略: %s",
|
||||
missing,
|
||||
)
|
||||
tables = [t for t in self.get_tables() if t.name in table_names]
|
||||
|
||||
lines = []
|
||||
lines.append(f"数据库: {self.schema.name}")
|
||||
lines.append(f"涉及表数: {len(tables)}")
|
||||
lines.append("=" * 60)
|
||||
lines.append("")
|
||||
|
||||
for table in tables:
|
||||
lines.append(f"【表】{table.name}")
|
||||
if table.comment:
|
||||
lines.append(f" 描述: {table.comment}")
|
||||
|
||||
if include_columns and table.columns:
|
||||
lines.append(f" 字段:")
|
||||
for col in table.columns[:max_columns_per_table]:
|
||||
pk_marker = " [PK]" if col.is_primary_key else ""
|
||||
null_marker = " NOT NULL" if not col.nullable else ""
|
||||
comment = f" -- {col.comment}" if col.comment else ""
|
||||
lines.append(f" {col.name}: {col.data_type}{pk_marker}{null_marker}{comment}")
|
||||
|
||||
if len(table.columns) > max_columns_per_table:
|
||||
lines.append(f" ... 还有 {len(table.columns) - max_columns_per_table} 个字段")
|
||||
|
||||
# 外键关系
|
||||
if table.foreign_keys:
|
||||
lines.append(" 外键:")
|
||||
for fk in table.foreign_keys[:5]:
|
||||
lines.append(
|
||||
f" {', '.join(fk.columns)} → {fk.ref_table}({', '.join(fk.ref_columns)})"
|
||||
)
|
||||
|
||||
lines.append("")
|
||||
|
||||
return "\n".join(lines)
|
||||
|
||||
def to_filtered_schema(self, table_names: List[str]) -> DatabaseSchema:
|
||||
"""
|
||||
根据表名过滤,生成子集Schema
|
||||
|
||||
Args:
|
||||
table_names: 要保留的表名列表
|
||||
|
||||
Returns:
|
||||
新的DatabaseSchema对象
|
||||
"""
|
||||
filtered_tables = [t for t in self.get_tables() if t.name in table_names]
|
||||
return DatabaseSchema(
|
||||
name=f"{self.schema.name}_filtered",
|
||||
tables=filtered_tables,
|
||||
)
|
||||
|
||||
# ========== 向量索引支持 ==========
|
||||
|
||||
def set_embeddings(self, table_embeddings: Dict[str, List[float]]):
|
||||
"""
|
||||
为表设置预计算的embedding向量
|
||||
|
||||
Args:
|
||||
table_embeddings: {table_name: embedding_vector}
|
||||
"""
|
||||
for table in self.get_tables():
|
||||
if table.name in table_embeddings:
|
||||
table.embedding = table_embeddings[table.name]
|
||||
|
||||
def get_table_for_embedding(
|
||||
self,
|
||||
max_columns: int = 48,
|
||||
max_col_comment_chars: int = 160,
|
||||
max_total_chars: int = 8000,
|
||||
) -> List[Dict]:
|
||||
"""
|
||||
获取用于embedding的表信息列表(单条文本同时承载 meta + structure 的可检索信息)。
|
||||
|
||||
- 表级:表名、表注释(来自 table_meta 合并后的 Table.comment)
|
||||
- 列级:字段名、类型、主键标记、列注释(来自 table_structure 解析后的 Column)
|
||||
|
||||
Returns:
|
||||
[{"name": "...", "text": "..."}, ...]
|
||||
"""
|
||||
result = []
|
||||
for table in self.get_tables():
|
||||
text_parts = [f"表名: {table.name}"]
|
||||
if table.comment:
|
||||
text_parts.append(f"描述: {table.comment}")
|
||||
|
||||
col_segments: List[str] = []
|
||||
shown = table.columns[:max_columns]
|
||||
for col in shown:
|
||||
seg = f"{col.name} {col.data_type}"
|
||||
if col.is_primary_key:
|
||||
seg += " PK"
|
||||
if col.comment:
|
||||
c = col.comment.strip().replace("\n", " ")
|
||||
if len(c) > max_col_comment_chars:
|
||||
c = c[: max_col_comment_chars - 1] + "…"
|
||||
seg += f" — {c}"
|
||||
col_segments.append(seg)
|
||||
|
||||
if col_segments:
|
||||
text_parts.append("字段: " + ";".join(col_segments))
|
||||
if len(table.columns) > max_columns:
|
||||
rest = len(table.columns) - max_columns
|
||||
text_parts.append(f"另有{rest}个字段未列出(共{len(table.columns)}列)")
|
||||
|
||||
text = "。".join(text_parts)
|
||||
if len(text) > max_total_chars:
|
||||
text = text[: max_total_chars - 1] + "…"
|
||||
|
||||
result.append({
|
||||
"name": table.name,
|
||||
"text": text,
|
||||
})
|
||||
return result
|
||||
|
||||
# ========== 缓存支持 ==========
|
||||
|
||||
def get_cached_result(self, key: str) -> Optional[str]:
|
||||
"""获取缓存结果"""
|
||||
return self._cache.get(key)
|
||||
|
||||
def set_cached_result(self, key: str, value: str, ttl: int = 3600):
|
||||
"""设置缓存结果"""
|
||||
self._cache[key] = value
|
||||
# TODO: 可扩展为Redis缓存
|
||||
|
||||
# ========== 统计信息 ==========
|
||||
|
||||
def get_statistics(self) -> Dict:
|
||||
"""获取Schema统计信息"""
|
||||
total_columns = sum(len(t.columns) for t in self.get_tables())
|
||||
tables_with_pk = sum(1 for t in self.get_tables() if t.primary_keys)
|
||||
tables_with_fk = sum(1 for t in self.get_tables() if t.foreign_keys)
|
||||
|
||||
return {
|
||||
"database": self.schema.name,
|
||||
"total_tables": len(self.get_tables()),
|
||||
"total_columns": total_columns,
|
||||
"avg_columns_per_table": total_columns / len(self.get_tables()) if self.get_tables() else 0,
|
||||
"tables_with_primary_key": tables_with_pk,
|
||||
"tables_with_foreign_key": tables_with_fk,
|
||||
"table_names": self.list_tables(),
|
||||
}
|
||||
|
||||
def __len__(self) -> int:
|
||||
return len(self.get_tables())
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"SchemaManager(database={self.schema.name}, tables={len(self)})"
|
||||
@@ -0,0 +1,179 @@
|
||||
"""
|
||||
Schema数据模型定义
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from typing import List, Optional, Dict, Any
|
||||
from datetime import datetime
|
||||
|
||||
|
||||
@dataclass
|
||||
class Column:
|
||||
"""列字段定义"""
|
||||
name: str
|
||||
data_type: str
|
||||
comment: Optional[str] = None
|
||||
nullable: bool = True
|
||||
is_primary_key: bool = False
|
||||
is_foreign_key: bool = False
|
||||
|
||||
def __str__(self):
|
||||
null_str = "NULL" if self.nullable else "NOT NULL"
|
||||
pk_str = " PK" if self.is_primary_key else ""
|
||||
comment_str = f" -- {self.comment}" if self.comment else ""
|
||||
return f"{self.name} {self.data_type}{null_str}{pk_str}{comment_str}"
|
||||
|
||||
|
||||
@dataclass
|
||||
class ForeignKey:
|
||||
"""外键关系"""
|
||||
columns: List[str] # 本表字段
|
||||
ref_table: str # 引用表
|
||||
ref_columns: List[str] # 引用字段
|
||||
|
||||
|
||||
@dataclass
|
||||
class Table:
|
||||
"""数据表定义"""
|
||||
name: str
|
||||
comment: Optional[str] = None
|
||||
columns: List[Column] = field(default_factory=list)
|
||||
primary_keys: List[str] = field(default_factory=list)
|
||||
foreign_keys: List[ForeignKey] = field(default_factory=list)
|
||||
# 向量embedding(可选,用于检索)
|
||||
embedding: Optional[Any] = None
|
||||
|
||||
@property
|
||||
def column_dict(self) -> Dict[str, Column]:
|
||||
"""字段名到Column对象的映射"""
|
||||
return {col.name: col for col in self.columns}
|
||||
|
||||
@property
|
||||
def column_names(self) -> List[str]:
|
||||
"""所有字段名列表"""
|
||||
return [col.name for col in self.columns]
|
||||
|
||||
def to_compact_string(self, max_columns: int = 20) -> str:
|
||||
"""
|
||||
生成紧凑的Schema字符串(用于LLM输入)
|
||||
|
||||
Args:
|
||||
max_columns: 最大字段数,超过时截断
|
||||
|
||||
Returns:
|
||||
紧凑的Schema描述字符串
|
||||
"""
|
||||
parts = [f"TABLE {self.name} ("]
|
||||
|
||||
# 只显示关键字段(主键、常见字段)
|
||||
display_columns = self.columns[:max_columns] if len(self.columns) > max_columns else self.columns
|
||||
|
||||
col_strs = []
|
||||
for col in display_columns:
|
||||
col_str = f" {col.name}: {col.data_type}"
|
||||
if col.is_primary_key:
|
||||
col_str += " [PK]"
|
||||
if col.comment:
|
||||
col_str += f" -- {col.comment}"
|
||||
col_strs.append(col_str)
|
||||
|
||||
parts.append(",\n".join(col_strs))
|
||||
|
||||
if len(self.columns) > max_columns:
|
||||
parts.append(f"\n -- 还有 {len(self.columns) - max_columns} 个字段未显示")
|
||||
|
||||
parts.append(")")
|
||||
return "".join(parts)
|
||||
|
||||
def to_dict(self) -> Dict:
|
||||
"""转换为字典格式"""
|
||||
return {
|
||||
"name": self.name,
|
||||
"comment": self.comment,
|
||||
"columns": [
|
||||
{
|
||||
"name": col.name,
|
||||
"data_type": col.data_type,
|
||||
"comment": col.comment,
|
||||
"nullable": col.nullable,
|
||||
"is_primary_key": col.is_primary_key,
|
||||
}
|
||||
for col in self.columns
|
||||
],
|
||||
"primary_keys": self.primary_keys,
|
||||
"foreign_keys": [
|
||||
{
|
||||
"columns": fk.columns,
|
||||
"ref_table": fk.ref_table,
|
||||
"ref_columns": fk.ref_columns,
|
||||
}
|
||||
for fk in self.foreign_keys
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class DatabaseSchema:
|
||||
"""数据库Schema(多个表的集合)"""
|
||||
name: str
|
||||
tables: List[Table] = field(default_factory=list)
|
||||
created_at: datetime = field(default_factory=datetime.now)
|
||||
|
||||
@property
|
||||
def table_dict(self) -> Dict[str, Table]:
|
||||
"""表名字典"""
|
||||
return {tbl.name: tbl for tbl in self.tables}
|
||||
|
||||
def get_table(self, table_name: str) -> Optional[Table]:
|
||||
"""获取指定表"""
|
||||
return self.table_dict.get(table_name)
|
||||
|
||||
def to_summary_string(self, include_columns: bool = True) -> str:
|
||||
"""
|
||||
生成Schema摘要字符串(用于LLM输入)
|
||||
|
||||
Args:
|
||||
include_columns: 是否包含字段信息
|
||||
|
||||
Returns:
|
||||
Schema描述字符串
|
||||
"""
|
||||
lines = [f"数据库: {self.name}\n"]
|
||||
lines.append(f"表总数: {len(self.tables)}\n")
|
||||
lines.append("=" * 60 + "\n")
|
||||
|
||||
for table in self.tables:
|
||||
lines.append(f"表名: {table.name}")
|
||||
if table.comment:
|
||||
lines.append(f"描述: {table.comment}")
|
||||
lines.append(f"字段数: {len(table.columns)}")
|
||||
|
||||
if include_columns and table.columns:
|
||||
lines.append("字段列表:")
|
||||
for col in table.columns[:30]: # 限制字段数
|
||||
pk_marker = " [PK]" if col.is_primary_key else ""
|
||||
null_marker = " NULL" if col.nullable else " NOT NULL"
|
||||
comment = f" -- {col.comment}" if col.comment else ""
|
||||
lines.append(f" {col.name}: {col.data_type}{pk_marker}{null_marker}{comment}")
|
||||
|
||||
if len(table.columns) > 30:
|
||||
lines.append(f" ... 还有 {len(table.columns) - 30} 个字段")
|
||||
|
||||
# 外键关系
|
||||
if table.foreign_keys:
|
||||
lines.append("外键关系:")
|
||||
for fk in table.foreign_keys[:5]:
|
||||
lines.append(
|
||||
f" {', '.join(fk.columns)} -> {fk.ref_table}({', '.join(fk.ref_columns)})"
|
||||
)
|
||||
|
||||
lines.append("") # 空行分隔
|
||||
|
||||
return "\n".join(lines)
|
||||
|
||||
def to_compact_dict(self) -> Dict:
|
||||
"""紧凑字典格式(用于序列化)"""
|
||||
return {
|
||||
"database": self.name,
|
||||
"tables": [tbl.to_dict() for tbl in self.tables],
|
||||
}
|
||||
@@ -0,0 +1 @@
|
||||
# utils 包初始化
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -0,0 +1,135 @@
|
||||
"""
|
||||
用户输入意图分类:区分「自然语言查数 / Text2SQL」与「寒暄、致谢、元问题」等不适合直接生成 SQL 的对话。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import re
|
||||
import unicodedata
|
||||
from enum import Enum
|
||||
from typing import NamedTuple, Optional
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class DialogIntent(str, Enum):
|
||||
TEXT2SQL = "text2sql"
|
||||
CONVERSATION = "conversation"
|
||||
|
||||
|
||||
class DialogClassifyResult(NamedTuple):
|
||||
intent: DialogIntent
|
||||
"""若为 CONVERSATION,可展示给用户的引导文案;TEXT2SQL 时为 None。"""
|
||||
reply_suggestion: Optional[str] = None
|
||||
|
||||
|
||||
DEFAULT_CONVERSATION_REPLY = (
|
||||
"您好,我是业务库 Text2SQL 助手。\n"
|
||||
"请用自然语言描述要查询或统计的内容(例如:查询某账户可用余额、按经纪商汇总未结算交易笔数)。\n"
|
||||
"输入 quit 或 exit 可退出。"
|
||||
)
|
||||
|
||||
_EMPTY_INPUT_REPLY = "请输入具体的业务查询问题,或输入 quit 退出。"
|
||||
|
||||
# 一旦出现,倾向于按「要查数据」处理(含常见业务词,避免误判)
|
||||
_SQL_OR_QUERY_HINT_RE = re.compile(
|
||||
r"(查|查询|查出|检索|统计|列出|汇总|求和|平均|分组|排序|排名|显示|导出|筛选|过滤|"
|
||||
r"多少|几个|几张|哪些|占比|同比|环比|"
|
||||
r"余额|交易|账户|持仓|报表|结算|合约|订单|流水|经纪商|对手方|证券|资金|"
|
||||
r"query|select|list|show|count|sum|avg|how\s+many|statistics|\bfrom\b|\bwhere\b|\btable\b)",
|
||||
re.IGNORECASE,
|
||||
)
|
||||
|
||||
_CHITCHAT_PHRASES = frozenset(
|
||||
{
|
||||
"你好",
|
||||
"您好",
|
||||
"嗨",
|
||||
"哈喽",
|
||||
"hello",
|
||||
"hi",
|
||||
"hey",
|
||||
"早上好",
|
||||
"下午好",
|
||||
"晚上好",
|
||||
"在吗",
|
||||
"在不在",
|
||||
"谢谢",
|
||||
"多谢",
|
||||
"感谢",
|
||||
"thanks",
|
||||
"thank you",
|
||||
"thx",
|
||||
"再见",
|
||||
"拜拜",
|
||||
"bye",
|
||||
"goodbye",
|
||||
"哈哈",
|
||||
"哈哈哈",
|
||||
"嗯",
|
||||
"嗯嗯",
|
||||
"好的",
|
||||
"好",
|
||||
"ok",
|
||||
"okay",
|
||||
"行",
|
||||
"收到",
|
||||
"👋",
|
||||
"😀",
|
||||
"哈哈谢谢",
|
||||
}
|
||||
)
|
||||
|
||||
_CHITCHAT_KEYS = frozenset(p.casefold() for p in _CHITCHAT_PHRASES)
|
||||
|
||||
_META_QUESTION_RE = re.compile(
|
||||
r"(你是谁|你是什么|你能(做|干)什么|你会什么|怎么用|如何使用|使用说明|帮助|help\b|"
|
||||
r"什么功能|干啥的)",
|
||||
re.IGNORECASE,
|
||||
)
|
||||
|
||||
|
||||
def _normalize(text: str) -> str:
|
||||
t = unicodedata.normalize("NFKC", text or "").strip()
|
||||
t = re.sub(r"\s+", " ", t)
|
||||
return t
|
||||
|
||||
|
||||
def _strip_trailing_punct(t: str) -> str:
|
||||
return re.sub(r"[!!。.??,,;;:~~…、]+$", "", t).strip()
|
||||
|
||||
|
||||
def classify_dialog(user_text: str) -> DialogClassifyResult:
|
||||
"""
|
||||
对用户一轮输入做粗分类。
|
||||
|
||||
策略:优先用「查询/业务」关键词锁定 TEXT2SQL;否则对短寒暄、致谢、元问题判为 CONVERSATION;
|
||||
其余默认 TEXT2SQL,避免漏判真实查询。
|
||||
"""
|
||||
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)")
|
||||
return DialogClassifyResult(DialogIntent.TEXT2SQL, None)
|
||||
|
||||
core = _strip_trailing_punct(t)
|
||||
if core.casefold() in _CHITCHAT_KEYS:
|
||||
logger.debug("[dialog] intent=conversation (chitchat phrase)")
|
||||
return DialogClassifyResult(
|
||||
DialogIntent.CONVERSATION, reply_suggestion=DEFAULT_CONVERSATION_REPLY
|
||||
)
|
||||
|
||||
if _META_QUESTION_RE.search(t):
|
||||
logger.debug("[dialog] intent=conversation (meta question)")
|
||||
return DialogClassifyResult(
|
||||
DialogIntent.CONVERSATION,
|
||||
reply_suggestion=DEFAULT_CONVERSATION_REPLY,
|
||||
)
|
||||
|
||||
logger.debug("[dialog] intent=text2sql (default)")
|
||||
return DialogClassifyResult(DialogIntent.TEXT2SQL, None)
|
||||
@@ -0,0 +1,530 @@
|
||||
"""
|
||||
Embedding 封装:本地 Qwen3-Embedding,或兼容 OpenAI /v1/embeddings 的远程 API
|
||||
(ModelScope 推理、阿里云 DashScope 等)。
|
||||
"""
|
||||
|
||||
import os
|
||||
from typing import Union, List, Optional, Any
|
||||
import numpy as np
|
||||
from pathlib import Path
|
||||
|
||||
try:
|
||||
from transformers import AutoModel, AutoTokenizer
|
||||
import torch
|
||||
_TRANSFORMERS_AVAILABLE = True
|
||||
except ImportError:
|
||||
_TRANSFORMERS_AVAILABLE = False
|
||||
|
||||
import logging
|
||||
from dataclasses import dataclass
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class _RemoteEmbeddingEnv:
|
||||
"""从环境变量解析出的远程 OpenAI-Compatible Embedding 配置。"""
|
||||
|
||||
api_key: str
|
||||
base_url: str
|
||||
model: str
|
||||
max_batch: int
|
||||
label: str
|
||||
|
||||
|
||||
def _remote_embedding_from_env() -> _RemoteEmbeddingEnv:
|
||||
"""
|
||||
优先 ModelScope(MODELSCOPE_*);未配置时回退 DashScope(DASHSCOPE_*)。
|
||||
"""
|
||||
ms_key = os.getenv("MODELSCOPE_API_KEY", "").strip()
|
||||
ms_base = os.getenv("MODELSCOPE_BASE_URL", "").strip()
|
||||
ds_key = os.getenv("DASHSCOPE_API_KEY", "").strip()
|
||||
ds_base = (
|
||||
os.getenv("DASHSCOPE_BASE_URL") or os.getenv("DASHSCOPE_base_url", "")
|
||||
).strip()
|
||||
|
||||
# 仅当配置了 API Key 时走 ModelScope(避免仅有 BASE_URL 时误判、阻断 DashScope)
|
||||
if ms_key:
|
||||
base_url = ms_base or "https://api-inference.modelscope.cn/v1"
|
||||
model = (
|
||||
os.getenv("MODELSCOPE_EMBEDDING_MODEL")
|
||||
or os.getenv("MODELSCOPE_MODEL", "Qwen/Qwen3-Embedding-8B")
|
||||
).strip()
|
||||
mb = os.getenv("MODELSCOPE_EMBEDDING_MAX_BATCH", "32").strip()
|
||||
max_batch = max(1, int(mb)) if mb.isdigit() else 32
|
||||
if not model:
|
||||
raise ValueError("未配置 MODELSCOPE_EMBEDDING_MODEL(或 MODELSCOPE_MODEL)")
|
||||
return _RemoteEmbeddingEnv(
|
||||
api_key=ms_key,
|
||||
base_url=base_url.rstrip("/"),
|
||||
model=model,
|
||||
max_batch=max_batch,
|
||||
label="ModelScope",
|
||||
)
|
||||
|
||||
# OpenAI 官方或兼容网关:OPENAI_API_KEY、OPENAI_EMBEDDING_MODEL、可选 OPENAI_BASE_URL
|
||||
oa_key = os.getenv("OPENAI_API_KEY", "").strip()
|
||||
if oa_key:
|
||||
if oa_key.startswith("http://") or oa_key.startswith("https://"):
|
||||
raise ValueError(
|
||||
"OPENAI_API_KEY 不能填写为 URL:请将网关地址写到 OPENAI_BASE_URL"
|
||||
"(例如 http://host:9080/v1),密钥单独写在 OPENAI_API_KEY"
|
||||
)
|
||||
oa_base = (
|
||||
os.getenv("OPENAI_BASE_URL", "").strip() or "https://api.openai.com/v1"
|
||||
)
|
||||
oa_model = os.getenv("OPENAI_EMBEDDING_MODEL", "").strip()
|
||||
if not oa_model:
|
||||
raise ValueError(
|
||||
"使用 OpenAI 兼容 Embedding 时请设置 OPENAI_EMBEDDING_MODEL"
|
||||
)
|
||||
mb = os.getenv("OPENAI_EMBEDDING_MAX_BATCH", "100").strip()
|
||||
max_batch = max(1, int(mb)) if mb.isdigit() else 100
|
||||
return _RemoteEmbeddingEnv(
|
||||
api_key=oa_key,
|
||||
base_url=oa_base.rstrip("/"),
|
||||
model=oa_model,
|
||||
max_batch=max_batch,
|
||||
label="OpenAI",
|
||||
)
|
||||
|
||||
if not ds_key:
|
||||
raise ValueError(
|
||||
"远程 Embedding 未配置:请设置 MODELSCOPE_API_KEY(及可选 BASE_URL),"
|
||||
"或 OPENAI_API_KEY / OPENAI_EMBEDDING_MODEL(及可选 OPENAI_BASE_URL),"
|
||||
"或 DASHSCOPE_API_KEY / DASHSCOPE_BASE_URL / DASHSCOPE_MODEL"
|
||||
)
|
||||
if not ds_base:
|
||||
raise ValueError("未配置 DASHSCOPE_BASE_URL(或 DASHSCOPE_base_url)")
|
||||
model = os.getenv("DASHSCOPE_MODEL", "").strip()
|
||||
if not model:
|
||||
raise ValueError("未配置 DASHSCOPE_MODEL")
|
||||
mb = os.getenv("DASHSCOPE_EMBEDDING_MAX_BATCH", "10").strip()
|
||||
max_batch = max(1, int(mb)) if mb.isdigit() else 10
|
||||
return _RemoteEmbeddingEnv(
|
||||
api_key=ds_key,
|
||||
base_url=ds_base.rstrip("/"),
|
||||
model=model,
|
||||
max_batch=max_batch,
|
||||
label="DashScope",
|
||||
)
|
||||
|
||||
|
||||
class Qwen3Embedding:
|
||||
"""
|
||||
Qwen3-Embedding-0.6B 向量化封装
|
||||
|
||||
使用 Mean Pooling 将token embeddings聚合为句子向量,
|
||||
并进行L2归一化以支持余弦相似度计算。
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model_path: Optional[str] = None,
|
||||
device: Optional[str] = None,
|
||||
use_fp16: bool = False
|
||||
):
|
||||
"""
|
||||
初始化 embedding 模型
|
||||
|
||||
Args:
|
||||
model_path: 本地模型路径,若为None则从环境变量或默认路径加载
|
||||
device: 推理设备('cpu', 'cuda', 'cuda:0'等),None则自动选择
|
||||
use_fp16: 是否使用FP16混合精度(GPU可用时建议开启,速度更快)
|
||||
"""
|
||||
if not _TRANSFORMERS_AVAILABLE:
|
||||
raise ImportError(
|
||||
"transformers 和 torch 未安装。请运行:\n"
|
||||
"pip install transformers torch sentencepiece accelerate"
|
||||
)
|
||||
|
||||
# 确定模型路径
|
||||
if model_path is None:
|
||||
model_path = os.getenv(
|
||||
"EMBEDDING_MODEL_PATH",
|
||||
"./data/models/Qwen3-Embedding-0.6B"
|
||||
)
|
||||
|
||||
model_path = Path(model_path)
|
||||
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}"
|
||||
)
|
||||
|
||||
# 确定设备
|
||||
if device is None:
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
|
||||
self.device = device
|
||||
logger.info(f"加载Qwen3-Embedding模型:{model_path},设备:{device}")
|
||||
|
||||
# 加载 tokenizer:fast(Rust) 解析 tokenizer.json 需较新 tokenizers;
|
||||
# 旧版本会报 ModelWrapper / untagged enum,回退到慢速 tokenizer 可恢复。
|
||||
try:
|
||||
self.tokenizer = AutoTokenizer.from_pretrained(
|
||||
str(model_path), trust_remote_code=True
|
||||
)
|
||||
except Exception as e:
|
||||
err = str(e).lower()
|
||||
if "modelwrapper" in err or "untagged enum" in err:
|
||||
logger.warning(
|
||||
"快速 tokenizer 解析 tokenizer.json 失败(多为 tokenizers 过旧),"
|
||||
"改用 use_fast=False:%s",
|
||||
e,
|
||||
)
|
||||
self.tokenizer = AutoTokenizer.from_pretrained(
|
||||
str(model_path), use_fast=False, trust_remote_code=True
|
||||
)
|
||||
else:
|
||||
raise
|
||||
|
||||
try:
|
||||
self.model = AutoModel.from_pretrained(
|
||||
str(model_path), trust_remote_code=True
|
||||
)
|
||||
except ValueError as e:
|
||||
msg = str(e)
|
||||
if "qwen3" in msg.lower() or "does not recognize this architecture" in msg:
|
||||
raise RuntimeError(
|
||||
"当前 transformers 版本不支持 Qwen3(model_type=qwen3)。"
|
||||
"请升级:pip install \"transformers>=4.51.0\" \"tokenizers>=0.21\""
|
||||
) from e
|
||||
raise
|
||||
|
||||
# 设置为评估模式并移动设备
|
||||
self.model.eval()
|
||||
self.model.to(device)
|
||||
|
||||
# 混合精度(仅GPU)
|
||||
self.use_fp16 = use_fp16 and device != "cpu"
|
||||
if self.use_fp16:
|
||||
self.model.half()
|
||||
|
||||
# 嵌入维度
|
||||
self.embedding_dim = self.model.config.hidden_size
|
||||
logger.info(f"[OK] 模型加载完成,嵌入维度:{self.embedding_dim}")
|
||||
|
||||
def encode(
|
||||
self,
|
||||
texts: Union[str, List[str]],
|
||||
batch_size: int = 32,
|
||||
normalize: bool = True,
|
||||
max_length: int = 8192,
|
||||
show_progress: bool = False
|
||||
) -> np.ndarray:
|
||||
"""
|
||||
编码文本为向量
|
||||
|
||||
Args:
|
||||
texts: 单个文本或文本列表
|
||||
batch_size: 批处理大小(根据显存调整)
|
||||
normalize: 是否L2归一化(余弦相似度必需)
|
||||
max_length: 最大序列长度(模型支持8192,建议512-1024平衡速度与精度)
|
||||
show_progress: 是否显示进度条(需安装tqdm)
|
||||
|
||||
Returns:
|
||||
numpy数组,shape=(len(texts), embedding_dim)
|
||||
"""
|
||||
if isinstance(texts, str):
|
||||
texts = [texts]
|
||||
|
||||
if not texts:
|
||||
return np.empty((0, self.embedding_dim), dtype=np.float32)
|
||||
|
||||
all_embeddings = []
|
||||
|
||||
# 可选进度条
|
||||
iterator = range(0, len(texts), batch_size)
|
||||
if show_progress:
|
||||
try:
|
||||
from tqdm import tqdm
|
||||
iterator = tqdm(iterator, desc="Embedding")
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
for i in iterator:
|
||||
batch = texts[i:i + batch_size]
|
||||
|
||||
# Tokenize
|
||||
inputs = self.tokenizer(
|
||||
batch,
|
||||
padding=True,
|
||||
truncation=True,
|
||||
max_length=max_length,
|
||||
return_tensors="pt"
|
||||
).to(self.device)
|
||||
|
||||
# Inference
|
||||
with torch.no_grad():
|
||||
outputs = self.model(**inputs)
|
||||
# Mean Pooling: 取序列维度的平均值
|
||||
# outputs.last_hidden_state shape: (batch, seq_len, hidden_size)
|
||||
embeddings = outputs.last_hidden_state.mean(dim=1)
|
||||
|
||||
# 转换为numpy(保持在CPU)
|
||||
if self.device != "cpu":
|
||||
embeddings = embeddings.cpu()
|
||||
embeddings = embeddings.numpy()
|
||||
|
||||
if normalize:
|
||||
# L2归一化(余弦相似度必需)
|
||||
norms = np.linalg.norm(embeddings, axis=1, keepdims=True)
|
||||
embeddings = embeddings / (norms + 1e-10)
|
||||
|
||||
all_embeddings.append(embeddings)
|
||||
|
||||
return np.vstack(all_embeddings).astype(np.float32)
|
||||
|
||||
def similarity(
|
||||
self,
|
||||
emb1: np.ndarray,
|
||||
emb2: np.ndarray
|
||||
) -> np.ndarray:
|
||||
"""
|
||||
计算两组embedding的余弦相似度
|
||||
|
||||
Args:
|
||||
emb1: 第一组向量 (n, dim)
|
||||
emb2: 第二组向量 (m, dim)
|
||||
|
||||
Returns:
|
||||
相似度矩阵 (n, m),值域[-1, 1](若已归一化则为[0, 1])
|
||||
"""
|
||||
# 确保已归一化
|
||||
return np.dot(emb1, emb2.T)
|
||||
|
||||
def encode_and_search(
|
||||
self,
|
||||
query: str,
|
||||
documents: List[str],
|
||||
top_k: int = 5
|
||||
) -> List[dict]:
|
||||
"""
|
||||
便捷方法:编码查询并检索最相似的文档
|
||||
|
||||
Args:
|
||||
query: 查询文本
|
||||
documents: 候选文档列表
|
||||
top_k: 返回前K个结果
|
||||
|
||||
Returns:
|
||||
[{"score": float, "document": str, "index": int}, ...]
|
||||
"""
|
||||
query_emb = self.encode([query], normalize=True)
|
||||
doc_embs = self.encode(documents, normalize=True)
|
||||
|
||||
scores = self.similarity(query_emb, doc_embs)[0]
|
||||
|
||||
# 获取top_k
|
||||
top_indices = np.argsort(scores)[::-1][:top_k]
|
||||
|
||||
results = []
|
||||
for idx in top_indices:
|
||||
results.append({
|
||||
"score": float(scores[idx]),
|
||||
"document": documents[idx],
|
||||
"index": int(idx)
|
||||
})
|
||||
|
||||
return results
|
||||
|
||||
|
||||
def _env_flag(name: str, default: str = "true") -> bool:
|
||||
return os.getenv(name, default).strip().lower() in ("1", "true", "yes", "on")
|
||||
|
||||
|
||||
class OpenAICompatibleRemoteEmbedding:
|
||||
"""
|
||||
通过 OpenAI 兼容接口获取文本向量(POST /v1/embeddings)。
|
||||
|
||||
环境变量(优先级:ModelScope → OpenAI 兼容 → DashScope):
|
||||
- ModelScope:MODELSCOPE_API_KEY、可选 MODELSCOPE_BASE_URL(默认
|
||||
https://api-inference.modelscope.cn/v1)、MODELSCOPE_EMBEDDING_MODEL
|
||||
或 MODELSCOPE_MODEL、可选 MODELSCOPE_EMBEDDING_MAX_BATCH
|
||||
- OpenAI 兼容:OPENAI_API_KEY、OPENAI_EMBEDDING_MODEL、可选 OPENAI_BASE_URL
|
||||
(默认 https://api.openai.com/v1)、可选 OPENAI_EMBEDDING_MAX_BATCH
|
||||
- DashScope:DASHSCOPE_API_KEY、DASHSCOPE_BASE_URL、DASHSCOPE_MODEL、
|
||||
可选 DASHSCOPE_EMBEDDING_MAX_BATCH
|
||||
|
||||
可选 VECTOR_DIM:在首次请求前确定空列表返回的维度。
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
api_key: Optional[str] = None,
|
||||
base_url: Optional[str] = None,
|
||||
model: Optional[str] = None,
|
||||
max_batch: Optional[int] = None,
|
||||
provider_label: Optional[str] = None,
|
||||
):
|
||||
try:
|
||||
from openai import OpenAI
|
||||
except ImportError as e:
|
||||
raise ImportError(
|
||||
"使用远程 Embedding 需要安装 openai:pip install openai"
|
||||
) from e
|
||||
|
||||
if api_key is not None and base_url is not None and model is not None:
|
||||
cfg = _RemoteEmbeddingEnv(
|
||||
api_key=api_key.strip(),
|
||||
base_url=base_url.strip().rstrip("/"),
|
||||
model=model.strip(),
|
||||
max_batch=max(1, int(max_batch)) if max_batch is not None else 32,
|
||||
label=provider_label or "custom",
|
||||
)
|
||||
else:
|
||||
cfg = _remote_embedding_from_env()
|
||||
|
||||
self.api_key = cfg.api_key
|
||||
self.base_url = cfg.base_url
|
||||
self.model = cfg.model
|
||||
self._api_max_batch = cfg.max_batch
|
||||
self._provider_label = cfg.label
|
||||
|
||||
self._client = OpenAI(api_key=self.api_key, base_url=self.base_url)
|
||||
vd = os.getenv("VECTOR_DIM", "").strip()
|
||||
self._embedding_dim: Optional[int] = int(vd) if vd.isdigit() else None
|
||||
|
||||
logger.info(
|
||||
"使用 %s Embedding API:model=%s,base_url=%s,max_batch=%s",
|
||||
self._provider_label,
|
||||
self.model,
|
||||
self.base_url,
|
||||
self._api_max_batch,
|
||||
)
|
||||
|
||||
@property
|
||||
def embedding_dim(self) -> int:
|
||||
if self._embedding_dim is None:
|
||||
raise RuntimeError(
|
||||
"尚未获知向量维度:请先执行一次 encode,或在 .env 中设置 VECTOR_DIM"
|
||||
)
|
||||
return self._embedding_dim
|
||||
|
||||
def _set_dim_from_vector(self, vec: List[float]) -> None:
|
||||
if self._embedding_dim is None:
|
||||
self._embedding_dim = len(vec)
|
||||
logger.info("[OK] Embedding 向量维度:%s", self._embedding_dim)
|
||||
|
||||
def encode(
|
||||
self,
|
||||
texts: Union[str, List[str]],
|
||||
batch_size: int = 10,
|
||||
normalize: bool = True,
|
||||
max_length: int = 8192,
|
||||
show_progress: bool = False,
|
||||
) -> np.ndarray:
|
||||
del max_length # API 侧截断,此处仅保持签名与本地实现一致
|
||||
|
||||
if isinstance(texts, str):
|
||||
texts = [texts]
|
||||
|
||||
if not texts:
|
||||
dim = self._embedding_dim
|
||||
if dim is None:
|
||||
vd = os.getenv("VECTOR_DIM", "").strip()
|
||||
dim = int(vd) if vd.isdigit() else 1024
|
||||
return np.empty((0, dim), dtype=np.float32)
|
||||
|
||||
# 无论调用方传多大,不能超过远端接口单次条数上限
|
||||
step = max(1, min(int(batch_size), self._api_max_batch))
|
||||
|
||||
all_embeddings: List[np.ndarray] = []
|
||||
iterator = range(0, len(texts), step)
|
||||
if show_progress:
|
||||
try:
|
||||
from tqdm import tqdm
|
||||
|
||||
iterator = tqdm(iterator, desc="Embedding (API)")
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
for i in iterator:
|
||||
batch = texts[i : i + step]
|
||||
resp = self._client.embeddings.create(
|
||||
model=self.model,
|
||||
input=batch,
|
||||
encoding_format="float",
|
||||
)
|
||||
rows = sorted(
|
||||
[(d.index, d.embedding) for d in resp.data],
|
||||
key=lambda x: x[0],
|
||||
)
|
||||
batch_embs = np.array([e for _, e in rows], dtype=np.float32)
|
||||
if batch_embs.size > 0:
|
||||
self._set_dim_from_vector(batch_embs[0].tolist())
|
||||
|
||||
if normalize:
|
||||
norms = np.linalg.norm(batch_embs, axis=1, keepdims=True)
|
||||
batch_embs = batch_embs / (norms + 1e-10)
|
||||
|
||||
all_embeddings.append(batch_embs)
|
||||
|
||||
return np.vstack(all_embeddings).astype(np.float32)
|
||||
|
||||
def similarity(self, emb1: np.ndarray, emb2: np.ndarray) -> np.ndarray:
|
||||
return np.dot(emb1, emb2.T)
|
||||
|
||||
def encode_and_search(
|
||||
self,
|
||||
query: str,
|
||||
documents: List[str],
|
||||
top_k: int = 5,
|
||||
) -> List[dict]:
|
||||
query_emb = self.encode([query], normalize=True)
|
||||
doc_embs = self.encode(documents, normalize=True)
|
||||
scores = self.similarity(query_emb, doc_embs)[0]
|
||||
top_indices = np.argsort(scores)[::-1][:top_k]
|
||||
return [
|
||||
{"score": float(scores[idx]), "document": documents[idx], "index": int(idx)}
|
||||
for idx in top_indices
|
||||
]
|
||||
|
||||
|
||||
# 向后兼容旧名称
|
||||
DashScopeOpenAIEmbedding = OpenAICompatibleRemoteEmbedding
|
||||
|
||||
# 全局单例(避免重复加载模型,节省显存/内存)
|
||||
_embedding_instance: Optional[Any] = None
|
||||
|
||||
|
||||
def get_embedder(
|
||||
model_path: Optional[str] = None,
|
||||
device: Optional[str] = None,
|
||||
force_reload: bool = False,
|
||||
) -> Any:
|
||||
"""
|
||||
获取 Embedding 单例:USE_LOCAL_EMBEDDING=true 时用本地 Qwen3,否则用远程
|
||||
OpenAI 兼容 API(优先级见 OpenAICompatibleRemoteEmbedding)。
|
||||
"""
|
||||
global _embedding_instance
|
||||
|
||||
if force_reload or _embedding_instance is None:
|
||||
if _env_flag("USE_LOCAL_EMBEDDING", "true"):
|
||||
_embedding_instance = Qwen3Embedding(
|
||||
model_path=model_path,
|
||||
device=device,
|
||||
)
|
||||
else:
|
||||
_embedding_instance = OpenAICompatibleRemoteEmbedding()
|
||||
|
||||
return _embedding_instance
|
||||
|
||||
|
||||
def clear_embedder():
|
||||
"""清空单例(用于测试或切换模型)"""
|
||||
global _embedding_instance
|
||||
_embedding_instance = None
|
||||
import gc
|
||||
|
||||
gc.collect()
|
||||
if _TRANSFORMERS_AVAILABLE:
|
||||
import torch
|
||||
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
@@ -0,0 +1,357 @@
|
||||
"""
|
||||
Few-shot示例选择器 - 基于经验数据集动态选择相关示例
|
||||
|
||||
与 Schema 向量检索一致,使用 utils.embedding.get_embedder()(本地 Qwen3 或远程 OpenAI 兼容 API,
|
||||
由 USE_LOCAL_EMBEDDING 及 OPENAI_* / MODELSCOPE_* / DASHSCOPE_* 等环境变量决定)。
|
||||
|
||||
用法:
|
||||
from utils.fewshot_selector import FewShotSelector
|
||||
|
||||
selector = FewShotSelector("data/experiences/all_samples.jsonl")
|
||||
examples = selector.select(question="查询2024年1月的销售额", top_k=3)
|
||||
|
||||
# 在Prompt中使用
|
||||
prompt = f"{examples}\n当前问题:{question}\nSchema:{schema}"
|
||||
"""
|
||||
|
||||
import json
|
||||
import os
|
||||
from pathlib import Path
|
||||
from typing import List, Dict, Optional
|
||||
from dataclasses import dataclass
|
||||
import numpy as np
|
||||
import logging
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_DEFAULT_LOCAL_EMBED_PATH = "./data/models/Qwen3-Embedding-0.6B"
|
||||
|
||||
|
||||
@dataclass
|
||||
class ExperienceSample:
|
||||
"""经验数据样本"""
|
||||
qid: str
|
||||
question_zh: str
|
||||
question_en: Optional[str]
|
||||
sql: str
|
||||
explanation: str
|
||||
rating: Optional[int]
|
||||
tags: List[str]
|
||||
difficulty: str
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, data: dict) -> "ExperienceSample":
|
||||
return cls(
|
||||
qid=data.get("qid", ""),
|
||||
question_zh=data.get("question_zh", ""),
|
||||
question_en=data.get("question_en"),
|
||||
sql=data.get("sql", ""),
|
||||
explanation=data.get("explanation", ""),
|
||||
rating=data.get("rating"),
|
||||
tags=data.get("tags", []),
|
||||
difficulty=data.get("difficulty", "medium")
|
||||
)
|
||||
|
||||
def to_fewshot_format(self, include_explanation: bool = True) -> str:
|
||||
"""转换为few-shot格式"""
|
||||
result = f"问题:{self.question_zh}\nSQL:\n{self.sql}"
|
||||
if include_explanation and self.explanation:
|
||||
result += f"\n说明:{self.explanation[:200]}"
|
||||
return result
|
||||
|
||||
def to_dict(self) -> dict:
|
||||
return {
|
||||
"qid": self.qid,
|
||||
"question": self.question_zh,
|
||||
"sql": self.sql,
|
||||
"rating": self.rating,
|
||||
"tags": self.tags,
|
||||
"difficulty": self.difficulty
|
||||
}
|
||||
|
||||
|
||||
class FewShotSelector:
|
||||
"""
|
||||
Few-shot示例选择器
|
||||
|
||||
根据用户问题,从经验数据集中检索最相似的示例,
|
||||
用于增强Prompt,提升LLM生成质量。
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
samples_path: str,
|
||||
embedding_model_path: Optional[str] = None,
|
||||
use_cache: bool = True,
|
||||
):
|
||||
"""
|
||||
Args:
|
||||
samples_path: 样本JSONL文件路径
|
||||
embedding_model_path: 本地 Embedding 模型目录;None 时用环境变量
|
||||
EMBEDDING_MODEL_PATH(仅 USE_LOCAL_EMBEDDING=true 时有效)
|
||||
use_cache: 是否缓存样本向量(按向量维度分文件,换模型会自动重建)
|
||||
"""
|
||||
self.samples_path = Path(samples_path)
|
||||
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._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()
|
||||
|
||||
def _load_samples(self):
|
||||
"""加载样本数据"""
|
||||
if not self.samples_path.exists():
|
||||
raise FileNotFoundError(f"样本文件不存在: {self.samples_path}")
|
||||
|
||||
logger.info(f"加载样本: {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)} 个样本")
|
||||
|
||||
def _build_index(self):
|
||||
"""用项目统一 Embedder 构建语义索引"""
|
||||
from utils.embedding import get_embedder
|
||||
|
||||
self._embedder = get_embedder(self._embedding_model_path)
|
||||
probe = self._embedder.encode(
|
||||
[" "],
|
||||
batch_size=1,
|
||||
normalize=True,
|
||||
show_progress=False,
|
||||
)
|
||||
dim = int(probe.shape[1])
|
||||
self.cache_path = (
|
||||
self.samples_path.parent / f"{self.samples_path.stem}.fewshot_dim{dim}.npy"
|
||||
)
|
||||
|
||||
if self.use_cache and self.cache_path.exists():
|
||||
try:
|
||||
self.embeddings = np.load(self.cache_path)
|
||||
if (
|
||||
self.embeddings.shape[0] == len(self.samples)
|
||||
and self.embeddings.shape[1] == dim
|
||||
):
|
||||
logger.info(f"[OK] 加载 Few-shot 向量缓存: {self.cache_path}")
|
||||
return
|
||||
except Exception as e:
|
||||
logger.warning(f"Few-shot 缓存加载失败: {e},将重新计算")
|
||||
|
||||
# 远程 API 通常不接受空字符串作 input
|
||||
questions = [
|
||||
(s.question_zh or "").strip() or " "
|
||||
for s in self.samples
|
||||
]
|
||||
if not questions:
|
||||
self.embeddings = np.empty((0, dim), dtype=np.float32)
|
||||
return
|
||||
|
||||
logger.info(f"计算 {len(questions)} 个 Few-shot 样本向量...")
|
||||
self.embeddings = self._embedder.encode(
|
||||
questions,
|
||||
batch_size=min(32, len(questions)),
|
||||
normalize=True,
|
||||
show_progress=True,
|
||||
)
|
||||
|
||||
if self.use_cache and self.cache_path is not None:
|
||||
np.save(self.cache_path, self.embeddings)
|
||||
logger.info(f"[OK] Few-shot 向量已缓存: {self.cache_path}")
|
||||
|
||||
def select(
|
||||
self,
|
||||
question: str,
|
||||
top_k: int = 3,
|
||||
min_rating: Optional[int] = None,
|
||||
required_tags: Optional[List[str]] = None,
|
||||
max_difficulty: str = "hard",
|
||||
exclude_qids: Optional[List[str]] = None
|
||||
) -> List[ExperienceSample]:
|
||||
"""
|
||||
选择最相关的few-shot示例
|
||||
|
||||
Args:
|
||||
question: 用户问题
|
||||
top_k: 返回示例数量
|
||||
min_rating: 最低评分(None表示不限制)
|
||||
required_tags: 必须包含的标签(如["aggregation", "join"])
|
||||
max_difficulty: 最大难度(过滤更难的示例)
|
||||
exclude_qids: 排除的QID(避免与当前问题相同)
|
||||
|
||||
Returns:
|
||||
排序后的示例列表(最相关优先)
|
||||
"""
|
||||
if self._embedder is None or self.embeddings is None:
|
||||
logger.error("Few-shot 索引未初始化")
|
||||
return []
|
||||
|
||||
if len(self.samples) == 0:
|
||||
return []
|
||||
|
||||
q_emb = self._embedder.encode(
|
||||
[question],
|
||||
batch_size=1,
|
||||
normalize=True,
|
||||
show_progress=False,
|
||||
)[0]
|
||||
|
||||
scores = np.dot(self.embeddings, q_emb)
|
||||
|
||||
candidates = []
|
||||
for idx, (score, sample) in enumerate(zip(scores, self.samples)):
|
||||
if exclude_qids and sample.qid in exclude_qids:
|
||||
continue
|
||||
if min_rating and sample.rating and sample.rating < min_rating:
|
||||
continue
|
||||
if max_difficulty == "easy" and sample.difficulty != "easy":
|
||||
continue
|
||||
if max_difficulty == "medium" and sample.difficulty == "hard":
|
||||
continue
|
||||
if required_tags and not all(tag in sample.tags for tag in required_tags):
|
||||
continue
|
||||
|
||||
candidates.append((idx, score, sample))
|
||||
|
||||
candidates.sort(key=lambda x: -x[1])
|
||||
|
||||
selected = [sample for _, _, sample in candidates[:top_k]]
|
||||
|
||||
logger.info(
|
||||
f"Few-shot选择: 问题='{question[:30]}...' "
|
||||
f"→ 选中{len(selected)}个示例 (top_k={top_k}, min_rating={min_rating})"
|
||||
)
|
||||
for s in selected:
|
||||
logger.debug(
|
||||
f" [{s.qid}] {s.question_zh[:50]}... (rating={s.rating}, tags={s.tags[:3]})"
|
||||
)
|
||||
|
||||
return selected
|
||||
|
||||
def get_examples_prompt(
|
||||
self,
|
||||
question: str,
|
||||
top_k: int = 3,
|
||||
min_rating: int = 7,
|
||||
**kwargs
|
||||
) -> str:
|
||||
"""
|
||||
生成few-shot prompt片段
|
||||
|
||||
Returns:
|
||||
格式化的示例字符串,可直接插入Prompt
|
||||
"""
|
||||
examples = self.select(question, top_k=top_k, min_rating=min_rating, **kwargs)
|
||||
|
||||
if not examples:
|
||||
return ""
|
||||
|
||||
lines = ["以下为相似问题的参考SQL示例:\n"]
|
||||
for i, ex in enumerate(examples, 1):
|
||||
lines.append(f"示例{i}:")
|
||||
lines.append(f"问题:{ex.question_zh}")
|
||||
lines.append(f"SQL:\n{ex.sql}")
|
||||
if ex.explanation:
|
||||
lines.append(f"说明:{ex.explanation[:150]}...")
|
||||
lines.append("") # 空行分隔
|
||||
|
||||
return "\n".join(lines)
|
||||
|
||||
def get_tagged_examples(self, tags: List[str], top_k_per_tag: int = 2) -> str:
|
||||
"""获取特定标签的示例"""
|
||||
tagged_samples = []
|
||||
for sample in self.samples:
|
||||
if any(tag in sample.tags for tag in tags):
|
||||
tagged_samples.append(sample)
|
||||
|
||||
tagged_samples.sort(key=lambda s: -(s.rating or 0))
|
||||
selected = tagged_samples[:top_k_per_tag * len(tags)]
|
||||
|
||||
lines = [f"# {tags} 相关示例\n"]
|
||||
for ex in selected:
|
||||
lines.append(f"## {ex.qid}. {ex.question_zh[:50]}")
|
||||
lines.append(f"评分: {ex.rating}/10")
|
||||
lines.append(f"标签: {', '.join(ex.tags)}")
|
||||
lines.append(f"```sql\n{ex.sql}\n```\n")
|
||||
|
||||
return "\n".join(lines)
|
||||
|
||||
def get_stats(self) -> Dict:
|
||||
"""获取数据集统计"""
|
||||
stats = {
|
||||
"total": len(self.samples),
|
||||
"by_rating": {},
|
||||
"by_difficulty": {},
|
||||
"by_tag": {},
|
||||
"avg_rating": 0.0
|
||||
}
|
||||
|
||||
ratings = [s.rating for s in self.samples 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)
|
||||
|
||||
for diff in ["easy", "medium", "hard"]:
|
||||
stats["by_difficulty"][diff] = sum(
|
||||
1 for s in self.samples if s.difficulty == diff
|
||||
)
|
||||
|
||||
tag_counts = {}
|
||||
for s in self.samples:
|
||||
for tag in s.tags:
|
||||
tag_counts[tag] = tag_counts.get(tag, 0) + 1
|
||||
stats["by_tag"] = tag_counts
|
||||
|
||||
return stats
|
||||
|
||||
|
||||
def load_fewshot_selector() -> FewShotSelector:
|
||||
"""加载默认的few-shot选择器"""
|
||||
# __file__ = backend/utils/fewshot_selector.py → 仓库根为 parents[2]
|
||||
default_path = Path(__file__).resolve().parents[2] / "data" / "experiences" / "all_samples.jsonl"
|
||||
return FewShotSelector(str(default_path))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import argparse
|
||||
|
||||
parser = argparse.ArgumentParser(description="Few-shot示例选择器")
|
||||
parser.add_argument("--samples", default="data/experiences/all_samples.jsonl")
|
||||
parser.add_argument("--question", help="测试问题")
|
||||
parser.add_argument("--top-k", type=int, default=3)
|
||||
parser.add_argument("--min-rating", type=int, default=7)
|
||||
parser.add_argument("--stats", action="store_true", help="显示数据集统计")
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
selector = FewShotSelector(args.samples)
|
||||
|
||||
if args.stats:
|
||||
stats = selector.get_stats()
|
||||
print("📊 数据集统计:")
|
||||
print(f" 总样本: {stats['total']}")
|
||||
print(f" 平均评分: {stats['avg_rating']:.1f}")
|
||||
print(f" 难度分布: {stats['by_difficulty']}")
|
||||
print(f"\n Top 10 标签:")
|
||||
sorted_tags = sorted(stats["by_tag"].items(), key=lambda x: -x[1])[:10]
|
||||
for tag, count in sorted_tags:
|
||||
print(f" {tag}: {count}")
|
||||
elif args.question:
|
||||
examples = selector.select(args.question, top_k=args.top_k, min_rating=args.min_rating)
|
||||
print(f"\n为问题 '{args.question}' 选择的示例:\n")
|
||||
for ex in examples:
|
||||
print(f"[{ex.qid}] 评分:{ex.rating} 难度:{ex.difficulty}")
|
||||
print(f"问题: {ex.question_zh}")
|
||||
print(f"SQL:\n{ex.sql}\n")
|
||||
else:
|
||||
print("请指定 --question 或 --stats")
|
||||
@@ -0,0 +1,23 @@
|
||||
"""自然语言问题语种启发式(用于是否走「英译中」再 Text2SQL)。"""
|
||||
|
||||
import re
|
||||
|
||||
|
||||
_CJK_RE = re.compile(r"[\u4e00-\u9fff]")
|
||||
_HANGUL_RE = re.compile(r"[\uac00-\ud7af]")
|
||||
_KANA_RE = re.compile(r"[\u3040-\u30ff]")
|
||||
|
||||
|
||||
def looks_like_english_only(text: str) -> bool:
|
||||
"""
|
||||
判断问题是否主要为英文(无中日韩表意文字),适合先译成中文再走检索/选表。
|
||||
|
||||
含中文、日文假名、韩文时不翻译,避免破坏中英混合问句。
|
||||
"""
|
||||
s = (text or "").strip()
|
||||
if not s or len(s) < 2:
|
||||
return False
|
||||
if _CJK_RE.search(s) or _HANGUL_RE.search(s) or _KANA_RE.search(s):
|
||||
return False
|
||||
latin = sum(1 for c in s if ("a" <= c <= "z") or ("A" <= c <= "Z"))
|
||||
return latin >= 3
|
||||
@@ -0,0 +1,364 @@
|
||||
"""
|
||||
SQL 解析与验证工具(基于 sqlglot)
|
||||
"""
|
||||
|
||||
import logging
|
||||
import re
|
||||
from typing import List, Tuple, Optional, Dict
|
||||
import sqlglot
|
||||
from sqlglot import exp, parse_one, ParseError
|
||||
|
||||
from schema.manager import SchemaManager
|
||||
from schema.models import Table, Column
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def validate_sql_syntax(sql: str, dialect: str = "tsql") -> Tuple[bool, List[str]]:
|
||||
"""
|
||||
验证SQL语法是否正确
|
||||
|
||||
Args:
|
||||
sql: SQL语句
|
||||
dialect: SQL方言
|
||||
|
||||
Returns:
|
||||
(是否有效, 错误信息列表)
|
||||
"""
|
||||
errors = []
|
||||
|
||||
try:
|
||||
# 尝试解析
|
||||
parsed = parse_one(sql, dialect=dialect)
|
||||
|
||||
if parsed is None:
|
||||
errors.append("SQL解析返回空结果")
|
||||
return False, errors
|
||||
|
||||
# 检查是否为只读查询(SELECT/CTE/SHOW等)
|
||||
# 动态获取可用表达式类型(兼容不同sqlglot版本)
|
||||
readable_ops = [exp.Select, exp.Union, exp.Intersect, exp.Except, exp.With]
|
||||
# 可选:添加 Show, Describe, Explain(如果存在)
|
||||
for op_name in ['Show', 'Describe', 'Explain']:
|
||||
if hasattr(exp, op_name):
|
||||
readable_ops.append(getattr(exp, op_name))
|
||||
|
||||
if not isinstance(parsed, tuple(readable_ops)):
|
||||
op_type = type(parsed).__name__
|
||||
errors.append(f"非查询操作({op_type}),只允许SELECT等只读语句")
|
||||
|
||||
return True, []
|
||||
|
||||
except ParseError as e:
|
||||
errors.append(f"SQL语法错误: {str(e)}")
|
||||
return False, errors
|
||||
except Exception as e:
|
||||
errors.append(f"解析异常: {str(e)}")
|
||||
return False, errors
|
||||
|
||||
|
||||
def extract_tables_from_sql(sql: str, dialect: str = "tsql") -> List[str]:
|
||||
"""
|
||||
从SQL中提取所有表名
|
||||
|
||||
Args:
|
||||
sql: SQL语句
|
||||
dialect: SQL方言
|
||||
|
||||
Returns:
|
||||
表名列表(去重)
|
||||
"""
|
||||
try:
|
||||
parsed = parse_one(sql, dialect=dialect)
|
||||
tables = []
|
||||
|
||||
# 遍历AST查找所有表名
|
||||
for node in parsed.walk():
|
||||
if isinstance(node, exp.Table):
|
||||
table_name = node.name
|
||||
if table_name and table_name not in tables:
|
||||
tables.append(table_name)
|
||||
|
||||
return tables
|
||||
except Exception as e:
|
||||
logger.warning(f"提取表名失败: {e}")
|
||||
return []
|
||||
|
||||
|
||||
def build_table_alias_map(parsed: exp.Expression) -> Dict[str, str]:
|
||||
"""
|
||||
从已解析的 AST 构建「别名/表名 -> 物理表名」映射。
|
||||
|
||||
FROM T a 时 a -> T,且 T -> T,便于将 a.col 解析到表 T 的列。
|
||||
"""
|
||||
alias_map: Dict[str, str] = {}
|
||||
for node in parsed.walk():
|
||||
if not isinstance(node, exp.Table):
|
||||
continue
|
||||
physical = node.name
|
||||
if not physical:
|
||||
continue
|
||||
alias_map[physical] = physical
|
||||
talias = node.args.get("alias")
|
||||
if talias is not None:
|
||||
aname = talias.name
|
||||
if aname:
|
||||
alias_map[aname] = physical
|
||||
return alias_map
|
||||
|
||||
|
||||
def extract_columns_from_sql(sql: str, dialect: str = "tsql") -> List[Tuple[str, str]]:
|
||||
"""
|
||||
从SQL中提取所有字段引用(表.字段)
|
||||
|
||||
Args:
|
||||
sql: SQL语句
|
||||
dialect: SQL方言
|
||||
|
||||
Returns:
|
||||
[(表名, 字段名), ...] 列表
|
||||
"""
|
||||
columns = []
|
||||
|
||||
try:
|
||||
parsed = parse_one(sql, dialect=dialect)
|
||||
|
||||
for node in parsed.walk():
|
||||
if isinstance(node, exp.Column):
|
||||
table_name = node.table
|
||||
col_name = node.name
|
||||
if table_name and col_name:
|
||||
columns.append((table_name, col_name))
|
||||
|
||||
return columns
|
||||
except Exception as e:
|
||||
logger.warning(f"提取字段失败: {e}")
|
||||
return []
|
||||
|
||||
|
||||
def validate_schema_consistency(
|
||||
sql: str,
|
||||
schema_manager: SchemaManager,
|
||||
dialect: str = "tsql"
|
||||
) -> Tuple[bool, List[str]]:
|
||||
"""
|
||||
验证SQL与Schema的一致性
|
||||
|
||||
检查:
|
||||
1. 所有表名存在于Schema
|
||||
2. 所有字段名属于对应的表
|
||||
3. JOIN条件字段存在且类型兼容
|
||||
|
||||
Args:
|
||||
sql: SQL语句
|
||||
schema_manager: Schema管理器
|
||||
dialect: SQL方言
|
||||
|
||||
Returns:
|
||||
(是否一致, 错误信息列表)
|
||||
"""
|
||||
errors = []
|
||||
|
||||
try:
|
||||
parsed = parse_one(sql, dialect=dialect)
|
||||
except Exception as e:
|
||||
logger.debug(f"Schema一致性检查跳过(解析失败): {e}")
|
||||
return True, []
|
||||
|
||||
alias_map = build_table_alias_map(parsed)
|
||||
|
||||
tables_used: List[str] = []
|
||||
for node in parsed.walk():
|
||||
if isinstance(node, exp.Table):
|
||||
tname = node.name
|
||||
if tname and tname not in tables_used:
|
||||
tables_used.append(tname)
|
||||
|
||||
columns_used: List[Tuple[str, str]] = []
|
||||
for node in parsed.walk():
|
||||
if isinstance(node, exp.Column):
|
||||
tref, cname = node.table, node.name
|
||||
if tref and cname:
|
||||
columns_used.append((tref, cname))
|
||||
|
||||
# 检查表存在性
|
||||
for tbl in tables_used:
|
||||
if not schema_manager.get_table(tbl):
|
||||
errors.append(f"表不存在: '{tbl}'")
|
||||
|
||||
# 检查字段存在性(表引用可为物理表名或别名)
|
||||
for tbl_name, col_name in columns_used:
|
||||
physical = alias_map.get(tbl_name, tbl_name)
|
||||
table = schema_manager.get_table(physical)
|
||||
if not table:
|
||||
errors.append(
|
||||
f"字段不存在: '{tbl_name}.{col_name}'"
|
||||
f"(无法将表引用解析到已加载Schema中的表)"
|
||||
)
|
||||
continue
|
||||
col_names = [c.name for c in table.columns]
|
||||
if col_name not in col_names:
|
||||
errors.append(
|
||||
f"字段不存在: '{tbl_name}.{col_name}'"
|
||||
f"(表 '{physical}' 可用字段: {col_names[:5]}...)"
|
||||
)
|
||||
|
||||
# 检查JOIN条件(外键匹配)
|
||||
try:
|
||||
for join in parsed.find_all(exp.Join):
|
||||
# 解析ON条件
|
||||
on_condition = join.args.get("on")
|
||||
if on_condition:
|
||||
# 检查ON条件中涉及的字段
|
||||
for eq in on_condition.find_all(exp.EQ):
|
||||
left = eq.left
|
||||
right = eq.right
|
||||
|
||||
# 提取左右两边的表.字段
|
||||
for side in [left, right]:
|
||||
if isinstance(side, exp.Column):
|
||||
tbl = side.table
|
||||
col = side.name
|
||||
physical = alias_map.get(tbl, tbl)
|
||||
table = schema_manager.get_table(physical)
|
||||
if table and col not in [c.name for c in table.columns]:
|
||||
errors.append(f"JOIN条件字段不存在: {tbl}.{col}")
|
||||
except Exception as e:
|
||||
logger.debug(f"JOIN条件检查异常: {e}")
|
||||
|
||||
return len(errors) == 0, errors
|
||||
|
||||
|
||||
def rewrite_mysql_builtins_for_tsql(sql: str) -> str:
|
||||
"""
|
||||
模型在 T-SQL 目标下仍常输出 MySQL 函数;sqlglot 转写也可能遗漏。
|
||||
SQL Server 无 CURDATE()/NOW(),需替换为 GETDATE 族。
|
||||
"""
|
||||
if not sql:
|
||||
return sql
|
||||
out = sql
|
||||
out = re.sub(
|
||||
r"\bCURDATE\s*\(\s*\)",
|
||||
"CAST(GETDATE() AS DATE)",
|
||||
out,
|
||||
flags=re.IGNORECASE,
|
||||
)
|
||||
out = re.sub(
|
||||
r"\bCURRENT_DATE\b",
|
||||
"CAST(GETDATE() AS DATE)",
|
||||
out,
|
||||
flags=re.IGNORECASE,
|
||||
)
|
||||
out = re.sub(r"\bNOW\s*\(\s*\)", "GETDATE()", out, flags=re.IGNORECASE)
|
||||
return out
|
||||
|
||||
|
||||
def normalize_sql_for_dialect(sql: str, dialect: str) -> str:
|
||||
"""
|
||||
将模型输出的 SQL 规范为目标方言。
|
||||
对于 T-SQL,主要进行 MySQL 函数替换(因为模型仍可能输出 CURDATE() 等)。
|
||||
"""
|
||||
sql = (sql or "").strip()
|
||||
if not sql:
|
||||
return sql
|
||||
|
||||
# 如果目标是 T-SQL,只做函数名替换,不再用 sqlglot 转写
|
||||
if dialect == "tsql":
|
||||
return rewrite_mysql_builtins_for_tsql(sql)
|
||||
|
||||
return sql
|
||||
|
||||
|
||||
def format_sql(sql: str, dialect: str = "tsql", indent: int = 2) -> str:
|
||||
"""
|
||||
格式化SQL(可读性)
|
||||
|
||||
Args:
|
||||
sql: SQL语句
|
||||
dialect: SQL方言
|
||||
indent: 缩进空格数
|
||||
|
||||
Returns:
|
||||
格式化后的SQL
|
||||
"""
|
||||
try:
|
||||
parsed = parse_one(sql, dialect=dialect)
|
||||
return parsed.sql(dialect=dialect, pretty=True, indent=indent)
|
||||
except Exception as e:
|
||||
logger.warning(f"SQL格式化失败: {e}")
|
||||
return sql
|
||||
|
||||
|
||||
def normalize_sql(sql: str, dialect: str = "tsql") -> str:
|
||||
"""
|
||||
标准化SQL(用于比较去重)
|
||||
|
||||
去除多余空格、统一引号、移除注释等
|
||||
|
||||
Args:
|
||||
sql: SQL语句
|
||||
dialect: SQL方言
|
||||
|
||||
Returns:
|
||||
标准化后的SQL
|
||||
"""
|
||||
try:
|
||||
# 解析后重新生成(会规范化格式)
|
||||
parsed = parse_one(sql, dialect=dialect)
|
||||
normalized = parsed.sql(dialect=dialect, pretty=False)
|
||||
# 转换为大写关键词
|
||||
return normalized.upper()
|
||||
except Exception:
|
||||
# 降级:简单处理
|
||||
import re
|
||||
# 移除多余空格
|
||||
sql = re.sub(r'\s+', ' ', sql.strip())
|
||||
# 移除注释
|
||||
sql = re.sub(r'--.*?$', '', sql, flags=re.MULTILINE)
|
||||
sql = re.sub(r'/\*.*?\*/', '', sql, flags=re.DOTALL)
|
||||
return sql.upper()
|
||||
|
||||
|
||||
def count_joins(sql: str, dialect: str = "tsql") -> int:
|
||||
"""统计JOIN数量"""
|
||||
try:
|
||||
parsed = parse_one(sql, dialect=dialect)
|
||||
joins = list(parsed.find_all(exp.Join))
|
||||
return len(joins)
|
||||
except Exception:
|
||||
return 0
|
||||
|
||||
|
||||
def has_subquery(sql: str, dialect: str = "tsql") -> bool:
|
||||
"""检查是否包含子查询"""
|
||||
try:
|
||||
parsed = parse_one(sql, dialect=dialect)
|
||||
# 检查嵌套的SELECT
|
||||
for select in parsed.find_all(exp.Select):
|
||||
if select is not parsed: # 不是最外层的SELECT
|
||||
return True
|
||||
return False
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
def get_query_complexity(sql: str, dialect: str = "tsql") -> Dict[str, int]:
|
||||
"""
|
||||
评估查询复杂度
|
||||
|
||||
Returns:
|
||||
复杂度指标字典
|
||||
"""
|
||||
try:
|
||||
parsed = parse_one(sql, dialect=dialect)
|
||||
|
||||
return {
|
||||
"join_count": len(list(parsed.find_all(exp.Join))),
|
||||
"subquery_count": len([s for s in parsed.find_all(exp.Select) if s is not parsed]),
|
||||
"where_conditions": len(list(parsed.find_all(exp.Predicate))),
|
||||
"aggregation_functions": len(list(parsed.find_all(exp.AggFunc))),
|
||||
"column_count": len(list(parsed.find_all(exp.Column))),
|
||||
}
|
||||
except Exception as e:
|
||||
logger.warning(f"复杂度评估失败: {e}")
|
||||
return {}
|
||||
@@ -0,0 +1,430 @@
|
||||
"""
|
||||
SQL 验证工具集
|
||||
"""
|
||||
|
||||
import re
|
||||
import logging
|
||||
from typing import Tuple, List, Dict
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# CJK Unified Ideographs + 兼容扩展(用于禁止中文业务词出现在 SQL 字符串字面量中)
|
||||
_CJK_IN_STRING_RE = re.compile(
|
||||
r"[\u3000-\u303f\u3400-\u4dbf\u4e00-\u9fff\uf900-\ufaff\uff00-\uffef]"
|
||||
)
|
||||
|
||||
|
||||
def check_no_cjk_in_sql_string_literals(sql: str) -> Tuple[bool, List[str]]:
|
||||
"""
|
||||
扫描 SQL 中单引号字符串(含 T-SQL N'…'),若字面量内出现 CJK 则判失败。
|
||||
|
||||
跳过 ``--`` 行注释与 ``/* */`` 块注释内的文本,避免误报。
|
||||
"""
|
||||
errors: List[str] = []
|
||||
i = 0
|
||||
n = len(sql)
|
||||
in_line_comment = False
|
||||
in_block_comment = False
|
||||
|
||||
def _read_single_quoted_string(start: int) -> Tuple[str, int]:
|
||||
"""从 start 指向的 opening `'` 之后开始读,返回 (内容, 闭合引号后下标)。"""
|
||||
j = start
|
||||
parts: List[str] = []
|
||||
while j < n:
|
||||
ch = sql[j]
|
||||
if ch == "'":
|
||||
if j + 1 < n and sql[j + 1] == "'":
|
||||
parts.append("'")
|
||||
j += 2
|
||||
continue
|
||||
return "".join(parts), j + 1
|
||||
parts.append(ch)
|
||||
j += 1
|
||||
return "".join(parts), j
|
||||
|
||||
while i < n:
|
||||
if in_line_comment:
|
||||
if sql[i] == "\n":
|
||||
in_line_comment = False
|
||||
i += 1
|
||||
continue
|
||||
if in_block_comment:
|
||||
if i + 1 < n and sql[i : i + 2] == "*/":
|
||||
in_block_comment = False
|
||||
i += 2
|
||||
else:
|
||||
i += 1
|
||||
continue
|
||||
|
||||
two = sql[i : i + 2]
|
||||
if two == "--":
|
||||
in_line_comment = True
|
||||
i += 2
|
||||
continue
|
||||
if two == "/*":
|
||||
in_block_comment = True
|
||||
i += 2
|
||||
continue
|
||||
|
||||
# N' 或 n' 前缀的 Unicode 字面量
|
||||
if i + 1 < n and sql[i] in "Nn" and sql[i + 1] == "'":
|
||||
body, i = _read_single_quoted_string(i + 2)
|
||||
if _CJK_IN_STRING_RE.search(body):
|
||||
prev = body[:48] + ("…" if len(body) > 48 else "")
|
||||
errors.append(
|
||||
"SQL 字符串字面量中含中文或与业务中文直接作为比对值(禁止)。"
|
||||
f"片段近似: …'{prev}'… — 请改用 Schema 注释中的代码/枚举,或通过维表 JOIN 用键列过滤。"
|
||||
)
|
||||
continue
|
||||
|
||||
if sql[i] == "'":
|
||||
body, i = _read_single_quoted_string(i + 1)
|
||||
if _CJK_IN_STRING_RE.search(body):
|
||||
prev = body[:48] + ("…" if len(body) > 48 else "")
|
||||
errors.append(
|
||||
"SQL 字符串字面量中含中文或与业务中文直接作为比对值(禁止)。"
|
||||
f"片段近似: …'{prev}'… — 请改用 Schema 注释中的代码/枚举,或通过维表 JOIN 用键列过滤。"
|
||||
)
|
||||
continue
|
||||
|
||||
i += 1
|
||||
|
||||
return len(errors) == 0, errors
|
||||
|
||||
|
||||
# 危险操作关键词(除非明确允许)
|
||||
DANGEROUS_KEYWORDS = [
|
||||
"DROP", "DELETE", "UPDATE", "INSERT", "ALTER", "TRUNCATE",
|
||||
"CREATE", "DROP DATABASE", "DROP TABLE", "DROP INDEX",
|
||||
"GRANT", "REVOKE", "PURGE", "FLUSH", "KILL"
|
||||
]
|
||||
|
||||
# 允许的操作(仅查询)
|
||||
ALLOWED_KEYWORDS = [
|
||||
"SELECT", "WITH", "FROM", "WHERE", "JOIN", "LEFT JOIN", "RIGHT JOIN",
|
||||
"INNER JOIN", "OUTER JOIN", "ON", "USING", "GROUP BY", "HAVING",
|
||||
"ORDER BY", "LIMIT", "OFFSET", "UNION", "UNION ALL", "EXCEPT", "INTERSECT",
|
||||
"AS", "CASE", "WHEN", "THEN", "ELSE", "END",
|
||||
"COUNT", "SUM", "AVG", "MIN", "MAX", "DISTINCT",
|
||||
"AND", "OR", "NOT", "IN", "EXISTS", "BETWEEN", "LIKE", "IS NULL", "IS NOT NULL",
|
||||
"CAST", "COALESCE", "NULLIF", "IFNULL",
|
||||
"DATE", "TIME", "TIMESTAMP", "EXTRACT", "DATE_FORMAT", "STR_TO_DATE",
|
||||
"CURRENT_DATE", "CURRENT_TIMESTAMP",
|
||||
]
|
||||
|
||||
|
||||
def check_dangerous_operations(sql: str) -> Tuple[bool, List[str]]:
|
||||
"""
|
||||
检查SQL是否包含危险操作
|
||||
|
||||
Args:
|
||||
sql: SQL语句(大小写不敏感)
|
||||
|
||||
Returns:
|
||||
(是否安全, 危险关键词列表)
|
||||
"""
|
||||
sql_upper = sql.upper()
|
||||
found_dangers = []
|
||||
|
||||
for keyword in DANGEROUS_KEYWORDS:
|
||||
# 使用正则避免部分匹配(如"DROP"不应匹配"DROPOUT")
|
||||
pattern = r'\b' + re.escape(keyword) + r'\b'
|
||||
if re.search(pattern, sql_upper):
|
||||
found_dangers.append(keyword)
|
||||
|
||||
is_safe = len(found_dangers) == 0
|
||||
|
||||
if not is_safe:
|
||||
logger.warning(f"检测到危险操作: {found_dangers}")
|
||||
|
||||
return is_safe, found_dangers
|
||||
|
||||
|
||||
def validate_no_dml(sql: str) -> Tuple[bool, str]:
|
||||
"""
|
||||
验证SQL不是DML/DDL操作(仅允许SELECT等查询)
|
||||
|
||||
Returns:
|
||||
(是否通过, 错误消息)
|
||||
"""
|
||||
sql_upper = sql.strip().upper()
|
||||
|
||||
# 检查是否以危险关键词开头
|
||||
first_word = sql_upper.split()[0] if sql_upper.split() else ""
|
||||
if first_word in ["INSERT", "UPDATE", "DELETE", "DROP", "ALTER", "CREATE", "TRUNCATE"]:
|
||||
return False, f"禁止的操作: {first_word}"
|
||||
|
||||
is_safe, dangers = check_dangerous_operations(sql)
|
||||
if not is_safe:
|
||||
return False, f"SQL包含危险操作: {', '.join(dangers)}"
|
||||
|
||||
return True, ""
|
||||
|
||||
|
||||
def check_sql_injection_patterns(sql: str) -> List[str]:
|
||||
"""
|
||||
检查明显的SQL注入模式
|
||||
|
||||
Args:
|
||||
sql: SQL语句
|
||||
|
||||
Returns:
|
||||
发现的注入模式列表
|
||||
"""
|
||||
patterns = {
|
||||
"union_all_injection": r"UNION\s+ALL\s+SELECT",
|
||||
"union_injection": r"UNION\s+SELECT",
|
||||
"comment_injection": r"(--|\#|/\*).*SELECT",
|
||||
"semicolon_injection": r";\s*(DROP|DELETE|UPDATE|INSERT)",
|
||||
"or_true_condition": r"OR\s+['\"]?\s*1\s*['\"]?\s*=\s*1",
|
||||
"always_true": r"1\s*=\s*1",
|
||||
}
|
||||
|
||||
findings = []
|
||||
sql_lower = sql.lower()
|
||||
|
||||
for name, pattern in patterns.items():
|
||||
if re.search(pattern, sql, re.IGNORECASE):
|
||||
findings.append(name)
|
||||
|
||||
return findings
|
||||
|
||||
|
||||
def validate_aggregation_groupby(sql: str, dialect: str = "tsql") -> Tuple[bool, List[str]]:
|
||||
"""
|
||||
验证聚合查询的GROUP BY正确性
|
||||
|
||||
检查:SELECT中的非聚合字段是否都在GROUP BY中
|
||||
|
||||
Args:
|
||||
sql: SQL语句
|
||||
dialect: SQL方言
|
||||
|
||||
Returns:
|
||||
(是否有效, 错误列表)
|
||||
"""
|
||||
from utils.sql_parser import parse_one, exp
|
||||
|
||||
errors = []
|
||||
|
||||
try:
|
||||
parsed = parse_one(sql, dialect=dialect)
|
||||
|
||||
# 只检查SELECT语句
|
||||
if not isinstance(parsed, exp.Select):
|
||||
return True, []
|
||||
|
||||
# 获取SELECT列表中的表达式
|
||||
select_exprs = parsed.expressions
|
||||
|
||||
# 获取GROUP BY字段
|
||||
group_by = parsed.args.get("group")
|
||||
if not group_by:
|
||||
# 没有GROUP BY但有聚合函数,通常是错误的
|
||||
has_agg = any(
|
||||
expr.find(exp.AggFunc) is not None
|
||||
for expr in select_exprs
|
||||
)
|
||||
if has_agg:
|
||||
errors.append("包含聚合函数但缺少GROUP BY子句")
|
||||
return len(errors) == 0, errors
|
||||
|
||||
group_by_exprs = group_by.expressions
|
||||
|
||||
# 提取GROUP BY的字段名(简单处理)
|
||||
group_by_cols = set()
|
||||
for expr in group_by_exprs:
|
||||
if isinstance(expr, exp.Column):
|
||||
group_by_cols.add(expr.name)
|
||||
elif isinstance(expr, exp.Ordered):
|
||||
# GROUP BY x ASC/DESC
|
||||
this = expr.this
|
||||
if isinstance(this, exp.Column):
|
||||
group_by_cols.add(this.name)
|
||||
|
||||
# 检查每个SELECT表达式
|
||||
for expr in select_exprs:
|
||||
# 如果是聚合函数,跳过
|
||||
if expr.find(exp.AggFunc):
|
||||
continue
|
||||
|
||||
# 如果是字面量或表达式,跳过
|
||||
if isinstance(expr, exp.Literal):
|
||||
continue
|
||||
|
||||
# 如果是列引用,检查是否在GROUP BY中
|
||||
if isinstance(expr, exp.Column):
|
||||
col_name = expr.name
|
||||
if col_name not in group_by_cols:
|
||||
errors.append(
|
||||
f"字段 '{col_name}' 在SELECT中但不在GROUP BY中"
|
||||
)
|
||||
elif isinstance(expr, exp.Alias):
|
||||
# 别名: column AS alias
|
||||
this = expr.this
|
||||
if isinstance(this, exp.Column):
|
||||
col_name = this.name
|
||||
if col_name not in group_by_cols:
|
||||
errors.append(
|
||||
f"字段 '{col_name}' (别名为'{expr.alias}') 在SELECT中但不在GROUP BY中"
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
logger.debug(f"GROUP BY验证异常: {e}")
|
||||
|
||||
return len(errors) == 0, errors
|
||||
|
||||
|
||||
def check_join_conditions(sql: str, dialect: str = "tsql") -> List[str]:
|
||||
"""
|
||||
检查JOIN条件是否完整
|
||||
|
||||
Args:
|
||||
sql: SQL语句
|
||||
dialect: SQL方言
|
||||
|
||||
Returns:
|
||||
问题列表(空表示无问题)
|
||||
"""
|
||||
from utils.sql_parser import parse_one, exp
|
||||
|
||||
issues = []
|
||||
|
||||
try:
|
||||
parsed = parse_one(sql, dialect=dialect)
|
||||
|
||||
# 遍历所有JOIN
|
||||
for join in parsed.find_all(exp.Join):
|
||||
# 检查是否有ON条件
|
||||
on_condition = join.args.get("on")
|
||||
if on_condition is None:
|
||||
# 检查是否使用USING
|
||||
using = join.args.get("using")
|
||||
if using is None:
|
||||
issues.append("JOIN缺少ON条件")
|
||||
else:
|
||||
# ON条件为空表达式
|
||||
if isinstance(on_condition, exp.Empty):
|
||||
issues.append("JOIN的ON条件为空")
|
||||
|
||||
except Exception as e:
|
||||
logger.debug(f"JOIN条件检查异常: {e}")
|
||||
|
||||
return issues
|
||||
|
||||
|
||||
def validate_order_by_fields(
|
||||
sql: str,
|
||||
schema_manager,
|
||||
dialect: str = "tsql"
|
||||
) -> List[str]:
|
||||
"""
|
||||
验证ORDER BY字段是否存在于对应表中
|
||||
|
||||
Args:
|
||||
sql: SQL语句
|
||||
schema_manager: Schema管理器
|
||||
dialect: SQL方言
|
||||
|
||||
Returns:
|
||||
问题列表
|
||||
"""
|
||||
from utils.sql_parser import parse_one, exp
|
||||
|
||||
issues = []
|
||||
|
||||
try:
|
||||
parsed = parse_one(sql, dialect=dialect)
|
||||
order = parsed.args.get("order")
|
||||
|
||||
if order:
|
||||
for ordered in order.expressions:
|
||||
expr = ordered.this
|
||||
|
||||
# 提取字段和表
|
||||
if isinstance(expr, exp.Column):
|
||||
tbl_name = expr.table
|
||||
col_name = expr.name
|
||||
|
||||
if tbl_name:
|
||||
table = schema_manager.get_table(tbl_name)
|
||||
if table:
|
||||
col_names = [c.name for c in table.columns]
|
||||
if col_name not in col_names:
|
||||
issues.append(
|
||||
f"ORDER BY字段不存在: {tbl_name}.{col_name}"
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
logger.debug(f"ORDER BY验证异常: {e}")
|
||||
|
||||
return issues
|
||||
|
||||
|
||||
def full_validation_pipeline(
|
||||
sql: str,
|
||||
schema_manager,
|
||||
dialect: str = "tsql",
|
||||
check_dangerous: bool = True
|
||||
) -> Dict:
|
||||
"""
|
||||
完整验证流水线
|
||||
|
||||
Args:
|
||||
sql: SQL语句
|
||||
schema_manager: Schema管理器
|
||||
dialect: SQL方言
|
||||
check_dangerous: 是否检查危险操作
|
||||
|
||||
Returns:
|
||||
验证结果字典
|
||||
"""
|
||||
result = {
|
||||
"valid": True,
|
||||
"errors": [],
|
||||
"warnings": [],
|
||||
"suggestions": []
|
||||
}
|
||||
|
||||
# 1. 语法验证
|
||||
syntax_ok, syntax_errors = validate_sql_syntax(sql, dialect)
|
||||
if not syntax_ok:
|
||||
result["valid"] = False
|
||||
result["errors"].extend(syntax_errors)
|
||||
|
||||
# 2. 危险操作检查
|
||||
if check_dangerous:
|
||||
safe, dangers = check_dangerous_operations(sql)
|
||||
if not safe:
|
||||
result["valid"] = False
|
||||
result["errors"].append(f"包含危险操作: {', '.join(dangers)}")
|
||||
|
||||
# 3. Schema一致性验证
|
||||
schema_ok, schema_errors = validate_schema_consistency(sql, schema_manager, dialect)
|
||||
if not schema_ok:
|
||||
result["valid"] = False
|
||||
result["errors"].extend(schema_errors)
|
||||
|
||||
# 4. GROUP BY验证
|
||||
groupby_ok, groupby_errors = validate_aggregation_groupby(sql, dialect)
|
||||
if not groupby_ok:
|
||||
result["valid"] = False
|
||||
result["errors"].extend(groupby_errors)
|
||||
|
||||
# 5. JOIN条件验证
|
||||
join_issues = check_join_conditions(sql, dialect)
|
||||
if join_issues:
|
||||
result["valid"] = False
|
||||
result["errors"].extend(join_issues)
|
||||
|
||||
# 6. ORDER BY验证
|
||||
order_issues = validate_order_by_fields(sql, schema_manager, dialect)
|
||||
if order_issues:
|
||||
result["warnings"].extend(order_issues)
|
||||
|
||||
# 7. SQL注入模式检查(警告)
|
||||
injection_patterns = check_sql_injection_patterns(sql)
|
||||
if injection_patterns:
|
||||
result["warnings"].append(f"检测到可疑模式: {', '.join(injection_patterns)}")
|
||||
|
||||
return result
|
||||
Reference in New Issue
Block a user