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

131 lines
5.5 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.
"""公共配置 + 回测组合 API 测试(TestClient + 内存 SQLite)。
覆盖:
- /api/config GET 返回默认值、PUT 持久化;
- /api/combos CRUD(name 唯一、原地更新保留 created_at、删除);
- /api/combos/{id}/run 在策略缺失时提前 400(不等到后台才失败)。
"""
from __future__ import annotations
import pytest
from app.api import deps
from app.infrastructure.persistence.sqlalchemy.base import Base
from app.main import app
from fastapi.testclient import TestClient
from sqlalchemy import create_engine
from sqlalchemy.orm import sessionmaker
@pytest.fixture()
def client(tmp_path) -> TestClient:
engine = create_engine(f"sqlite:///{tmp_path / 'combo_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 TestConfigApi:
def test_get_returns_defaults(self, client: TestClient) -> None:
cfg = client.get("/api/config").json()
assert cfg["id"] == "default"
assert cfg["price_adjustment"] == "hfq"
assert cfg["min_commission"] == 5.0
def test_put_persists(self, client: TestClient) -> None:
body = {
"commission_rate": 0.00025, "stamp_tax_rate": 0.0005,
"slippage_rate": 0.0008, "min_commission": 3.0,
"price_adjustment": "qfq", "benchmark": "000905.SH",
}
resp = client.put("/api/config", json=body)
assert resp.status_code == 200
got = client.get("/api/config").json()
assert got["commission_rate"] == pytest.approx(0.00025)
assert got["price_adjustment"] == "qfq"
assert got["benchmark"] == "000905.SH"
def test_put_rejects_unknown_field(self, client: TestClient) -> None:
"""拼错键名必须报错:否则「以为改了滑点、其实没生效」是静默降级。"""
resp = client.put("/api/config", json={"slippage": 0.001})
assert resp.status_code == 422
assert "slippage" in resp.text
# 值未被改动(仍是默认滑点)
assert client.get("/api/config").json()["slippage_rate"] == pytest.approx(0.001)
class TestCombosApi:
def _make_strategy(self, client: TestClient, name: str) -> str:
resp = client.post(
"/api/strategies",
json={"name": name, "factors": [{"name": "dividend_yield", "weight": 1}]},
)
assert resp.status_code == 200
return resp.json()["id"]
def test_crud(self, client: TestClient) -> None:
sid = self._make_strategy(client, "高股息")
body = {
"name": "组合A", "strategy_ids": [sid],
"initial_capital": 500000, "hold_count": 10,
"hold_min_days": 5, "hold_max_days": 30,
"rebalance_freq": "weekly", "period": ["2024-01-01", "2024-06-01"],
}
created = client.post("/api/combos", json=body)
assert created.status_code == 200
cid = created.json()["id"]
assert cid.startswith("CMB-")
assert created.json()["hold_max_days"] == 30
assert len(client.get("/api/combos").json()) == 1
detail = client.get(f"/api/combos/{cid}").json()
assert detail["strategy_ids"] == [sid]
assert detail["rebalance_freq"] == "weekly"
# 原地更新保留 id 与 created_at
created_at = detail["created_at"]
upd = client.put(f"/api/combos/{cid}", json={**body, "name": "组合A", "hold_count": 15})
assert upd.status_code == 200
assert upd.json()["id"] == cid
assert upd.json()["created_at"] == created_at
assert upd.json()["hold_count"] == 15
assert client.delete(f"/api/combos/{cid}").status_code == 200
assert client.get(f"/api/combos/{cid}").status_code == 404
def test_duplicate_name_400(self, client: TestClient) -> None:
sid = self._make_strategy(client, "S")
body = {"name": "重名", "strategy_ids": [sid], "hold_count": 5,
"rebalance_freq": "monthly", "period": ["2024-01-01", "2024-02-01"]}
assert client.post("/api/combos", json=body).status_code == 200
assert client.post("/api/combos", json=body).status_code == 400
def test_run_missing_strategy_400(self, client: TestClient) -> None:
"""引用不存在的策略 → 提交时即 400,而非等后台 Job 才失败。"""
body = {"name": "缺策略组合", "strategy_ids": ["STG-NOT-EXIST"], "hold_count": 5,
"rebalance_freq": "monthly", "period": ["2024-01-01", "2024-02-01"]}
resp = client.post("/api/combos/run", json=body)
assert resp.status_code == 400
assert "不存在" in resp.json()["detail"]
def test_unknown_param_rejected(self, client: TestClient) -> None:
"""回测参数写错键名(capital/hold_days)→ 422,而不是静默用默认值。"""
sid = self._make_strategy(client, "S2")
body = {
"name": "拼错参数", "strategy_ids": [sid], "hold_count": 5,
"rebalance_freq": "monthly", "period": ["2024-01-01", "2024-02-01"],
"capital": 500000, "hold_days": 30,
}
resp = client.post("/api/combos", json=body)
assert resp.status_code == 422
assert "capital" in resp.text and "hold_days" in resp.text
assert client.get("/api/combos").json() == []