0.1.1 暂存

This commit is contained in:
陈辅元
2026-04-14 10:28:22 +08:00
parent 4a0638ba2e
commit cc83fe963a
83 changed files with 4005 additions and 179 deletions
+5
View File
@@ -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.
+1
View File
@@ -0,0 +1 @@
# agents 包初始化
Binary file not shown.
Binary file not shown.
+703
View File
@@ -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,
}
+163
View File
@@ -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)
+206
View File
@@ -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
+299
View File
@@ -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条
+1
View File
@@ -0,0 +1 @@
# config 包初始化
Binary file not shown.
Binary file not shown.
Binary file not shown.
+352
View File
@@ -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}
仅输出一句中文:"""
+52
View File
@@ -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()
+31
View File
@@ -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.
+87
View File
@@ -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]
+268
View File
@@ -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
+46
View File
@@ -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
+1
View File
@@ -0,0 +1 @@
# llm 包初始化
Binary file not shown.
+361
View File
@@ -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
View File
@@ -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()
+287
View File
@@ -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()
+1
View File
@@ -0,0 +1 @@
# schema 包初始化
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
+229
View File
@@ -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),
}
+319
View File
@@ -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)
+315
View File
@@ -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)})"
+179
View File
@@ -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],
}
+1
View File
@@ -0,0 +1 @@
# utils 包初始化
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
+135
View File
@@ -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)
+530
View File
@@ -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()
+357
View File
@@ -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")
+23
View File
@@ -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
+364
View File
@@ -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 {}
+430
View File
@@ -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