Files
ai-g3sb-backman2.0/agents/orchestrator.py
T
2026-04-10 16:52:07 +08:00

559 lines
19 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
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,
}