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,77 @@
|
||||
"""LLM 客户端抽象(Phase 5)。
|
||||
|
||||
实现约定:真实 Key 来自 .env(LLM_API_KEY / LLM_BASE_URL / LLM_MODEL,见 config.yaml 引用)。
|
||||
测试注入 FakeLLMClient 走完整编排链路,不触网。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Protocol
|
||||
|
||||
import httpx
|
||||
|
||||
from app.agent.tools import tools_schema
|
||||
from app.agent.tools_impl import build_tools
|
||||
from app.core.config import get_settings
|
||||
|
||||
SYSTEM_TEMPLATE = """你是个人 A 股量化研究助手(Research Assistant),不是系统管理员。
|
||||
|
||||
你可以调用以下工具(每次只能输出一个动作):
|
||||
{tools}
|
||||
|
||||
输出规则:只输出一行 JSON,两种形态之一:
|
||||
1. 需要调用工具:{{"tool": "<工具名>", "args": {{...}}}}
|
||||
2. 给出结论:{{"final": "结论文本"}}
|
||||
|
||||
研究纪律(必须遵守):
|
||||
- 先提出假设 → 用工具做因子测试或回测 → 基于实验事实分析,再给结论
|
||||
- 不得仅凭单次回测高收益就宣布策略有效;要主动说明样本外、过拟合、
|
||||
look-ahead bias、交易成本、参数敏感性等风险(未做验证的项要明说「未验证」)
|
||||
- 全程只读:不得要求删除/修改数据或执行任意命令(你也没有这类工具)
|
||||
- 回答使用简体中文
|
||||
"""
|
||||
|
||||
|
||||
class LLMClient(Protocol):
|
||||
def chat(self, messages: list[dict]) -> str: ...
|
||||
|
||||
|
||||
class OpenAICompatibleClient:
|
||||
"""OpenAI 兼容 Chat Completions(qwen/dashscope、deepseek、openai 等均适用)。"""
|
||||
|
||||
def __init__(self, *, base_url: str, api_key: str, model: str, timeout: float = 60.0) -> None:
|
||||
self._base_url = base_url.rstrip("/")
|
||||
self._api_key = api_key
|
||||
self._model = model
|
||||
self._timeout = timeout
|
||||
|
||||
def chat(self, messages: list[dict]) -> str:
|
||||
with httpx.Client(timeout=self._timeout) as client:
|
||||
resp = client.post(
|
||||
f"{self._base_url}/chat/completions",
|
||||
headers={"Authorization": f"Bearer {self._api_key}"},
|
||||
json={
|
||||
"model": self._model,
|
||||
"messages": messages,
|
||||
"temperature": 0.2,
|
||||
},
|
||||
)
|
||||
resp.raise_for_status()
|
||||
data = resp.json()
|
||||
return data["choices"][0]["message"]["content"]
|
||||
|
||||
|
||||
def build_llm_from_settings():
|
||||
"""从配置构造 LLM;未配置 Key 时返回 None(调用方给出引导提示)。"""
|
||||
settings = get_settings()
|
||||
if not settings.llm_api_key:
|
||||
return None
|
||||
return OpenAICompatibleClient(
|
||||
base_url=settings.llm_base_url or "https://dashscope.aliyuncs.com/compatible-mode/v1",
|
||||
api_key=settings.llm_api_key,
|
||||
model=settings.llm_model,
|
||||
)
|
||||
|
||||
|
||||
def system_prompt() -> str:
|
||||
return SYSTEM_TEMPLATE.format(tools=tools_schema(build_tools()))
|
||||
@@ -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}
|
||||
@@ -0,0 +1,45 @@
|
||||
"""AI Research Agent(Phase 5)。
|
||||
|
||||
Agent 是 Research Assistant(AGENT.md §28/§29):
|
||||
- 只能调用本目录 tools 提供的**白名单受控工具**(只读研究能力)
|
||||
- 禁止:执行任意 shell / 修改删除数据 / 修改配置与凭证
|
||||
- 研究行为:提出假设 → 建立实验 → 运行测试 → 分析 → 下一步,禁止以单次高收益宣告策略有效
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
TOOL_CALL_TAG = "__tool__"
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class Tool:
|
||||
"""受控工具元数据(LLM 可见的 JSON Schema + 调用实现)。"""
|
||||
|
||||
name: str
|
||||
description: str
|
||||
parameters: dict[str, Any]
|
||||
handler: Callable[[dict[str, Any]], str] = field(repr=False)
|
||||
|
||||
def schema(self) -> dict[str, Any]:
|
||||
return {
|
||||
"name": self.name,
|
||||
"description": self.description,
|
||||
"parameters": self.parameters,
|
||||
}
|
||||
|
||||
def invoke(self, args: dict[str, Any]) -> str:
|
||||
"""执行工具;任何异常都转为可读错误串(不向上抛,避免中断整轮对话)。"""
|
||||
try:
|
||||
result = self.handler(args)
|
||||
return result if isinstance(result, str) else json.dumps(result, ensure_ascii=False)
|
||||
except Exception as exc: # noqa: BLE001 —— 工具异常反馈给 LLM 而非崩溃
|
||||
return json.dumps({"error": f"{type(exc).__name__}: {exc}"}, ensure_ascii=False)
|
||||
|
||||
|
||||
def tools_schema(tools: list[Tool]) -> str:
|
||||
return json.dumps([t.schema() for t in tools], ensure_ascii=False, indent=1)
|
||||
@@ -1 +0,0 @@
|
||||
"""Agent 受控工具集:每个 Tool 是后端服务的只读 / 沙箱化入口。"""
|
||||
@@ -0,0 +1,248 @@
|
||||
"""受控工具集实现(AGENT.md §28):Agent 只能调用这里的白名单工具。
|
||||
|
||||
全部工具经 Job/Experiment 链路或只读查询执行:
|
||||
- 不提供 shell / 任意代码执行 / 修改配置与凭证 / 删除数据
|
||||
- 任何研究都会产出 Experiment 归档(可复现)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from datetime import date
|
||||
|
||||
from app.agent.tools import Tool
|
||||
from app.application.services.job_executor import default_factories, submit_and_run
|
||||
from app.domain.entities.research import (
|
||||
BacktestResult,
|
||||
FactorTestReport,
|
||||
ResearchSpec,
|
||||
)
|
||||
|
||||
|
||||
def _day(text: str) -> date:
|
||||
return date.fromisoformat(text)
|
||||
|
||||
|
||||
def _pick(mapping: dict, key: str, default=None):
|
||||
val = mapping.get(key, default)
|
||||
if isinstance(val, str):
|
||||
val = val.strip()
|
||||
if val == "":
|
||||
return default
|
||||
return val
|
||||
|
||||
|
||||
def build_tools(factories: dict | None = None) -> list[Tool]:
|
||||
facts = factories or default_factories()
|
||||
session_factory = facts["session_factory"]
|
||||
stock_repo_f = facts["stock_repo_factory"]
|
||||
daily_repo_f = facts["daily_repo_factory"]
|
||||
exp_repo_f = facts["experiment_repo_factory"]
|
||||
|
||||
def search_stocks(args: dict) -> str:
|
||||
q = str(_pick(args, "q", "") or "").upper()
|
||||
with session_factory() as session:
|
||||
stocks = stock_repo_f(session).list()
|
||||
rows = [
|
||||
s for s in stocks if (not q) or q in s.symbol.upper() or q in (s.name or "").upper()
|
||||
][:15]
|
||||
if not rows:
|
||||
return "未找到匹配股票"
|
||||
return "\n".join(
|
||||
f"{s.symbol} {s.name} 行业={s.industry or '-'} 上市={s.list_date}" for s in rows
|
||||
)
|
||||
|
||||
def get_market_data(args: dict) -> str:
|
||||
symbol = str(_pick(args, "symbol", "")).upper()
|
||||
start = _day(str(_pick(args, "start", "2024-01-01")))
|
||||
end = _day(str(_pick(args, "end", date.today().isoformat())))
|
||||
with session_factory() as session:
|
||||
bars = daily_repo_f(session).get_range(symbol, start, end)
|
||||
if not bars:
|
||||
return f"{symbol} 在 {start}~{end} 无日线数据(可能未同步)"
|
||||
head, tail = bars[0], bars[-1]
|
||||
last = "\n".join(f"{b.trade_date} close={b.close}" for b in bars[-8:])
|
||||
change = float(tail.close) / float(head.close) - 1 if head.close and tail.close else None
|
||||
return (
|
||||
f"{symbol} {start}~{end} 共 {len(bars)} 根日线;"
|
||||
f"区间 {head.trade_date}→{tail.trade_date} 收盘 {head.close}→{tail.close}"
|
||||
f"(涨跌 {change * 100:.2f}% 若数据完整);最近 8 根:\n{last}"
|
||||
)
|
||||
|
||||
def _run_spec(spec: ResearchSpec, desc: str) -> str:
|
||||
job = submit_and_run(spec, factories=facts)
|
||||
if job.status != "success":
|
||||
return f"{desc} 执行失败:{job.error}"
|
||||
if spec.type == "backtest":
|
||||
result = BacktestResult.model_validate_json(job.result_json or "{}")
|
||||
s = result.summary
|
||||
return (
|
||||
f"回测完成(Experiment {job.experiment_id},代码版本 {_code_version(job, exp_repo_f)})。"
|
||||
f"总收益 {s.total_return_pct:.2f}%,年化 {s.annual_return_pct:.2f}%,"
|
||||
f"Sharpe {s.sharpe:.2f},最大回撤 {s.max_drawdown_pct:.2f}%,"
|
||||
f"交易 {s.total_trades} 笔,平均换手 {s.avg_turnover_pct:.1f}%。"
|
||||
f"未建模约束 {len(result.unimplemented)} 项(成本/涨跌停近似见实验详情)。"
|
||||
)
|
||||
report = FactorTestReport.model_validate_json(job.result_json or "{}")
|
||||
qs = ", ".join(f"Q{q.quantile + 1}: {q.return_pct:.2f}%" for q in report.quantile_returns)
|
||||
return (
|
||||
f"因子测试完成(Experiment {job.experiment_id})。IC {report.ic_mean:.4f},"
|
||||
f"RankIC {report.rank_ic_mean:.4f},ICIR {report.icir:.2f},正收益占比 "
|
||||
f"{report.positive_ratio_pct:.1f}%,样本 {report.sample_days} 日;分层未来收益 {qs}。"
|
||||
f"注意:单因子测试不代表策略有效,需结合稳健性分析。"
|
||||
)
|
||||
|
||||
def test_factor(args: dict) -> str:
|
||||
name = str(_pick(args, "name", ""))
|
||||
start = _day(str(_pick(args, "start", "2024-01-01")))
|
||||
end = _day(str(_pick(args, "end", "2024-12-31")))
|
||||
spec = ResearchSpec(
|
||||
type="factor_test",
|
||||
universe={"exclude_st": True, "min_listing_days": 0},
|
||||
factors=[{"name": name, "weight": 1.0}],
|
||||
selection={"top_n": 10},
|
||||
rebalance="monthly",
|
||||
period=(start, end),
|
||||
)
|
||||
return _run_spec(spec, f"因子 {name} 测试")
|
||||
|
||||
def run_backtest(args: dict) -> str:
|
||||
factor_names = [f.strip() for f in str(_pick(args, "factors", "momentum_60")).split(",")]
|
||||
top_n = int(_pick(args, "top_n", 5) or 5)
|
||||
rebalance = str(_pick(args, "rebalance", "monthly"))
|
||||
exclude_st = bool(_pick(args, "exclude_st", True))
|
||||
start = _day(str(_pick(args, "start", "2024-01-01")))
|
||||
end = _day(str(_pick(args, "end", "2024-12-31")))
|
||||
spec = ResearchSpec(
|
||||
type="backtest",
|
||||
universe={"exclude_st": exclude_st, "min_listing_days": 0},
|
||||
factors=[{"name": n, "weight": 1.0} for n in factor_names],
|
||||
selection={"top_n": top_n},
|
||||
rebalance=rebalance,
|
||||
period=(start, end),
|
||||
)
|
||||
return _run_spec(spec, "回测")
|
||||
|
||||
def get_experiment(args: dict) -> str:
|
||||
exp_id = str(_pick(args, "experiment_id", "")).upper()
|
||||
with session_factory() as session:
|
||||
exp = exp_repo_f(session).get(exp_id)
|
||||
if exp is None:
|
||||
return f"Experiment {exp_id} 不存在(可用列表:GET /api/experiments)"
|
||||
spec = json.loads(exp.spec_json)
|
||||
return (
|
||||
f"Experiment {exp.id} [{exp.kind}] 因子={[f['name'] for f in spec.get('factors', [])]} "
|
||||
f"区间={spec.get('period')} 调仓={spec.get('rebalance')};摘要:{exp.summary_text or '-'} "
|
||||
f"代码版本={exp.code_version or '-'} 创建={exp.created_at}"
|
||||
)
|
||||
|
||||
def compare_experiments(args: dict) -> str:
|
||||
ids = [
|
||||
x.strip().upper()
|
||||
for x in str(_pick(args, "experiment_ids", "")).split(",")
|
||||
if x.strip()
|
||||
]
|
||||
if not ids:
|
||||
return "请提供 experiment_ids(逗号分隔)"
|
||||
with session_factory() as session:
|
||||
repo = exp_repo_f(session)
|
||||
rows = [(i, repo.get(i)) for i in ids]
|
||||
out = []
|
||||
for exp_id, exp in rows:
|
||||
if exp is None:
|
||||
out.append(f"{exp_id}: 不存在")
|
||||
else:
|
||||
spec = json.loads(exp.spec_json)
|
||||
out.append(
|
||||
f"{exp.id}: 因子={[f['name'] for f in spec.get('factors', [])]} "
|
||||
f"区间={spec.get('period')} → {exp.summary_text or '-'}"
|
||||
)
|
||||
return "\n".join(out)
|
||||
|
||||
return [
|
||||
Tool(
|
||||
"search_stocks",
|
||||
"按代码或名称搜索股票,返回基础信息(只读)",
|
||||
{
|
||||
"type": "object",
|
||||
"properties": {"q": {"type": "string", "description": "代码或名称关键字"}},
|
||||
},
|
||||
search_stocks,
|
||||
),
|
||||
Tool(
|
||||
"get_market_data",
|
||||
"读取一只股票一段区间的日线行情摘要(只读,不复权)",
|
||||
{
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"symbol": {"type": "string", "description": "如 600519.SH"},
|
||||
"start": {"type": "string", "description": "YYYY-MM-DD"},
|
||||
"end": {"type": "string", "description": "YYYY-MM-DD"},
|
||||
},
|
||||
"required": ["symbol"],
|
||||
},
|
||||
get_market_data,
|
||||
),
|
||||
Tool(
|
||||
"test_factor",
|
||||
"对单个因子做 IC/RankIC/分层测试并归档 Experiment",
|
||||
{
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"name": {
|
||||
"type": "string",
|
||||
"description": "因子名(momentum_60 / volatility_20 等)",
|
||||
},
|
||||
"start": {"type": "string"},
|
||||
"end": {"type": "string"},
|
||||
},
|
||||
"required": ["name"],
|
||||
},
|
||||
test_factor,
|
||||
),
|
||||
Tool(
|
||||
"run_backtest",
|
||||
"运行 TopK 低频回测并归档 Experiment(成本/涨跌停近似建模)",
|
||||
{
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"factors": {"type": "string", "description": "逗号分隔的因子名"},
|
||||
"top_n": {"type": "integer"},
|
||||
"rebalance": {"type": "string", "enum": ["monthly", "weekly"]},
|
||||
"exclude_st": {"type": "boolean"},
|
||||
"start": {"type": "string"},
|
||||
"end": {"type": "string"},
|
||||
},
|
||||
},
|
||||
run_backtest,
|
||||
),
|
||||
Tool(
|
||||
"get_experiment",
|
||||
"读取已归档实验的摘要",
|
||||
{
|
||||
"type": "object",
|
||||
"properties": {"experiment_id": {"type": "string"}},
|
||||
"required": ["experiment_id"],
|
||||
},
|
||||
get_experiment,
|
||||
),
|
||||
Tool(
|
||||
"compare_experiments",
|
||||
"对比多个实验(因子/区间/收益摘要)",
|
||||
{
|
||||
"type": "object",
|
||||
"properties": {"experiment_ids": {"type": "string"}},
|
||||
"required": ["experiment_ids"],
|
||||
},
|
||||
compare_experiments,
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
def _code_version(job, exp_repo_f) -> str:
|
||||
try:
|
||||
with default_factories()["session_factory"]() as session:
|
||||
exp = exp_repo_f(session).get(job.experiment_id or "")
|
||||
return exp.code_version or "-" if exp else "-"
|
||||
except Exception: # noqa: BLE001
|
||||
return "-"
|
||||
Reference in New Issue
Block a user