first commit

This commit is contained in:
陈辅元
2026-04-10 16:52:07 +08:00
commit 84fe545640
87 changed files with 20842 additions and 0 deletions
+163
View File
@@ -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)