字段库(本次新增的表与接口): - `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。
319 lines
16 KiB
Python
319 lines
16 KiB
Python
"""字段库测试(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 # 非默认字段留给用户按需添加 |