feat: Phase 5 — AI Research Agent(受控工具白名单 + LLM 编排 + API)
- agent/tools.py:Tool 元数据(JSON Schema)+ 白名单调用(异常转可读反馈,不中断对话) - agent/tools_impl.py:6 个受控工具 search_stocks / get_market_data / test_factor / run_backtest / get_experiment / compare_experiments —— 全部只读经 Job/Experiment 链路,研究自动归档;无 shell/任意执行/写删数据能力 - agent/llm.py:LLMClient 抽象 + OpenAI 兼容客户端(LLM_API_KEY/LLM_BASE_URL/LLM_MODEL 走 .env,未配置给出引导提示)+ 研究纪律 system prompt(反过拟合/样本外/成本) - agent/service.py:编排循环(tool/final JSON 决策 → 执行 → 回喂 → 结论),轮次上限兜底,未知工具拒绝 - /api/agent/chat;httpx 移至主依赖;Job 默认工厂抽取(api/agent/executor 复用) - 测试 6 项(白名单无 shell、完整研究循环产出、未知工具拒绝、轮次兜底),全量 79 passed / ruff clean
This commit is contained in:
@@ -0,0 +1,98 @@
|
||||
"""Agent 编排:自然语言 → 受控工具调用循环 → 结论(AGENT.md §29 假设-实验-分析循环)。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import re
|
||||
from typing import Any
|
||||
|
||||
from app.agent.llm import SYSTEM_TEMPLATE, LLMClient
|
||||
from app.agent.tools import Tool, tools_schema
|
||||
from app.agent.tools_impl import build_tools
|
||||
|
||||
MAX_TOOL_ROUNDS = 5
|
||||
_DECISION_PATTERN = re.compile(r"\{.*\}", re.DOTALL)
|
||||
|
||||
|
||||
def _parse_decision(text: str) -> dict[str, Any]:
|
||||
"""容忍 LLM 输出中的代码块/前后缀,提取首个 JSON 对象。"""
|
||||
match = _DECISION_PATTERN.search(text)
|
||||
if not match:
|
||||
raise ValueError(f"无法从模型输出中解析动作 JSON:{text[:200]}")
|
||||
try:
|
||||
return json.loads(match.group(0))
|
||||
except json.JSONDecodeError as exc:
|
||||
raise ValueError(f"模型输出的 JSON 不合法:{text[:200]}") from exc
|
||||
|
||||
|
||||
class AgentService:
|
||||
"""受控研究 Agent:每轮让 LLM 决策 tool/final,执行工具并把结果回喂,直至 final。"""
|
||||
|
||||
def __init__(self, llm: LLMClient, tools: list[Tool] | None = None) -> None:
|
||||
self._llm = llm
|
||||
self._tools = {t.name: t for t in (tools or build_tools())}
|
||||
|
||||
def _by_name(self, name: str) -> Tool:
|
||||
tool = self._tools.get(name)
|
||||
if tool is None:
|
||||
raise ValueError(f"工具不存在:{name}(可用 {sorted(self._tools)})")
|
||||
return tool
|
||||
|
||||
def chat(self, message: str, max_rounds: int = MAX_TOOL_ROUNDS) -> dict:
|
||||
messages: list[dict] = [
|
||||
{
|
||||
"role": "system",
|
||||
"content": SYSTEM_TEMPLATE.format(tools=tools_schema(list(self._tools.values()))),
|
||||
},
|
||||
{"role": "user", "content": message},
|
||||
]
|
||||
actions: list[dict] = []
|
||||
|
||||
for _round in range(max_rounds):
|
||||
try:
|
||||
decision = _parse_decision(self._llm.chat(messages))
|
||||
except ValueError as exc:
|
||||
# LLM 回复格式异常:回传错误要求重试一次结构化输出
|
||||
messages.append(
|
||||
{
|
||||
"role": "user",
|
||||
"content": f"输出格式错误,请只输出一行 JSON(tool 或 final):{exc}",
|
||||
}
|
||||
)
|
||||
continue
|
||||
if "final" in decision:
|
||||
return {"reply": str(decision["final"]), "actions": actions}
|
||||
tool_name = str(decision.get("tool", ""))
|
||||
args = decision.get("args") or {}
|
||||
if not isinstance(args, dict):
|
||||
args = {}
|
||||
tool = self._by_name(tool_name) # 白名单之外的调用直接报错
|
||||
output = tool.invoke(args)
|
||||
actions.append({"tool": tool_name, "args": args, "output": output[:1000]})
|
||||
messages.append(
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": f'{{"tool": "{tool_name}", "args": {json.dumps(args, ensure_ascii=False)}}}',
|
||||
}
|
||||
)
|
||||
messages.append(
|
||||
{
|
||||
"role": "user",
|
||||
"content": f"工具 {tool_name} 返回:\n{output}\n请继续(如需再调用输出 tool,否则输出 final)。",
|
||||
}
|
||||
)
|
||||
|
||||
# 轮次耗尽:强制收尾
|
||||
messages.append(
|
||||
{
|
||||
"role": "user",
|
||||
"content": "已达工具调用轮次上限,请直接给出基于已有事实的最终结论(只输出 final JSON)。",
|
||||
}
|
||||
)
|
||||
try:
|
||||
decision = _parse_decision(self._llm.chat(messages))
|
||||
except ValueError:
|
||||
decision = {
|
||||
"final": "研究轮次耗尽且模型未给出结构化结论,请人工查看 actions 中的实验输出。"
|
||||
}
|
||||
return {"reply": str(decision.get("final", decision)), "actions": actions}
|
||||
Reference in New Issue
Block a user