Files
qlib/backend/tests/test_strategies.py
T
Simon 2e90f3eeac feat(backend): 字段库(condition_field)+ 因子参数化(模板/受控参数)+ 单位换算底座
字段库(本次新增的表与接口):
- `condition_field` 表 + `/api/condition-fields`:中文名/说明可编辑、可停用;
  `kind`/单位阶梯/`base_unit` 由代码注册表收敛(改类型 422,伪字段 422,
  越界单位 422),停用的字段不再进条件下拉,但既有策略仍按名字解析。
- 说明书里的数值条件按字段注册表补**基准单位**后缀(字段间比较不加,不猜单位)。

因子参数化(键即身份,冻结口径):
- 模板 + 参数注册表(`quant/factors.py`):`ParamSpec`(类型/范围/枚举/默认值/说明)+
  `FactorTemplate`(公式/依赖列/参数);规范键把**全部**参数写进名字,如
  `momentum(window=90,direction=lower_is_better)`,所以改参数 = 新建一个身份,
  旧因子/既有策略/已归档实验都不变义;`momentum(window=90)`(缺参数)明确拒绝 ——
  缺项要靠模板默认值补齐,而默认值是可改的代码细节,一旦改动会追溯性改义。
- 参数只在受控范围内取值(窗口 2~500、方向二选一),越界/未知模板/多给参数一律 422
  并列出允许范围,不静默截断、不悄悄取默认值;内置实例的启用开关由代码决定(422)。
- `/api/factors` 暴露 `template`/`params`/`param_specs`/`label`/`source`/`enabled`/
  `resolvable`;新增 `/api/factors/templates`、`POST /api/factors`、`PATCH /api/factors`;
  `get_factor = resolve_factor` 兼容全部旧调用点,参数化键也是一等条件字段。
- 迁移链:c5d6(存量策略陈旧说明重算)→ d6e7(condition_field)→ a7c1
  (factor_definition.enabled + name varchar(128))。

测试:新增 test_condition_fields.py / test_factor_params.py;全量 pytest 500 passed。
2026-10-01 16:33:32 +08:00

198 lines
8.2 KiB
Python

"""选股策略(SelectionStrategy)测试:repo CRUD + /api/strategies + 说明生成。
2026-09 重构后:策略库只存「选股条件组合」(股票池 + 因子 + 条件),
不再持有 selection/rebalance/costs/portfolio/区间 —— 那些移到回测组合与公共配置。
旧的 `to_research_spec` / `/expand` 已移除。
"""
from __future__ import annotations
import json
from datetime import datetime
import pytest
from app.api import deps
from app.domain.entities.strategy import SelectionStrategy
from app.infrastructure.persistence.sqlalchemy.base import Base
from app.infrastructure.persistence.sqlalchemy.models.strategy import StrategyModel
from app.infrastructure.persistence.sqlalchemy.repositories.strategy_impl import (
SqlAlchemyStrategyRepository,
)
from app.main import app
from fastapi.testclient import TestClient
from pydantic import ValidationError
from sqlalchemy import create_engine
from sqlalchemy.orm import sessionmaker
def _st() -> SelectionStrategy:
return SelectionStrategy(
name="质量成长动量",
description="ROE+动量(演示)",
factors=[{"name": "momentum_60", "weight": 1.0}],
conditions=[{"field": "dv_ratio", "op": "lte", "value": 30}],
)
@pytest.fixture()
def session(tmp_path):
engine = create_engine(f"sqlite:///{tmp_path / 'st.db'}", future=True)
Base.metadata.create_all(engine)
Session = sessionmaker(bind=engine, expire_on_commit=False)
with Session() as s:
yield s
class TestStrategyRepository:
def test_save_get_list_delete(self, session) -> None:
repo = SqlAlchemyStrategyRepository(session)
repo.save(_st().model_copy(update={"id": "STG-T1"}))
session.commit()
got = repo.get("STG-T1")
assert got is not None and got.name == "质量成长动量"
# 选股策略只保留选股相关字段
assert got.factors[0].name == "momentum_60"
assert len(got.conditions) == 1 and got.conditions[0].field == "dv_ratio"
assert len(repo.list()) == 1
assert repo.get_by_name("质量成长动量") is not None
assert repo.delete("STG-T1") is True
session.commit()
assert repo.get("STG-T1") is None
def test_duplicate_name(self, session) -> None:
repo = SqlAlchemyStrategyRepository(session)
repo.save(_st().model_copy(update={"id": "STG-A"}))
session.commit()
with pytest.raises(ValueError):
repo.save(_st().model_copy(update={"id": "STG-B"}))
def test_no_backtest_params_in_entity(self) -> None:
"""选股策略实体不应再有回测执行参数字段(重构的核心约束)。"""
st = _st()
dumped = st.model_dump()
for forbidden in (
"selection", "rebalance", "costs", "portfolio",
"initial_capital", "period", "price_adjustment",
"selection_interval_months", "rebalance_interval_months",
):
assert forbidden not in dumped, f"选股策略不应含回测参数字段 {forbidden}"
def test_rejects_legacy_backtest_params(self) -> None:
"""混入旧版回测参数必须**报错**,不能静默丢弃(否则调用方以为设上了)。"""
with pytest.raises(ValidationError) as exc:
SelectionStrategy(
name="带旧参数的策略",
factors=[{"name": "momentum_60", "weight": 1}],
costs={"commission_rate": 0.001},
rebalance="monthly",
initial_capital=500_000,
)
# 三个未知键都应被点名(便于调用方知道该搬去哪里)
for key in ("costs", "rebalance", "initial_capital"):
assert key in str(exc.value)
def test_spec_type_only_selection(self) -> None:
"""spec_type 取值域收敛为 selection(本实体只表示选股策略)。"""
with pytest.raises(ValidationError):
SelectionStrategy(
name="旧类型", spec_type="backtest",
factors=[{"name": "momentum_60", "weight": 1}],
)
assert _st().spec_type == "selection"
def test_legacy_rows_with_extra_keys_still_readable(self, session) -> None:
"""历史行 config_json 残留旧键时仍能读出(仓储读出前剔除),forbid 不影响兼容。"""
session.add(
StrategyModel(
id="STG-LEGACY", name="历史策略", description="历史说明",
spec_type="backtest",
config_json=json.dumps({
"universe": {"exclude_st": True},
"factors": [{"name": "dividend_yield", "weight": 1}],
"conditions": [],
"selection": {"top_n": 20},
"rebalance": "monthly",
"costs": {"commission_rate": 0.0003},
}, ensure_ascii=False),
version="1", created_at=datetime(2026, 1, 1),
)
)
session.commit()
got = SqlAlchemyStrategyRepository(session).get("STG-LEGACY")
assert got is not None
assert got.factors[0].name == "dividend_yield"
assert got.spec_type == "selection" # 旧列值不参与实体(回落到默认)
@pytest.fixture()
def client(tmp_path):
engine = create_engine(f"sqlite:///{tmp_path / 'api.db'}", future=True)
Base.metadata.create_all(engine)
Session = sessionmaker(bind=engine, expire_on_commit=False)
def _session_override():
with Session() as s:
yield s
app.dependency_overrides[deps.get_session] = _session_override
with TestClient(app) as c:
yield c
app.dependency_overrides.clear()
class TestStrategiesApi:
def test_crud_and_describe(self, client) -> None:
body = {
"name": "演示策略",
"description": "动量",
"factors": [{"name": "momentum_60", "weight": 1}],
"conditions": [{"field": "dv_ratio", "op": "lte", "value": 30}],
}
created = client.post("/api/strategies", json=body)
assert created.status_code == 200
sid = created.json()["id"]
assert sid.startswith("STG-")
# 回读不含回测参数字段
detail = client.get(f"/api/strategies/{sid}").json()
assert detail["name"] == "演示策略"
assert "selection" not in detail and "costs" not in detail
assert len(client.get("/api/strategies").json()) == 1
# 说明生成:选股策略走专用路径,不假装知道回测参数
doc = client.get(f"/api/strategies/{sid}/describe").json()
assert "选股策略" in doc["summary"]
assert any("回测组合" in w for w in doc["warnings"])
assert client.delete(f"/api/strategies/{sid}").status_code == 200
assert client.get(f"/api/strategies/{sid}").status_code == 404
def test_duplicate_name_400(self, client) -> None:
body = {"name": "A", "factors": [{"name": "momentum_60", "weight": 1}]}
assert client.post("/api/strategies", json=body).status_code == 200
assert client.post("/api/strategies", json=body).status_code == 400
def test_expand_endpoint_removed(self, client) -> None:
"""/expand 已随重构移除(回测改由「回测组合」驱动,不再从单策略展开 ResearchSpec)。"""
body = {"name": "B", "factors": [{"name": "momentum_60", "weight": 1}]}
sid = client.post("/api/strategies", json=body).json()["id"]
resp = client.post(
f"/api/strategies/{sid}/expand",
json={"period": ["2024-01-01", "2024-06-01"]},
)
assert resp.status_code in (404, 405)
def test_create_with_legacy_params_422(self, client) -> None:
"""旧调用方把回测参数塞进策略 → 422 并点名未知字段(而非 200 静默丢弃)。"""
body = {
"name": "旧调用方",
"factors": [{"name": "momentum_60", "weight": 1}],
"rebalance": "monthly",
"costs": {"commission_rate": 0.001},
}
resp = client.post("/api/strategies", json=body)
assert resp.status_code == 422
assert "costs" in resp.text and "rebalance" in resp.text
# 未落库:策略列表仍为空
assert client.get("/api/strategies").json() == []