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。
This commit is contained in:
Simon
2026-10-01 16:33:32 +08:00
parent 40bd603b44
commit 2e90f3eeac
39 changed files with 3280 additions and 244 deletions
+22 -1
View File
@@ -105,7 +105,18 @@ class TestV3Tools:
out = _invoke(tools, "inspect_factor", {"name": "momentum_60"})
assert "momentum_60" in out and "公式" in out and "lookback" in out
miss = _invoke(tools, "inspect_factor", {"name": "no_such"})
assert "不在目录" in miss
# 走引擎解析:报「不可用 + 可用列表」比含糊的「不在目录」有用
# (参数化因子常常还没进目录就能算,所以目录不是准入门槛)
assert "不可用" in miss
def test_inspect_parameterized_factor(self, tools) -> None:
"""参数化键(还没进目录)也要能问出参数:这是「暴露真实筛选参数」的一环。"""
key = "momentum(window=90,direction=lower_is_better)"
out = _invoke(tools, "inspect_factor", {"name": key})
assert key in out
assert "窗口" in out or "window=90" in out
assert "越低越好" in out and "参数:window=90" in out
assert "90" in out
def test_create_composite_factor(self, tools) -> None:
out = _invoke(
@@ -117,6 +128,16 @@ class TestV3Tools:
{"name": "x", "factors": "no_such:1"})
assert "无法创建" in bad
def test_composite_accepts_parameterized_factor(self, tools) -> None:
"""参数化键里有逗号:逗号切分必须括号感知,否则会被劈成两个「不存在的因子」。"""
key = "momentum(window=90,direction=lower_is_better)"
out = _invoke(
tools, "create_composite_factor",
{"name": "参数化动量组合", "factors": f"{key}:0.7,volatility_60:0.3"},
)
assert "组合已保存" in out, out
assert key in out and "volatility_60" in out
def test_get_backtest_result(self, tools) -> None:
out = _invoke(tools, "get_backtest_result", {"experiment_id": "EXP-D4"})
assert "总收益" in out and "成交" in out
+21
View File
@@ -53,6 +53,14 @@ class TestConfigApi:
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:
@@ -107,3 +115,16 @@ class TestCombosApi:
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() == []
-1
View File
@@ -12,7 +12,6 @@ from datetime import date, timedelta
import pandas as pd
import pytest
from app.domain.entities.combo import BacktestCombo
from app.domain.entities.research import CostSpec
from app.quant.combo_engine import HoldingBandRunner, borda_combine
+319
View File
@@ -0,0 +1,319 @@
"""字段库测试(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 # 非默认字段留给用户按需添加
+54
View File
@@ -142,3 +142,57 @@ class TestRegistrySyncRegression:
names2 = {r["name"] for r in client.get("/api/factors").json()}
assert "my_custom_factor" in names2, "只补不删:自定义因子必须保留"
assert "dividend_yield" in names2
class TestCatalogIsRegistryProjection:
"""目录是注册表的投影(2026-10 明确语义):口径字段按代码纠正,自定义行不碰。
原先是「只在缺名字时才 upsert 全量」,于是手改内置因子文案能存活到「代码里出现
新因子」那一刻再被无声覆盖 —— 行为不确定。现在改成按字段差集确定性收敛。
"""
def test_hand_edited_registry_row_converges_back_to_code(self, client, tmp_path) -> None:
from app.quant.factors import get_factor
from sqlalchemy import create_engine, text
engine = create_engine(f"sqlite:///{tmp_path / 'api.db'}", future=True)
with engine.begin() as conn:
conn.execute(
text(
"UPDATE factor_definition SET description = '手改的假口径', "
"lookback = 999, direction = 'lower_is_better' "
"WHERE name = 'dividend_yield'"
)
)
rows = {r["name"]: r for r in client.get("/api/factors").json()}
defn, _ = get_factor("dividend_yield")
got = rows["dividend_yield"]
assert got["description"] == defn.description, "口径文案必须按代码改回(否则说明书会撒谎)"
assert got["lookback"] == defn.lookback
assert got["direction"] == defn.direction, "方向被手改会让读者以为越大越差"
def test_steady_state_writes_nothing(self, session) -> None:
from app.application.services.factor_catalog import sync_registry_factors
repo = SqlAlchemyFactorRepository(session)
assert sync_registry_factors(repo, session) == len(list_factors()), "首次:全部补齐"
assert sync_registry_factors(repo, session) == 0, "稳态:目录 == 注册表 → 零写入"
def test_custom_row_untouched_and_never_deleted(self, session) -> None:
from app.application.services.factor_catalog import sync_registry_factors
repo = SqlAlchemyFactorRepository(session)
sync_registry_factors(repo, session)
repo.upsert_many(
[
FactorDefinition(
name="my_note_factor", description="人工登记的备注", formula="x",
brief="b", lookback=5, requires=[],
)
]
)
session.commit()
assert sync_registry_factors(repo, session) == 0, "自定义行不算 stale,不触发写入"
got = repo.get("my_note_factor")
assert got is not None and got.description == "人工登记的备注"
assert got.name not in {d.name for d in list_factors()}, "前提:它确实不在注册表里"
+330
View File
@@ -0,0 +1,330 @@
"""因子参数化测试(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)
+109
View File
@@ -2,6 +2,7 @@
from __future__ import annotations
import json
import sqlite3
from pathlib import Path
@@ -38,6 +39,7 @@ def test_upgrade_head_creates_phase1_tables(tmp_path) -> None:
"trading_calendar",
"financial_indicator",
"sync_log",
"condition_field",
"alembic_version",
}
assert expected <= tables
@@ -63,3 +65,110 @@ def test_upgrade_head_idempotent(tmp_path) -> None:
cfg = _alembic_config(db_path)
command.upgrade(cfg, "head")
command.upgrade(cfg, "head") # 二次执行不报错
def test_refresh_stale_strategy_descriptions(tmp_path) -> None:
"""存量策略的陈旧说明被重算,人工说明保留,空说明补全(b4c5… → c5d6… 全链路)。
模拟真实历史:在「策略表已建、还没做组合重构」的版本上插入带全套回测参数的旧行,
一路 upgrade head —— 既验证 b4c5d6e7f8a9 剥掉配置键,也验证 c5d6e7f8a9 重算说明。
"""
db_path = tmp_path / "refresh.db"
cfg = _alembic_config(db_path)
command.upgrade(cfg, "e1f2a3b4c5d6") # strategy 表建好、尚未重构
stale_desc = (
"全市场(剔除 ST),按股息率排序取出前 20 只等权持有,每 6 个月重新择股、"
"每 6 个月调仓,后复权口径、按调仓日收盘价成交(含佣金 0.03%/印花税 0.05%/滑点 0.1%)。"
)
base_cfg = {
"universe": {"market": "CN_A", "exclude_st": True},
"factors": [{"name": "dividend_yield", "weight": 1}],
"conditions": [],
}
legacy_cfg = {**base_cfg, "selection": {"top_n": 20}, "rebalance": "monthly",
"costs": {"commission_rate": 0.0003}, "price_adjustment": "hfq"}
con = sqlite3.connect(db_path)
con.execute(
"INSERT INTO strategy (id,name,description,spec_type,config_json,version,created_at) "
"VALUES (?,?,?,?,?,?,?)",
("STG-STALE", "高股息 Top20(案例口径)", stale_desc, "backtest",
json.dumps(legacy_cfg, ensure_ascii=False), "1", "2026-01-01 00:00:00"),
)
con.execute(
"INSERT INTO strategy (id,name,description,spec_type,config_json,version,created_at) "
"VALUES (?,?,?,?,?,?,?)",
("STG-HUMAN", "我的成长股", "只看 ROE 与动量,人工撰写的说明不要被覆盖。", "backtest",
json.dumps({**base_cfg, "factors": [{"name": "momentum_60", "weight": 1}]},
ensure_ascii=False), "1", "2026-01-02 00:00:00"),
)
con.execute(
"INSERT INTO strategy (id,name,description,spec_type,config_json,version,created_at) "
"VALUES (?,?,?,?,?,?,?)",
("STG-EMPTY", "空说明策略", "", "backtest",
json.dumps(base_cfg, ensure_ascii=False), "1", "2026-01-03 00:00:00"),
)
con.commit()
con.close()
command.upgrade(cfg, "head")
con = sqlite3.connect(db_path)
try:
rows = {
r[0]: {"desc": r[1], "spec": r[2], "cfg": r[3]}
for r in con.execute(
"SELECT id, description, spec_type, config_json FROM strategy"
).fetchall()
}
tables = {
r[0] for r in con.execute(
"select name from sqlite_master where type='table'"
).fetchall()
}
finally:
con.close()
# ① 陈旧自动说明被重算:不再提佣金/印花税/滑点/调仓,且是新口径文案
stale = rows["STG-STALE"]
assert stale["desc"] != stale_desc
for marker in ("佣金", "印花税", "滑点", "调仓", "择股", "复权口径"):
assert marker not in stale["desc"], f"陈旧说明仍含 {marker}: {stale['desc']}"
assert "选股策略" in stale["desc"] and "回测组合" in stale["desc"]
# ② 人工撰写的说明原样保留(迁移不覆盖用户文本)
assert rows["STG-HUMAN"]["desc"] == "只看 ROE 与动量,人工撰写的说明不要被覆盖。"
# ③ 空说明按当前口径补全
assert rows["STG-EMPTY"]["desc"].strip()
assert "选股策略" in rows["STG-EMPTY"]["desc"]
# ④ b4c5… 的职责仍在:旧回测参数从 config_json 剥掉、spec_type 收敛
stripped = json.loads(rows["STG-STALE"]["cfg"])
assert "costs" not in stripped and "rebalance" not in stripped and "selection" not in stripped
assert stripped["factors"][0]["name"] == "dividend_yield"
assert rows["STG-STALE"]["spec"] == "selection"
assert {"global_config", "backtest_combo"} <= tables
def test_condition_field_table_columns(tmp_path) -> None:
"""字段库表结构:name 主键 + 中文名/含义/类型/分组/来源/启用状态(2026-10)。"""
db_path = tmp_path / "fields.db"
command.upgrade(_alembic_config(db_path), "head")
con = sqlite3.connect(db_path)
try:
cols = {row[1] for row in con.execute("pragma table_info(condition_field)").fetchall()}
pk = [
row[1]
for row in con.execute("pragma table_info(condition_field)").fetchall()
if row[5]
]
finally:
con.close()
assert {
"name", "label", "description", "kind", "group_name",
"unit", "source", "enabled", "sort_order", "created_at", "updated_at",
} <= cols
assert pk == ["name"]
+65
View File
@@ -7,15 +7,20 @@
from __future__ import annotations
import json
from datetime import datetime
import pytest
from app.api import deps
from app.domain.entities.strategy import SelectionStrategy
from app.infrastructure.persistence.sqlalchemy.base import Base
from app.infrastructure.persistence.sqlalchemy.models.strategy import StrategyModel
from app.infrastructure.persistence.sqlalchemy.repositories.strategy_impl import (
SqlAlchemyStrategyRepository,
)
from app.main import app
from fastapi.testclient import TestClient
from pydantic import ValidationError
from sqlalchemy import create_engine
from sqlalchemy.orm import sessionmaker
@@ -72,6 +77,52 @@ class TestStrategyRepository:
):
assert forbidden not in dumped, f"选股策略不应含回测参数字段 {forbidden}"
def test_rejects_legacy_backtest_params(self) -> None:
"""混入旧版回测参数必须**报错**,不能静默丢弃(否则调用方以为设上了)。"""
with pytest.raises(ValidationError) as exc:
SelectionStrategy(
name="带旧参数的策略",
factors=[{"name": "momentum_60", "weight": 1}],
costs={"commission_rate": 0.001},
rebalance="monthly",
initial_capital=500_000,
)
# 三个未知键都应被点名(便于调用方知道该搬去哪里)
for key in ("costs", "rebalance", "initial_capital"):
assert key in str(exc.value)
def test_spec_type_only_selection(self) -> None:
"""spec_type 取值域收敛为 selection(本实体只表示选股策略)。"""
with pytest.raises(ValidationError):
SelectionStrategy(
name="旧类型", spec_type="backtest",
factors=[{"name": "momentum_60", "weight": 1}],
)
assert _st().spec_type == "selection"
def test_legacy_rows_with_extra_keys_still_readable(self, session) -> None:
"""历史行 config_json 残留旧键时仍能读出(仓储读出前剔除),forbid 不影响兼容。"""
session.add(
StrategyModel(
id="STG-LEGACY", name="历史策略", description="历史说明",
spec_type="backtest",
config_json=json.dumps({
"universe": {"exclude_st": True},
"factors": [{"name": "dividend_yield", "weight": 1}],
"conditions": [],
"selection": {"top_n": 20},
"rebalance": "monthly",
"costs": {"commission_rate": 0.0003},
}, ensure_ascii=False),
version="1", created_at=datetime(2026, 1, 1),
)
)
session.commit()
got = SqlAlchemyStrategyRepository(session).get("STG-LEGACY")
assert got is not None
assert got.factors[0].name == "dividend_yield"
assert got.spec_type == "selection" # 旧列值不参与实体(回落到默认)
@pytest.fixture()
def client(tmp_path):
@@ -130,3 +181,17 @@ class TestStrategiesApi:
json={"period": ["2024-01-01", "2024-06-01"]},
)
assert resp.status_code in (404, 405)
def test_create_with_legacy_params_422(self, client) -> None:
"""旧调用方把回测参数塞进策略 → 422 并点名未知字段(而非 200 静默丢弃)。"""
body = {
"name": "旧调用方",
"factors": [{"name": "momentum_60", "weight": 1}],
"rebalance": "monthly",
"costs": {"commission_rate": 0.001},
}
resp = client.post("/api/strategies", json=body)
assert resp.status_code == 422
assert "costs" in resp.text and "rebalance" in resp.text
# 未落库:策略列表仍为空
assert client.get("/api/strategies").json() == []
+38
View File
@@ -253,6 +253,44 @@ class TestStepsAndWarnings:
)
assert any("替补" in s for s in doc.steps)
def test_selection_steps_carry_no_own_numbering(self) -> None:
"""选股策略的执行步骤不带自带序号:前端渲染进 <ol>,后端再写「1.」会双重编号。"""
st = SelectionStrategy(
name="高股息",
description="x",
factors=[FactorSpec(name="dividend_yield", weight=1)],
conditions=[],
)
doc = describe_strategy(st)
assert doc.steps
for s in doc.steps:
assert not s.lstrip().startswith(("1.", "2.", "3.", "4.", "5.")), s
# 步骤内容本身仍要在(去掉的只是序号)
assert "股票池" in doc.steps[0]
assert "因子" in doc.steps[-1]
def test_condition_literals_carry_base_unit(self) -> None:
"""字面量条件必须带**基准单位**:库里存的就是基准单位值,裸数字会被读成别的量级。
单位只做界面换算(见 quant/condition_fields 的单位阶梯),所以文档里写「50000 万元」
才是引擎真正比较的口径;字段间比较(ref)两侧同单位,不加后缀。
"""
st = SelectionStrategy(
name="单位口径",
description="x",
factors=[FactorSpec(name="dividend_yield", weight=1)],
conditions=[
ConditionSpec(field="total_mv", op="gte", value=50000),
ConditionSpec(field="close", op="gt", ref="ma60"),
ConditionSpec(field="industry", op="in", value=["银行", "白酒"]),
],
)
doc = describe_strategy(st)
joined = "\n".join([doc.formula, *doc.steps])
assert "total_mv >= 50000 万元" in joined
assert "close > 字段 ma60" in joined and "ma60 元" not in joined
assert "['银行', '白酒']" in joined
def test_universe_step_conditional_and_time_accurate(self) -> None:
"""股票池步骤只写实际生效的过滤,且时点口径必须正确(起始日过滤一次)。