first commit
This commit is contained in:
@@ -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)
|
||||
Reference in New Issue
Block a user