"""字段库测试(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 # 非默认字段留给用户按需添加