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

319 lines
16 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)。
覆盖四件事:
1. **注册表与引擎不漂移**:内置字段必须全部是引擎真能算的(is_supported_field);
2. **拒绝伪字段**:日期字段/不存在的字段/拼错的字段一律拒绝并给出理由;
3. **目录 API**:seed 幂等、只补不删(改过的文案不被覆盖)、自定义增删改、内置不可删;
4. **比较加固**:类型不匹配的条件返回 False 而不是抛异常(否则用户手滑 → 500)。
"""
from __future__ import annotations
from datetime import date
import pytest
from app.api import deps
from app.application.services import condition_field_catalog as svc
from app.domain.entities.condition_field import ConditionField
from app.infrastructure.persistence.sqlalchemy.base import Base
from app.infrastructure.persistence.sqlalchemy.repositories.condition_field_impl import (
SqlAlchemyConditionFieldRepository,
)
from app.main import app
from app.quant.condition_fields import (
builtin_fields,
curated_fields,
is_supported_field,
reason_unsupported,
unit_allowed,
unit_options,
)
from app.quant.selection import _compare
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 / 'fields.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 TestRegistry:
def test_every_builtin_field_is_engine_computable(self) -> None:
"""注册表里写了的字段,引擎必须真能算 —— 否则字段库就是在骗人。"""
bad = [d.name for d in builtin_fields() if not is_supported_field(d.name)]
assert bad == []
def test_builtin_names_unique(self) -> None:
names = [d.name for d in builtin_fields()]
assert len(names) == len(set(names))
def test_curated_are_subset_of_builtin(self) -> None:
assert {d.name for d in curated_fields()} <= {d.name for d in builtin_fields()}
def test_every_field_has_label_and_description(self) -> None:
"""下拉要显示中文名,含义提示不能是空白(否则「字段说明含义」没达成)。"""
for d in builtin_fields():
assert d.label.strip(), d.name
assert d.description.strip(), d.name
assert d.kind in ("num", "str"), d.name
assert d.ops, d.name
def test_str_fields_only_equality_ops(self) -> None:
"""字符串字段不能比大小:可选比较符里不该出现 >/≥/<(那些恒为假)。"""
for d in builtin_fields():
if d.kind == "str":
assert set(d.ops) <= {"eq", "ne", "in", "not_in"}, d.name
def test_unit_ladders_are_well_formed(self) -> None:
"""单位阶梯:首项必须是基准单位且系数 1.0,系数递增(否则换算会反向)。"""
checked = 0
for d in builtin_fields():
if not d.units:
assert d.unit_options == ((d.unit, 1.0),), d.name
continue
checked += 1
base, factor = d.units[0]
assert (base, factor) == (d.unit, 1.0), d.name
factors = [f for _, f in d.units]
assert factors == sorted(factors) and len(set(factors)) == len(factors), d.name
assert checked >= 6, "至少金额/股数/股本的字段要有可选单位"
def test_unit_ladder_matches_data_source_scale(self) -> None:
"""换算系数必须与数据源落库口径一致(写错了就是静默的 10000 倍误差)。"""
assert unit_options("total_mv") == [("万元", 1.0), ("亿元", 10000.0)]
assert unit_options("amount") == [("元", 1.0), ("万元", 10000.0), ("亿元", 100000000.0)]
assert unit_options("volume") == [("股", 1.0), ("手", 100.0), ("万手", 1000000.0)]
# 百分数/倍数没有备选:换个说法只会制造误读
assert unit_options("dv_ratio") == [("%", 1.0)]
assert unit_options("pe") == [("倍", 1.0)]
def test_unit_allowed_only_inside_ladder(self) -> None:
assert unit_allowed("total_mv", "亿元")
assert unit_allowed("total_mv", "万元")
assert not unit_allowed("total_mv", "元")
assert not unit_allowed("dv_ratio", "小数")
assert not unit_allowed("close", "分")
@pytest.mark.parametrize(
"name",
["static.list_date", "static.delist_date", "static.nope", "fundamental.report_date",
"fundamental.announce_date", "totally_made_up", ""],
)
def test_rejects_uncomputable_fields(self, name: str) -> None:
assert not is_supported_field(name)
assert reason_unsupported(name).strip()
@pytest.mark.parametrize(
"name",
["close", "volume", "amount", "ma20", "ma60", "dv_ratio", "pe", "total_mv",
"static.industry", "static.market", "fundamental.roe", "momentum_60", "dividend_yield"],
)
def test_accepts_real_fields(self, name: str) -> None:
assert is_supported_field(name), reason_unsupported(name)
class TestCompareHardening:
"""条件求值不能因为类型不匹配抛异常(用户手滑的字段名不该 500)。"""
def test_date_field_vs_number_returns_false(self) -> None:
assert _compare(date(2020, 1, 1), 20200101, "gt") is False
assert _compare(date(2020, 1, 1), "2020-01-01", "lte") is False
def test_in_with_non_container_returns_false(self) -> None:
assert _compare("银行", 5, "in") is False
assert _compare("银行", 5, "not_in") is False
def test_normal_comparisons_still_work(self) -> None:
assert _compare(10.0, 5.0, "gt") is True
assert _compare("银行", ["银行", "白酒"], "in") is True
assert _compare("银行", ["白酒"], "not_in") is True
assert _compare(None, 5.0, "lt") is False
assert _compare(None, 5.0, "ne") is True
class TestFieldCatalogApi:
def test_get_seeds_builtin_fields(self, client: TestClient) -> None:
rows = client.get("/api/condition-fields").json()
assert len(rows) == len(curated_fields())
names = {r["name"] for r in rows}
assert {"close", "dv_ratio", "static.industry", "fundamental.roe"} <= names
dv = next(r for r in rows if r["name"] == "dv_ratio")
assert dv["kind"] == "num"
assert dv["unit"] == "%"
assert dv["source"] == "builtin"
assert "股息率" in (dv["label"] + dv["description"])
def test_seed_is_idempotent(self, client: TestClient) -> None:
first = client.get("/api/condition-fields").json()
second = client.get("/api/condition-fields").json()
assert len(first) == len(second)
def test_available_lists_supported_but_uncurated(self, client: TestClient) -> None:
client.get("/api/condition-fields") # 先 seed
rows = client.get("/api/condition-fields/available").json()
assert rows, "至少应有若干「引擎支持但未默认展示」的字段可选"
for r in rows:
assert is_supported_field(r["name"]), r["name"]
assert r["label"] and r["description"]
def test_response_carries_base_unit_and_unit_ladder(self, client: TestClient) -> None:
"""响应必须给出「基准单位」与可选界面单位(前端据此换算输入/回显)。"""
rows = {r["name"]: r for r in client.get("/api/condition-fields").json()}
mv = rows["total_mv"]
assert mv["base_unit"] == "万元", "基准单位 = 引擎存储单位"
assert mv["unit"] == "万元", "初始界面单位 = 基准单位"
assert mv["units"] == [
{"unit": "万元", "factor": 1.0},
{"unit": "亿元", "factor": 10000.0},
]
# 没有备选单位的字段:只有一个选项,界面上不给选择
assert rows["close"]["units"] == [{"unit": "元", "factor": 1.0}]
def test_update_unit_only_inside_ladder(self, client: TestClient) -> None:
"""单位只能在阶梯里选:自由文本 → 422(否则就是标签与实际口径不一致的静默错误)。"""
client.get("/api/condition-fields") # 先 seed(total_mv 是内置字段)
bad = client.put("/api/condition-fields/total_mv", json={"unit": "亿亿元"})
assert bad.status_code == 422, bad.text
assert "万元" in bad.json()["detail"] and "亿元" in bad.json()["detail"]
# 不在该字段阶梯里的别的单位也要拒(元 是 amount 的单位,不是 total_mv 的)
assert client.put("/api/condition-fields/total_mv", json={"unit": "元"}).status_code == 422
assert client.put("/api/condition-fields/dv_ratio", json={"unit": "小数"}).status_code == 422
ok = client.put("/api/condition-fields/total_mv", json={"unit": "亿元"})
assert ok.status_code == 200, ok.text
assert ok.json()["unit"] == "亿元"
assert ok.json()["base_unit"] == "万元", "选界面单位不改基准单位(引擎口径不动)"
# 空串 = 回到基准单位
reset = client.put("/api/condition-fields/total_mv", json={"unit": ""})
assert reset.status_code == 200 and reset.json()["unit"] == "万元"
def test_create_unit_only_inside_ladder(self, client: TestClient) -> None:
client.get("/api/condition-fields")
bad = client.post("/api/condition-fields", json={"name": "ps", "unit": "亿亿元"})
assert bad.status_code == 422, bad.text
ok = client.post("/api/condition-fields", json={"name": "ps", "unit": "倍"})
assert ok.status_code == 200 and ok.json()["unit"] == "倍"
def test_create_custom_field(self, client: TestClient) -> None:
created = client.post(
"/api/condition-fields",
json={"name": "ps", "label": "市销率(我的叫法)", "description": "自定义说明"},
)
assert created.status_code == 200, created.text
row = created.json()
assert row["source"] == "custom"
assert row["kind"] == "num" # 类型来自引擎,不是调用方说了算
assert row["label"] == "市销率(我的叫法)"
# available 里不应再出现
rest = {r["name"] for r in client.get("/api/condition-fields/available").json()}
assert "ps" not in rest
def test_create_rejects_uncomputable_field(self, client: TestClient) -> None:
resp = client.post("/api/condition-fields", json={"name": "static.list_date"})
assert resp.status_code == 422
assert "日期" in resp.text
def test_create_rejects_duplicate(self, client: TestClient) -> None:
client.post("/api/condition-fields", json={"name": "ps"})
again = client.post("/api/condition-fields", json={"name": "ps"})
assert again.status_code == 422
assert "已在字段库" in again.text
def test_create_rejects_unknown_key(self, client: TestClient) -> None:
assert client.post("/api/condition-fields", json={"name": "ps", "knd": "num"}).status_code == 422
def test_update_edits_label_and_disables(self, client: TestClient) -> None:
client.get("/api/condition-fields")
resp = client.put(
"/api/condition-fields/close",
json={"label": "收盘价(我改的)", "description": "自定义口径说明", "enabled": False},
)
assert resp.status_code == 200, resp.text
row = resp.json()
assert row["label"] == "收盘价(我改的)"
assert row["enabled"] is False
assert row["source"] == "builtin" # 内置字段改了文案仍是内置
def test_update_cannot_change_name_or_kind(self, client: TestClient) -> None:
"""name/kind 是引擎事实:请求里带上它们必须报错,而不是被静默忽略。"""
assert client.put("/api/condition-fields/close", json={"kind": "str"}).status_code == 422
assert client.put("/api/condition-fields/close", json={"name": "pe"}).status_code == 422
def test_update_unknown_returns_404(self, client: TestClient) -> None:
assert client.put("/api/condition-fields/nope", json={"label": "x"}).status_code == 404
def test_update_rejects_empty_label(self, client: TestClient) -> None:
client.get("/api/condition-fields")
assert client.put("/api/condition-fields/close", json={"label": " "}).status_code == 422
def test_disabled_hidden_from_picker(self, client: TestClient) -> None:
client.get("/api/condition-fields")
client.put("/api/condition-fields/close", json={"enabled": False})
picker = client.get("/api/condition-fields?include_disabled=false").json()
assert "close" not in {r["name"] for r in picker}
allrows = client.get("/api/condition-fields").json()
assert "close" in {r["name"] for r in allrows} # 管理页仍看得到,才能重新启用
def test_builtin_cannot_be_deleted(self, client: TestClient) -> None:
client.get("/api/condition-fields")
resp = client.delete("/api/condition-fields/close")
assert resp.status_code == 400
assert "停用" in resp.text
assert "close" in {r["name"] for r in client.get("/api/condition-fields").json()}
def test_custom_can_be_deleted(self, client: TestClient) -> None:
client.post("/api/condition-fields", json={"name": "ps"})
assert client.delete("/api/condition-fields/ps").status_code == 200
assert "ps" not in {r["name"] for r in client.get("/api/condition-fields").json()}
assert client.delete("/api/condition-fields/ps").status_code == 404
class TestSeedDoesNotClobber:
def test_edited_builtin_text_survives_reseed(self, tmp_path) -> None:
"""seed「只补不删」:用户改过的中文名/含义不能被下次读取冲掉。"""
engine = create_engine(f"sqlite:///{tmp_path / 'reseed.db'}", future=True)
Base.metadata.create_all(engine)
Session = sessionmaker(bind=engine, expire_on_commit=False)
with Session() as session:
repo = SqlAlchemyConditionFieldRepository(session)
svc.sync_builtin_fields(repo, session)
svc.update_field(repo, session, "close", label="我的收盘价")
added = svc.sync_builtin_fields(repo, session) # 再 seed 一次
assert added == 0
assert repo.get("close").label == "我的收盘价"
def test_new_builtin_is_backfilled(self, tmp_path) -> None:
"""代码里新增的默认字段下次读取要补进来(否则字段库会永远缺新字段)。"""
engine = create_engine(f"sqlite:///{tmp_path / 'backfill.db'}", future=True)
Base.metadata.create_all(engine)
Session = sessionmaker(bind=engine, expire_on_commit=False)
with Session() as session:
repo = SqlAlchemyConditionFieldRepository(session)
repo.save(ConditionField(name="close", label="仅此一条", kind="num"))
session.commit()
assert svc.sync_builtin_fields(repo, session) == len(curated_fields()) - 1
def test_uncurated_fields_are_not_seeded(self, tmp_path) -> None:
"""curated=False 的字段默认不进库,否则「新增字段」永远无字段可选。"""
engine = create_engine(f"sqlite:///{tmp_path / 'uncurated.db'}", future=True)
Base.metadata.create_all(engine)
Session = sessionmaker(bind=engine, expire_on_commit=False)
with Session() as session:
repo = SqlAlchemyConditionFieldRepository(session)
svc.sync_builtin_fields(repo, session)
seeded = {f.name for f in repo.list()}
assert seeded == {d.name for d in curated_fields()}
assert "ps" not in seeded # 非默认字段留给用户按需添加