Files
qlib/backend/tests/test_agent_selection_tools.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

95 lines
3.4 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.
"""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