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

330 lines
15 KiB
Python
Raw Permalink 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.
"""因子参数化测试(2026-10):模板 + 参数 + 参数化键 + 目录/API 行为。
覆盖四件事(每件都是「错了会静默算错」的地方):
1. **参数真生效**:窗口改了因子值真的变、方向改了排序真的反过来;
2. **键即身份**:规范键可解析、非规范/越界/缺参一律当场拒绝(不靠默认值兜底);
3. **目录分工**:内置实例口径按代码收敛、参数化实例的开关是人配的(不被同步冲掉)、
算不出来的历史行保留但标 resolvable=False;
4. **过滤条件**:参数化因子能当条件字段用,且单位后缀不会瞎猜。
"""
from __future__ import annotations
import pytest
from app.api import deps
from app.application.services.factor_catalog import (
create_parameterized_factor,
sync_registry_factors,
)
from app.domain.entities.factor import FactorDefinition
from app.domain.entities.research import FactorSpec
from app.infrastructure.persistence.sqlalchemy.base import Base
from app.infrastructure.persistence.sqlalchemy.repositories.factor_impl import (
SqlAlchemyFactorRepository,
)
from app.main import app
from app.quant.composite import build_score_panel
from app.quant.condition_fields import get_field, reason_unsupported
from app.quant.factors import (
FactorError,
canonical_key,
compute_factor,
get_factor,
get_template,
is_resolvable,
list_factors,
list_templates,
parse_factor_key,
)
from app.quant.selection import condition_needed_columns
from fastapi.testclient import TestClient
from sqlalchemy import create_engine
from sqlalchemy.orm import sessionmaker
from conftest_quant import synthetic_daily
@pytest.fixture()
def session(tmp_path):
engine = create_engine(f"sqlite:///{tmp_path / 'param.db'}", future=True)
Base.metadata.create_all(engine)
Session = sessionmaker(bind=engine, expire_on_commit=False)
with Session() as s:
yield s
@pytest.fixture()
def client(tmp_path):
engine = create_engine(f"sqlite:///{tmp_path / 'param_api.db'}", future=True)
Base.metadata.create_all(engine)
Session = sessionmaker(bind=engine, expire_on_commit=False)
def _override():
with Session() as s:
yield s
app.dependency_overrides[deps.get_session] = _override
with TestClient(app) as c:
yield c
app.dependency_overrides.clear()
class TestParamsTakeEffect:
"""参数不是装饰:改了必须真的影响引擎算出来的东西。"""
def test_window_changes_values_and_lookback(self) -> None:
daily = synthetic_daily({"AAA": 0.002, "BBB": -0.001}, n=200)
d20, p20 = compute_factor("momentum_20", daily)
d90, p90 = compute_factor(
"momentum(window=90,direction=higher_is_better)", daily
)
assert d20.lookback == 20 and d90.lookback == 90
# 90 日动量覆盖更长区间:上涨股的累计涨幅必须更大
assert float(p90.iloc[-1]["AAA"]) > float(p20.iloc[-1]["AAA"])
# 且两者不是同一序列(窗口参数真的进了计算,而不是只写进名字)
assert not p20.iloc[-1].equals(p90.iloc[-1])
def test_direction_flips_the_score_ranking(self) -> None:
daily = synthetic_daily({"AAA": 0.002, "BBB": 0.0, "CCC": -0.002}, n=160)
hi = build_score_panel(
daily, [FactorSpec(name="momentum(window=90,direction=higher_is_better)")]
).iloc[-1]
lo = build_score_panel(
daily, [FactorSpec(name="momentum(window=90,direction=lower_is_better)")]
).iloc[-1]
assert float(hi["AAA"]) > float(hi["CCC"])
assert float(lo["CCC"]) > float(lo["AAA"]), "方向改成越低越好后,排序必须反过来"
def test_formula_and_description_render_the_params(self) -> None:
defn, _fn = get_factor("momentum(window=90,direction=higher_is_better)")
assert "90" in defn.description and "90" in defn.formula
assert defn.label == "动量(窗口 90,越高越好)"
assert defn.template == "momentum"
assert defn.params == {"window": 90, "direction": "higher_is_better"}
assert defn.source == "custom"
def test_builtin_instances_keep_their_legacy_metadata(self) -> None:
"""老名字必须还按原来的口径算:归档/策略/文档里到处是它们。"""
assert get_factor("momentum_60")[0].lookback == 60
assert get_factor("momentum_60")[0].direction == "higher_is_better"
assert get_factor("volatility_20")[0].direction == "lower_is_better"
assert get_factor("volume_ratio_5_60")[0].lookback == 60
assert get_factor("dividend_yield")[0].lookback == 0
assert {d.name for d in list_factors()} >= {"momentum_20", "dividend_yield_ttm"}
class TestCanonicalKeys:
def test_canonical_key_lists_every_param(self) -> None:
key = canonical_key("momentum", {"window": 90})
assert key == "momentum(window=90,direction=higher_is_better)"
# 写全参数的意义:键自解释,不依赖模板默认值(默认值改了也不会让老键变义)
assert canonical_key("volatility", {"window": 30}) == (
"volatility(window=30,direction=lower_is_better)"
)
assert canonical_key("volume_ratio", {"fast": 10, "slow": 120}) == (
"volume_ratio(fast=10,slow=120,direction=higher_is_better)"
)
def test_roundtrip(self) -> None:
key = "volume_ratio(fast=10,slow=120,direction=lower_is_better)"
template, params = parse_factor_key(key)
assert template.name == "volume_ratio"
assert canonical_key(template, params) == key
assert is_resolvable(key)
@pytest.mark.parametrize(
"bad,why",
[
("momentum(window=90)", "缺 direction(键不许依赖默认值)"),
("momentum(window=1,direction=higher_is_better)", "窗口低于下限"),
("momentum(window=999,direction=higher_is_better)", "窗口高于上限"),
("momentum(window=abc,direction=higher_is_better)", "窗口不是整数"),
("momentum(window=90,direction=upper)", "方向不是枚举值"),
("momentum(window=90,direction=higher_is_better,foo=1)", "多给了参数"),
("volume_ratio(fast=60,slow=5,direction=higher_is_better)", "快线不小于慢线"),
("no_such_template(window=5,direction=higher_is_better)", "模板不存在"),
("momentum_20(window=5,direction=higher_is_better)", "实例名不能带参数"),
("not_a_factor", "既不是实例名也不是参数化键"),
],
)
def test_bad_keys_are_rejected(self, bad: str, why: str) -> None:
with pytest.raises(FactorError):
get_factor(bad) # noqa: PT011 - 只关心「拒绝」,文案另有断言
def test_error_messages_tell_the_allowed_range(self) -> None:
with pytest.raises(FactorError, match="2 ~ 500"):
get_factor("momentum(window=999,direction=higher_is_better)")
with pytest.raises(FactorError, match="规范写法"):
get_factor("momentum(window=90)")
with pytest.raises(FactorError, match="快线窗口"):
get_factor("volume_ratio(fast=9,slow=9,direction=higher_is_better)")
def test_key_length_fits_the_column(self) -> None:
"""最长的合法键也要能进 factor_definition.name(String(128))—— 不然是运行期报错。"""
longest = canonical_key("volume_ratio", {"fast": 499, "slow": 500})
assert len(longest) < 128
def test_templates_expose_editable_params(self) -> None:
tpls = {t.name: t for t in list_templates()}
assert "momentum" in tpls and "dividend_yield" in tpls
specs = {s.name: s for s in tpls["momentum"].specs()}
assert specs["window"].kind == "int" and specs["window"].maximum == 500
assert specs["direction"].choices == ("higher_is_better", "lower_is_better")
# 股息率没有窗口参数:只有方向可编辑(不许凭空造出「窗口」)
dy = {s.name for s in get_template("dividend_yield").specs()}
assert dy == {"direction"}
class TestCatalogAndApi:
def test_create_then_read_rows_are_engine_projection(self, client) -> None:
rows = {r["name"]: r for r in client.get("/api/factors").json()}
assert rows["momentum_60"]["label"] == "动量(窗口 60,越高越好)"
assert rows["momentum_60"]["params"] == {
"window": 60,
"direction": "higher_is_better",
}
assert rows["momentum_60"]["source"] == "builtin"
assert {s["name"] for s in rows["momentum_60"]["param_specs"]} == {
"window",
"direction",
}
created = client.post("/api/factors", json={"template": "momentum", "params": {"window": 90}})
assert created.status_code == 201, created.text
body = created.json()
key = "momentum(window=90,direction=higher_is_better)"
assert body["name"] == key and body["label"] == "动量(窗口 90,越高越好)"
assert body["source"] == "custom" and body["resolvable"] is True
again = client.get("/api/factors").json()
assert key in {r["name"] for r in again}
def test_create_rejects_bad_params_and_duplicates(self, client) -> None:
client.get("/api/factors")
bad = client.post("/api/factors", json={"template": "momentum", "params": {"window": 0}})
assert bad.status_code == 422 and "2 ~ 500" in bad.json()["detail"]
unknown = client.post("/api/factors", json={"template": "nope", "params": {}})
assert unknown.status_code == 422 and "未知因子模板" in unknown.json()["detail"]
first = client.post("/api/factors", json={"template": "reversal", "params": {"window": 6}})
assert first.status_code == 201
dup = client.post("/api/factors", json={"template": "reversal", "params": {"window": 6}})
assert dup.status_code == 422 and "已存在" in dup.json()["detail"]
def test_templates_endpoint(self, client) -> None:
tpls = {t["name"]: t for t in client.get("/api/factors/templates").json()}
assert tpls["momentum"]["defaults"] == {
"window": 20,
"direction": "higher_is_better",
}
assert tpls["momentum"]["instances"] == ["momentum_20", "momentum_60", "momentum_120"]
def test_disable_only_affects_pickability(self, client) -> None:
client.get("/api/factors")
key = "momentum(window=77,direction=higher_is_better)"
client.post("/api/factors", json={"template": "momentum", "params": {"window": 77}})
off = client.patch("/api/factors", json={"name": key, "enabled": False})
assert off.status_code == 200 and off.json()["enabled"] is False
# 再读一次:同步不会把人的开关冲掉
rows = {r["name"]: r for r in client.get("/api/factors").json()}
assert rows[key]["enabled"] is False
# 停用不影响引擎解析(历史策略/归档照样能算)
assert get_factor(key)[0].lookback == 77
def test_builtin_cannot_be_disabled(self, client) -> None:
client.get("/api/factors")
resp = client.patch("/api/factors", json={"name": "momentum_60", "enabled": False})
assert resp.status_code == 422 and "内置因子" in resp.json()["detail"]
def test_patch_unknown_factor_is_404(self, client) -> None:
client.get("/api/factors")
resp = client.patch(
"/api/factors",
json={"name": "ma_bias(window=13,direction=higher_is_better)", "enabled": False},
)
assert resp.status_code == 404
def test_engine_text_converges_for_parameterized_rows(self, session) -> None:
"""参数化实例的口径文案同样按代码收敛,但 enabled 是人的配置、不许冲掉。"""
repo = SqlAlchemyFactorRepository(session)
sync_registry_factors(repo, session)
key = "momentum(window=88,direction=higher_is_better)"
row = create_parameterized_factor(repo, session, template="momentum", params={"window": 88})
assert row.name == key
# 手改文案 + 停用 → 同步后:文案回到代码文本,开关保留
repo.upsert_many(
[row.model_copy(update={"description": "手改的假口径", "enabled": False})]
)
session.commit()
assert sync_registry_factors(repo, session) >= 1
got = repo.get(key)
assert got.description == get_factor(key)[0].description
assert got.enabled is False, "停用是人配的,不能被代码投影冲回 True"
assert sync_registry_factors(repo, session) == 0, "再来一次应稳态零写入"
def test_unresolvable_hand_row_is_kept_and_flagged(self, session) -> None:
repo = SqlAlchemyFactorRepository(session)
sync_registry_factors(repo, session)
repo.upsert_many(
[
FactorDefinition(
name="someone_typo_factor",
description="人工登记的备注",
formula="x",
brief="b",
lookback=5,
requires=[],
)
]
)
session.commit()
assert sync_registry_factors(repo, session) == 0, "算不出来的行不参与收敛"
got = repo.get("someone_typo_factor")
assert got is not None and got.description == "人工登记的备注"
assert not is_resolvable("someone_typo_factor")
class TestParameterizedFactorAsCondition:
"""因子既能打分也能过滤:参数化实例在条件路径上同样要能用、文案不撒谎。"""
def test_resolvable_and_no_bogus_unit(self) -> None:
key = "momentum(window=45,direction=lower_is_better)"
assert reason_unsupported(key) == ""
field = get_field(key)
assert field is not None
assert field.kind == "num"
assert field.unit == "", "因子是无量纲量:不许给它挂单位(挂错了会误导输入)"
assert "动量" in field.label
def test_condition_needed_columns_accepts_parameterized_key(self) -> None:
from app.domain.entities.research import ConditionSpec
class _Q:
def __init__(self, conds):
self.conditions = conds
need = condition_needed_columns(
_Q([ConditionSpec(field="volume_ratio(fast=3,slow=9,direction=higher_is_better)", op="gte", value=1)])
)
assert "volume" in need, "参数化量比因子必须把 volume 列带进装配"
def test_unknown_field_message_mentions_parameterized_keys(self) -> None:
from app.domain.entities.research import ConditionSpec
class _Q:
def __init__(self, conds):
self.conditions = conds
with pytest.raises(ValueError, match="参数化因子键"):
condition_needed_columns(_Q([ConditionSpec(field="nope", op="gte", value=1)]))
def test_docs_render_the_parameterized_name(self) -> None:
from app.domain.entities.strategy import SelectionStrategy
from app.quant.strategy_doc import describe_strategy
key = "momentum(window=45,direction=higher_is_better)"
doc = describe_strategy(
SelectionStrategy(name="参数化因子演示", factors=[FactorSpec(name=key, weight=1.0)])
)
joined = "\n".join([doc.formula, *doc.steps])
assert key in joined, "说明书要写清用的是哪个参数版本(否则读者不知道窗口是多少)"
assert not any("未知因子" in w for w in doc.warnings)