164 lines
4.8 KiB
Python
164 lines
4.8 KiB
Python
"""
|
||
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)
|