Files
qlib/backend/tests/test_agent_v3_tools.py
T
Simon ed54096331 feat(agent): D4 Agent 补齐至 v3 §25(10 → 14 工具)
- inspect_factor:因子目录元数据(公式/方向/lookback/输入列)
- create_composite_factor:解析 name:weight 组件并保存(方向由注册表填充,未注册 400 语义)
- get_backtest_result:回测 Experiment 详细结果(收益/回撤/交易/意图与成交统计)
- create_experiment:成功 Job 兜底归档为 Experiment(幂等提示)
- 白名单工具总数 14(= v3 §25 清单 + get_market_data);tests/test_agent_v3_tools.py
  5 例(数量/各工具行为/幂等);全量 pytest 通过
2026-09-09 07:40:45 +08:00

133 lines
5.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.
"""D4 Agent 补齐至 v3 §25(14 工具)测试:inspect_factor / create_composite_factor /
get_backtest_result / create_experiment。tmp SQLite + 真实 repo factories。"""
from __future__ import annotations
from datetime import date, datetime
import pytest
from app.agent.tools_impl import build_tools
from app.domain.entities.factor import FactorDefinition
from app.domain.entities.market import Stock
from app.domain.entities.research import (
ExperimentRecord,
JobRecord,
JobStatus,
ResearchSpec,
)
from app.infrastructure.persistence.sqlalchemy.base import Base
from app.infrastructure.persistence.sqlalchemy.repositories.factor_impl import (
SqlAlchemyFactorRepository,
)
from app.infrastructure.persistence.sqlalchemy.repositories.jobs_impl import (
SqlAlchemyExperimentRepository,
SqlAlchemyJobRepository,
)
from app.infrastructure.persistence.sqlalchemy.repositories.market_impl import (
SqlAlchemyDailyBarRepository,
SqlAlchemyStockRepository,
)
from app.quant.factors import list_factors
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"]
@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))
SqlAlchemyFactorRepository(session).upsert_many(
[FactorDefinition.from_registry_def(d) for d in list_factors()]
)
# 造一条成功 Job(无 experiment 关联,供 create_experiment 补档)
spec = ResearchSpec(
type="backtest",
universe={"exclude_st": False, "min_listing_days": 0, "symbols": _SYMS},
factors=[{"name": "momentum_60", "weight": 1.0}],
selection={"top_n": 1},
rebalance="monthly",
period=(date(2024, 5, 1), date(2024, 8, 31)),
)
from app.quant.engine import LocalEngine
result = LocalEngine().run_backtest(df, spec)
SqlAlchemyJobRepository(session).create(
JobRecord(id="JOB-D4", kind="backtest",
spec_json=spec.model_dump_json(), status=JobStatus.SUCCESS,
result_json=result.model_dump_json(),
experiment_id=None,
created_at=datetime.now(), finished_at=datetime.now())
)
SqlAlchemyExperimentRepository(session).save(
ExperimentRecord(id="EXP-D4", kind="backtest", spec_json=spec.model_dump_json(),
result_json=result.model_dump_json(),
summary_text="d4", job_id="JOB-D4",
created_at=datetime.now())
)
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: SqlAlchemyExperimentRepository(s),
"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 TestV3Tools:
def test_tool_count_14(self, tools) -> None:
names = {t.name for t in tools}
assert len(names) == 14
assert {
"inspect_factor", "create_composite_factor", "get_backtest_result", "create_experiment"
} <= names
def test_inspect_factor(self, tools) -> None:
out = _invoke(tools, "inspect_factor", {"name": "momentum_60"})
assert "momentum_60" in out and "公式" in out and "lookback" in out
miss = _invoke(tools, "inspect_factor", {"name": "no_such"})
assert "不在目录" in miss
def test_create_composite_factor(self, tools) -> None:
out = _invoke(
tools, "create_composite_factor",
{"name": "动量低波", "factors": "momentum_60:0.7,volatility_60:0.3"},
)
assert "组合已保存" in out and "CF-" in out
bad = _invoke(tools, "create_composite_factor",
{"name": "x", "factors": "no_such:1"})
assert "无法创建" in bad
def test_get_backtest_result(self, tools) -> None:
out = _invoke(tools, "get_backtest_result", {"experiment_id": "EXP-D4"})
assert "总收益" in out and "成交" in out
miss = _invoke(tools, "get_backtest_result", {"experiment_id": "EXP-NONE"})
assert "不存在" in miss
def test_create_experiment_backfills_job(self, tools) -> None:
out = _invoke(tools, "create_experiment", {"job_id": "JOB-D4"})
assert "已归档" in out and "JOB-D4" in out
again = _invoke(tools, "create_experiment", {"job_id": "JOB-D4"})
assert "已归档为" in again # 幂等提示
bad = _invoke(tools, "create_experiment", {"job_id": "JOB-NONE"})
assert "不存在" in bad