Files
ai-g3sb-backman2.0/backend/agents/schema_linker.py
T
2026-04-14 10:28:22 +08:00

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