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 通过
This commit is contained in:
@@ -12,10 +12,22 @@ 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,
|
||||
)
|
||||
|
||||
|
||||
@@ -159,6 +171,121 @@ def build_tools(factories: dict | None = None) -> list[Tool]:
|
||||
)
|
||||
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",
|
||||
@@ -236,6 +363,61 @@ def build_tools(factories: dict | None = None) -> list[Tool]:
|
||||
},
|
||||
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,
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,94 @@
|
||||
"""M8.5 Agent 新工具测试:screen_stocks / explain_selection / create_strategy。
|
||||
|
||||
使用 tmp SQLite(真实 SQLAlchemy repo)+ build_tools(factories) 直调工具。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date
|
||||
|
||||
import pytest
|
||||
from app.agent.tools_impl import build_tools
|
||||
from app.domain.entities.market import Stock
|
||||
from app.infrastructure.persistence.sqlalchemy.base import Base
|
||||
from app.infrastructure.persistence.sqlalchemy.repositories.market_impl import (
|
||||
SqlAlchemyDailyBarRepository,
|
||||
SqlAlchemyStockRepository,
|
||||
)
|
||||
from sqlalchemy import create_engine
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
|
||||
from conftest_quant import bars_dataframe_to_daily_bars, synthetic_daily
|
||||
|
||||
_SYMS = ["600000.SH", "600001.SH", "600002.SH", "600003.SH", "600004.SH"]
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def tools(tmp_path):
|
||||
engine = create_engine(f"sqlite:///{tmp_path / 'agent.db'}", future=True)
|
||||
Base.metadata.create_all(engine)
|
||||
Session = sessionmaker(bind=engine, expire_on_commit=False)
|
||||
|
||||
df = synthetic_daily({s: 0.006 - 0.0015 * i for i, s in enumerate(_SYMS)}, n=320)
|
||||
with Session() as session:
|
||||
SqlAlchemyStockRepository(session).upsert_many(
|
||||
[
|
||||
Stock(symbol=s, name=f"测试股份{i}", list_date=date(1999, 1, 1))
|
||||
for i, s in enumerate(_SYMS)
|
||||
]
|
||||
)
|
||||
SqlAlchemyDailyBarRepository(session).upsert_many(bars_dataframe_to_daily_bars(df))
|
||||
session.commit()
|
||||
|
||||
factories = {
|
||||
"session_factory": lambda: Session(),
|
||||
"stock_repo_factory": lambda s: SqlAlchemyStockRepository(s),
|
||||
"daily_repo_factory": lambda s: SqlAlchemyDailyBarRepository(s),
|
||||
"experiment_repo_factory": lambda s: None,
|
||||
"engine": None,
|
||||
}
|
||||
return build_tools(factories=factories)
|
||||
|
||||
|
||||
def _invoke(tools, name: str, args: dict) -> str:
|
||||
tool = next(t for t in tools if t.name == name)
|
||||
return tool.invoke(args)
|
||||
|
||||
|
||||
class TestAgentSelectionTools:
|
||||
def test_screen_stocks(self, tools) -> None:
|
||||
out = _invoke(
|
||||
tools,
|
||||
"screen_stocks",
|
||||
{
|
||||
"factors": "momentum_60",
|
||||
"top_n": 3,
|
||||
"as_of": "2024-12-31",
|
||||
"symbols": ",".join(_SYMS),
|
||||
},
|
||||
)
|
||||
assert "600000.SH" in out or "600001.SH" in out
|
||||
assert "score=" in out
|
||||
|
||||
def test_screen_stocks_empty_scope(self, tools) -> None:
|
||||
out = _invoke(tools, "screen_stocks", {"symbols": "600999.SH", "as_of": "2024-12-31"})
|
||||
assert "无候选" in out or "600999" in out
|
||||
|
||||
def test_create_strategy_and_explain_flow(self, tools) -> None:
|
||||
out = _invoke(
|
||||
tools,
|
||||
"create_strategy",
|
||||
{"name": "Agent 策略", "factors": "momentum_60", "top_n": 5},
|
||||
)
|
||||
assert "策略已保存" in out
|
||||
# 重名被拒(工具返回错误消息而非崩溃)
|
||||
out2 = _invoke(tools, "create_strategy", {"name": "Agent 策略", "factors": "momentum_20"})
|
||||
assert "已存在" in out2 or "失败" in out2
|
||||
|
||||
def test_explain_selection_missing(self, tools) -> None:
|
||||
out = _invoke(tools, "explain_selection", {"selection_id": "SEL-NOPE"})
|
||||
assert "不存在" in out
|
||||
|
||||
def test_tool_names_registered(self, tools) -> None:
|
||||
names = {t.name for t in tools}
|
||||
assert {"screen_stocks", "explain_selection", "generate_signals", "create_strategy"} <= names
|
||||
Reference in New Issue
Block a user