Files
qlib/backend/app/agent/tools_impl.py
T
Simon 8b2f8ac35c feat(agent): M8.5 Agent 工具补齐(screen_stocks/explain_selection/generate_signals/create_strategy)
- tools_impl 新增 4 工具(Agent 共 10 个):
  screen_stocks(因子评分选股,symbols 白名单防全市场长任务)、
  explain_selection(读回选股结果并解释理由)、
  generate_signals(BUY/WATCH/SELL + 规则)、create_strategy(命名策略入库)
- 全部经白名单 Tool + Repository/Session,无 shell/写删权限扩张
- tests/test_agent_selection_tools.py 5 例(选股/策略保存+重名/解释 404/注册表);全量 pytest 通过
2026-09-09 00:39:42 +08:00

431 lines
18 KiB
Python
Raw 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.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.application.services.selection_service import SelectionService
from app.application.services.signal_service import SignalService
from app.domain.entities.research import (
BacktestResult,
FactorTestReport,
ResearchSpec,
UniverseSpec,
)
from app.domain.entities.selection import SelectionQuery
from app.domain.entities.signal import SignalRules
from app.domain.entities.strategy import StrategyDefinition
from app.infrastructure.persistence.sqlalchemy.repositories.selection_impl import (
SqlAlchemySelectionRepository,
)
from app.infrastructure.persistence.sqlalchemy.repositories.strategy_impl import (
SqlAlchemyStrategyRepository,
)
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)
def _scope_symbols(raw: str | None) -> list[str]:
"""白名单(可选):避免全市场长任务拖垮同步对话(全市场可用 Web 页异步)。"""
if not raw:
return []
return [x.strip().upper() for x in raw.split(",") if x.strip()][:60]
def screen_stocks(args: dict) -> str:
factors = [x.strip() for x in str(_pick(args, "factors", "momentum_60")).split(",") if x.strip()]
top_n = int(_pick(args, "top_n", 10) or 10)
as_of = _day(str(_pick(args, "as_of", date.today().isoformat())))
symbols = _scope_symbols(str(_pick(args, "symbols", "") or ""))
if not symbols:
return (
"为避免全市场长任务(>1 分钟),请传 symbols 白名单(≤60,逗号分隔)"
"或使用 Web 选股页执行全市场选股。"
)
query = SelectionQuery(
universe=UniverseSpec(
exclude_st=bool(_pick(args, "exclude_st", True)),
min_listing_days=0,
symbols=symbols,
),
factors=[{"name": f, "weight": 1.0} for f in factors],
top_n=top_n,
as_of=as_of,
)
with session_factory() as session:
service = SelectionService(
stock_repo_f(session), daily_repo_f(session)
)
result = service.select(query)
if not result.candidates:
return (
f"{as_of} 无候选(范围 {result.statistics.universe_size} 只,"
f"可评分 {result.statistics.evaluated})。如需白名单可传 symbols(≤60)。"
)
lines = [f"as_of={result.as_of_date} 选出 Top{len(result.candidates)}:"]
for c in result.candidates:
vals = ", ".join(f"{k}={v:.4f}" for k, v in c.factor_values.items())
lines.append(f" #{c.rank} {c.symbol} score={c.score:.4f}({vals})")
lines.append("入选理由见 explain_selection(selection_id)。")
return "\n".join(lines)
def explain_selection(args: dict) -> str:
sel_id = str(_pick(args, "selection_id", "")).upper()
with session_factory() as session:
repo = SqlAlchemySelectionRepository(session)
result = repo.get(sel_id)
if result is None:
return f"选股记录 {sel_id} 不存在(先通过 Web 选股页或 screen_stocks 生成)"
out = [f"选股 {sel_id} as_of={result.as_of_date}({result.method},选出 {len(result.candidates)} 只)"]
for c in result.candidates[:10]:
reasons = "; ".join(c.selection_reason[:3])
out.append(f" #{c.rank} {c.symbol} score={c.score:.4f} — {reasons}")
return "\n".join(out)
def generate_signals(args: dict) -> str:
factors = [x.strip() for x in str(_pick(args, "factors", "momentum_60")).split(",") if x.strip()]
as_of = _day(str(_pick(args, "as_of", date.today().isoformat())))
symbols = _scope_symbols(str(_pick(args, "symbols", "") or ""))
query = SelectionQuery(
universe=UniverseSpec(
exclude_st=bool(_pick(args, "exclude_st", True)),
min_listing_days=0,
symbols=symbols,
),
factors=[{"name": f, "weight": 1.0} for f in factors],
top_n=int(_pick(args, "top_n", 50) or 50),
as_of=as_of,
)
rules = SignalRules(
buy_rank_threshold=int(_pick(args, "buy_rank", 20) or 20),
sell_rank_threshold=int(_pick(args, "sell_rank", 50) or 50),
)
with session_factory() as session:
res = SignalService(stock_repo_f(session), daily_repo_f(session)).signal(query, rules)
out = [
f"信号 as_of={res.as_of_date}: BUY {res.statistics.buy} / WATCH {res.statistics.watch} / "
f"SELL {res.statistics.sell}(前 8 条)"
]
for e in res.events[:8]:
out.append(f" {e.signal_type} {e.symbol} score={e.score:.4f} — {e.trigger_reason[0] if e.trigger_reason else ''}")
return "\n".join(out)
def create_strategy(args: dict) -> str:
name = str(_pick(args, "name", ""))
if not name:
return "请提供 name"
factors = [
{"name": x.strip(), "weight": 1.0}
for x in str(_pick(args, "factors", "momentum_60")).split(",")
if x.strip()
]
if not factors:
return "请提供至少一个 factors(逗号分隔)"
description = str(_pick(args, "description", "") or "")
st = StrategyDefinition(
name=name,
description=description,
universe=UniverseSpec(
exclude_st=bool(_pick(args, "exclude_st", True)), min_listing_days=0
),
factors=factors,
selection={"top_n": int(_pick(args, "top_n", 10) or 10)},
rebalance=str(_pick(args, "rebalance", "monthly")),
)
from app.application.services.job_executor import new_id
with session_factory() as session:
saved = SqlAlchemyStrategyRepository(session).save(
st.model_copy(update={"id": new_id("STG")})
)
session.commit()
return f"策略已保存:{saved.id} {saved.name}(factors={[f.name for f in saved.factors]})"
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,
),
Tool(
"screen_stocks",
"按因子评分筛选股票(TopN;传 symbols 白名单避免全市场长任务)",
{
"type": "object",
"properties": {
"factors": {"type": "string", "description": "逗号分隔因子名"},
"top_n": {"type": "integer"},
"as_of": {"type": "string", "description": "YYYY-MM-DD"},
"symbols": {"type": "string", "description": "逗号分隔白名单(可选,≤60)"},
},
},
screen_stocks,
),
Tool(
"explain_selection",
"解释一次选股结果:为什么选这些股票(含因子值与理由)",
{
"type": "object",
"properties": {"selection_id": {"type": "string"}},
"required": ["selection_id"],
},
explain_selection,
),
Tool(
"generate_signals",
"基于选股评分+趋势生成 BUY/WATCH/SELL 信号",
{
"type": "object",
"properties": {
"factors": {"type": "string"},
"as_of": {"type": "string"},
"symbols": {"type": "string", "description": "白名单(可选)"},
"buy_rank": {"type": "integer"},
"sell_rank": {"type": "integer"},
},
},
generate_signals,
),
Tool(
"create_strategy",
"创建/保存一个命名策略(可随后展开为回测)",
{
"type": "object",
"properties": {
"name": {"type": "string"},
"description": {"type": "string"},
"factors": {"type": "string"},
"top_n": {"type": "integer"},
"rebalance": {"type": "string", "enum": ["monthly", "weekly"]},
},
"required": ["name", "factors"],
},
create_strategy,
),
]
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 "-"