Files
ai-g3sb-backman2.0/agents/orchestrator.py
T

559 lines
19 KiB
Python
Raw Normal View History

2026-04-10 16:52:07 +08:00
"""
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. Schema Linker:粗筛 + LLM精筛,选出相关表
2. 外键扩展:自动包含关联表
3. SQL Generator:生成SQL
4. Validator:验证SQL,不通过则重试(最多max_retry次)
"""
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,
):
"""
初始化编排器
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: 是否使用向量检索粗筛
"""
self.schema_manager = schema_manager
self.max_retry = max_retry
self.use_vector_search = use_vector_search
# 初始化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 "")
)
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"
) -> str:
"""
SQL生成(SQL Generator Agent)
Args:
question: 用户问题
schema_str: Schema描述字符串
dialect: SQL方言
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 英文别名。"
)
messages = [
{"role": "system", "content": SQL_GENERATOR_SYSTEM},
{"role": "user", "content": user_content},
]
response = self.deepseek.chat(messages)
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",
) -> Tuple[bool, List[str], List[str]]:
"""
SQL验证(Validator Agent + 程序验证)
Args:
sql: SQL语句
schema_str: Schema描述
dialect: 与生成一致的 SQL 方言(sqlglot 名,默认 tsql)
Returns:
(是否通过, 错误列表, 警告列表)
"""
errors = []
warnings = []
# === 阶段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
danger_ok, danger_errors = check_dangerous_operations(sql)
if not danger_ok:
errors.extend(danger_errors)
# === 阶段2: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}")
is_valid = len(errors) == 0
return is_valid, errors, warnings
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对象
"""
logger.info(f"[GEN] 开始生成SQL:{question[:50]}...")
attempt = 0
last_sql = None
last_errors = []
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
logger.info(f" 重试使用之前的Schema({len(tables_used)}张表)")
# === Step 2: SQL生成 ===
try:
sql = self._generate_sql(question, filtered_schema_str, dialect)
last_sql = sql
except Exception as e:
last_errors = [f"SQL生成失败: {str(e)}"]
attempt += 1
continue
# === Step 3: 验证 ===
is_valid, errors, warnings = self._validate_sql(
sql, filtered_schema_str, dialect=dialect
)
if is_valid:
logger.info(f"[OK] SQL生成并验证通过({attempt + 1}次尝试)")
result = GenerationResult(
sql=sql,
valid=True,
errors=[],
warnings=warnings,
tables_used=tables_used,
attempts=attempt + 1,
reasoning=reasoning if attempt == 0 else None,
)
if include_schema_in_result:
result.metadata["schema"] = filtered_schema_str
return result
# 验证失败,准备重试
last_errors = errors
logger.warning(f" [FAIL] 验证失败:{errors}")
attempt += 1
# 达到最大重试次数
logger.error(f"[FAIL] 达到最大重试次数({self.max_retry}),生成失败")
return GenerationResult(
sql=last_sql or "",
valid=False,
errors=last_errors,
tables_used=tables_used,
attempts=attempt,
)
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,
}