- 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 通过
95 lines
3.4 KiB
Python
95 lines
3.4 KiB
Python
"""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
|