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:
Simon
2026-09-06 17:22:05 +08:00
parent 0ea229d766
commit d9be75a98f
14 changed files with 722 additions and 23 deletions
+248
View File
@@ -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 "-"