From 8b2f8ac35cfef5d6870457bb6718840e53fe85eb Mon Sep 17 00:00:00 2001 From: Simon Date: Wed, 9 Sep 2026 00:39:42 +0800 Subject: [PATCH] =?UTF-8?q?feat(agent):=20M8.5=20Agent=20=E5=B7=A5?= =?UTF-8?q?=E5=85=B7=E8=A1=A5=E9=BD=90=EF=BC=88screen=5Fstocks/explain=5Fs?= =?UTF-8?q?election/generate=5Fsignals/create=5Fstrategy=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 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 通过 --- backend/app/agent/tools_impl.py | 182 ++++++++++++++++++++ backend/tests/test_agent_selection_tools.py | 94 ++++++++++ 2 files changed, 276 insertions(+) create mode 100644 backend/tests/test_agent_selection_tools.py diff --git a/backend/app/agent/tools_impl.py b/backend/app/agent/tools_impl.py index 06fe554..87420d9 100644 --- a/backend/app/agent/tools_impl.py +++ b/backend/app/agent/tools_impl.py @@ -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, + ), ] diff --git a/backend/tests/test_agent_selection_tools.py b/backend/tests/test_agent_selection_tools.py new file mode 100644 index 0000000..c32ffc7 --- /dev/null +++ b/backend/tests/test_agent_selection_tools.py @@ -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