Files
Simon d9be75a98f 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
2026-09-06 17:22:05 +08:00

99 lines
3.8 KiB
Python
Raw Permalink 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.
"""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}