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
+67 -18
View File
@@ -43,6 +43,46 @@ def _day(text: str) -> date:
return date.fromisoformat(text) return date.fromisoformat(text)
def _split_factor_list(raw: str) -> list[str]:
"""按逗号切因子列表,但**不切参数化因子键里的逗号**。
参数化因子的名字把参数写全了(`momentum(window=90,direction=higher_is_better)`),
直接 `.split(",")` 会把它劈成「momentum(window=90」和「direction=…):0.7」两段,
模型与用户只会收到「因子不存在」这种看不懂的错。括号深度感知的切分让两种写法都能用:
momentum_60,volatility_60
momentum(window=90,direction=lower_is_better),volatility_60
"""
out: list[str] = []
depth = 0
buf: list[str] = []
for ch in raw:
if ch == "(":
depth += 1
elif ch == ")":
depth = max(0, depth - 1)
if ch == "," and depth == 0:
out.append("".join(buf).strip())
buf = []
else:
buf.append(ch)
out.append("".join(buf).strip())
return [x for x in out if x]
def _split_name_weight(part: str) -> tuple[str, str]:
"""把 `name:weight` 按**括号外**的第一个冒号切开(参数化键里的 `=`/`,` 不受影响)。"""
depth = 0
for i, ch in enumerate(part):
if ch == "(":
depth += 1
elif ch == ")":
depth = max(0, depth - 1)
elif ch == ":" and depth == 0:
return part[:i].strip(), part[i + 1 :].strip()
return part.strip(), ""
def _pick(mapping: dict, key: str, default=None): def _pick(mapping: dict, key: str, default=None):
val = mapping.get(key, default) val = mapping.get(key, default)
if isinstance(val, str): if isinstance(val, str):
@@ -147,7 +187,7 @@ def build_tools(factories: dict | None = None) -> list[Tool]:
return _run_spec(spec, f"因子 {name} 测试") return _run_spec(spec, f"因子 {name} 测试")
def run_backtest(args: dict) -> str: def run_backtest(args: dict) -> str:
factor_names = [f.strip() for f in str(_pick(args, "factors", "momentum_60")).split(",")] factor_names = _split_factor_list(str(_pick(args, "factors", "momentum_60")))
top_n = int(_pick(args, "top_n", 5) or 5) top_n = int(_pick(args, "top_n", 5) or 5)
rebalance = str(_pick(args, "rebalance", "monthly")) rebalance = str(_pick(args, "rebalance", "monthly"))
exclude_st = bool(_pick(args, "exclude_st", True)) exclude_st = bool(_pick(args, "exclude_st", True))
@@ -206,7 +246,7 @@ def build_tools(factories: dict | None = None) -> list[Tool]:
return [x.strip().upper() for x in raw.split(",") if x.strip()][:60] return [x.strip().upper() for x in raw.split(",") if x.strip()][:60]
def screen_stocks(args: dict) -> str: def screen_stocks(args: dict) -> str:
factors = [x.strip() for x in str(_pick(args, "factors", "momentum_60")).split(",") if x.strip()] factors = _split_factor_list(str(_pick(args, "factors", "momentum_60")))
top_n = int(_pick(args, "top_n", 10) or 10) top_n = int(_pick(args, "top_n", 10) or 10)
as_of = _day(str(_pick(args, "as_of", date.today().isoformat()))) as_of = _day(str(_pick(args, "as_of", date.today().isoformat())))
symbols = _scope_symbols(str(_pick(args, "symbols", "") or "")) symbols = _scope_symbols(str(_pick(args, "symbols", "") or ""))
@@ -256,7 +296,7 @@ def build_tools(factories: dict | None = None) -> list[Tool]:
return "\n".join(out) return "\n".join(out)
def generate_signals(args: dict) -> str: def generate_signals(args: dict) -> str:
factors = [x.strip() for x in str(_pick(args, "factors", "momentum_60")).split(",") if x.strip()] factors = _split_factor_list(str(_pick(args, "factors", "momentum_60")))
as_of = _day(str(_pick(args, "as_of", date.today().isoformat()))) as_of = _day(str(_pick(args, "as_of", date.today().isoformat())))
symbols = _scope_symbols(str(_pick(args, "symbols", "") or "")) symbols = _scope_symbols(str(_pick(args, "symbols", "") or ""))
query = SelectionQuery( query = SelectionQuery(
@@ -288,9 +328,8 @@ def build_tools(factories: dict | None = None) -> list[Tool]:
if not name: if not name:
return "请提供 name" return "请提供 name"
factors = [ factors = [
{"name": x.strip(), "weight": 1.0} {"name": x, "weight": 1.0}
for x in str(_pick(args, "factors", "momentum_60")).split(",") for x in _split_factor_list(str(_pick(args, "factors", "momentum_60")))
if x.strip()
] ]
if not factors: if not factors:
return "请提供至少一个 factors(逗号分隔)" return "请提供至少一个 factors(逗号分隔)"
@@ -316,28 +355,38 @@ def build_tools(factories: dict | None = None) -> list[Tool]:
def inspect_factor(args: dict) -> str: def inspect_factor(args: dict) -> str:
name = str(_pick(args, "name", "")) name = str(_pick(args, "name", ""))
# 先问引擎:目录里有没有这行是「管理」问题,引擎算不算得出来才是「能不能用」。
# 参数化因子(momentum(window=90,direction=…))经常还没进目录就被引用,也能算。
try:
defn, _fn = get_factor(name)
except FactorError as exc:
return f"因子不可用:{exc}"
with session_factory() as session: with session_factory() as session:
row = SqlAlchemyFactorRepository(session).get(name) row = SqlAlchemyFactorRepository(session).get(name)
if row is None: params = ",".join(f"{k}={v}" for k, v in defn.params.items())
return f"因子 {name} 不在目录(可用列表:GET /api/factors)" head = f"{defn.label}({defn.name})" if defn.label else defn.name
return ( return (
f"{row.name}:{row.description}\n公式:{row.formula}\n方向:" f"{head}:{defn.description}\n公式:{defn.formula}\n方向:"
f"{'越高越好' if row.direction == 'higher_is_better' else '越低越好'}" f"{'越高越好' if defn.direction == 'higher_is_better' else '越低越好'}"
f"(lookback {row.lookback},输入 {row.requires})\n简介:{row.brief}" f"(lookback {defn.lookback},输入 {defn.requires})\n"
f"参数:{params or '(无:内置实例名固定口径)'}\n"
f"来源:{'代码注册表内置' if defn.source == 'builtin' else '目录里的参数化实例'}"
f"{';在目录中已停用(仍可被引用)' if row is not None and not row.enabled else ''}\n"
f"简介:{defn.brief}"
) )
def create_composite_factor(args: dict) -> str: def create_composite_factor(args: dict) -> str:
name = str(_pick(args, "name", "")) name = str(_pick(args, "name", ""))
raw = str(_pick(args, "factors", "")) raw = str(_pick(args, "factors", ""))
if not name or not raw: if not name or not raw:
return "请提供 name 与 factors(格式:momentum_60:0.7,volatility_60:0.3)" return (
"请提供 name 与 factors(格式:momentum_60:0.7,volatility_60:0.3;"
"参数化因子写成 momentum(window=90,direction=lower_is_better):0.7)"
)
comps: list[CompositeComponent] = [] comps: list[CompositeComponent] = []
for part in raw.split(","): for part in _split_factor_list(raw):
if not part.strip(): fname, weight_text = _split_name_weight(part)
continue weight = float(weight_text) if weight_text else 1.0
seg = part.strip().split(":")
fname = seg[0].strip()
weight = float(seg[1]) if len(seg) > 1 and seg[1].strip() else 1.0
if not fname: if not fname:
continue continue
try: try:
+3 -1
View File
@@ -17,6 +17,8 @@ POST /api/combos/run 提交临时组合(不保存)为异步 Job
from __future__ import annotations from __future__ import annotations
from datetime import datetime
from fastapi import APIRouter, BackgroundTasks, HTTPException from fastapi import APIRouter, BackgroundTasks, HTTPException
from app.api.deps import ComboRepoDep, DbSession, StrategyRepoDep from app.api.deps import ComboRepoDep, DbSession, StrategyRepoDep
@@ -108,7 +110,7 @@ def _submit_combo_job(
_ensure_strategies_exist(combo, strategy_repo) _ensure_strategies_exist(combo, strategy_repo)
job = JobRecord( job = JobRecord(
id=new_id("JOB"), kind="combo", status=JobStatus.QUEUED, id=new_id("JOB"), kind="combo", status=JobStatus.QUEUED,
spec_json=combo.model_dump_json(), spec_json=combo.model_dump_json(), created_at=datetime.now(),
) )
SqlAlchemyJobRepository(session).create(job) SqlAlchemyJobRepository(session).create(job)
session.commit() session.commit()
+177
View File
@@ -0,0 +1,177 @@
"""字段库 API:/api/condition-fields(2026-10)。
GET /api/condition-fields 字段库列表(首次读取自动 seed 内置字段)
GET /api/condition-fields/available 引擎支持但尚未进库的字段(「新增」可选项)
POST /api/condition-fields 新增自定义字段(必须指向引擎真能算的字段)
PUT /api/condition-fields/{name} 改中文名/含义/分组/单位/启用状态
DELETE /api/condition-fields/{name} 删除自定义字段(内置字段只能停用)
设计要点(AGENT.md §24 不假装支持):字段能不能算由 `quant.condition_fields` 对着引擎域
判定;库里登记不出来的字段会被拒绝(422),否则用户会建出「永远选不出股票」的空策略。
"""
from __future__ import annotations
from fastapi import APIRouter, HTTPException
from pydantic import BaseModel, ConfigDict, Field
from app.api.deps import ConditionFieldRepoDep, DbSession
from app.application.services import condition_field_catalog as svc
from app.domain.entities.condition_field import ConditionField
from app.quant.condition_fields import FieldDef, get_field, unit_options
router = APIRouter(prefix="/condition-fields", tags=["condition-fields"])
class UnitOption(BaseModel):
"""可选**界面单位** + 它到**基准单位**的换算系数(提交前 ×factor,回显时 ÷factor)。
``factor`` 由注册表给出(如 总市值:万元=1、亿元=10000),前端据此换算。
"""
model_config = ConfigDict(extra="forbid")
unit: str
factor: float
class ConditionFieldOut(ConditionField):
"""字段库响应:目录字段 + **从注册表派生**的单位信息。
``base_unit`` / ``units`` 不落库:它们是引擎口径(注册表)的投影,若存进表里就会
和代码漂移。``unit`` 才是库里存的那一项(当前界面单位,必须落在 ``units`` 里)。
"""
base_unit: str = ""
units: list[UnitOption] = Field(default_factory=list)
@classmethod
def from_item(cls, item: ConditionField) -> ConditionFieldOut:
d = get_field(item.name)
return cls(
**item.model_dump(exclude={"ops"}), # ops 是 computed_field,不参与构造
base_unit=(d.unit if d else item.unit),
units=[UnitOption(unit=u, factor=f) for u, f in unit_options(item.name)],
)
class FieldOption(BaseModel):
"""「可新增字段」的建议项:来自代码注册表,附带默认中文名与含义。"""
model_config = ConfigDict(extra="forbid")
name: str
label: str
description: str
kind: str
group_name: str
unit: str = ""
units: list[UnitOption] = Field(default_factory=list)
@classmethod
def from_def(cls, d: FieldDef) -> FieldOption:
return cls(
name=d.name,
label=d.label,
description=d.description,
kind=d.kind,
group_name=d.group_name,
unit=d.unit,
units=[UnitOption(unit=u, factor=f) for u, f in d.unit_options],
)
class ConditionFieldCreate(BaseModel):
"""新增请求。kind 不收:类型是引擎事实,由服务端按注册表填。"""
model_config = ConfigDict(extra="forbid")
name: str = Field(min_length=1, max_length=64)
label: str = Field(default="", max_length=64)
description: str = Field(default="", max_length=500)
group_name: str = Field(default="", max_length=32)
unit: str = Field(default="", max_length=16)
enabled: bool = True
class ConditionFieldUpdate(BaseModel):
"""编辑请求。name/kind/source 有意不可改(name 是引擎字段名,改了就换字段了)。"""
model_config = ConfigDict(extra="forbid")
label: str | None = Field(default=None, max_length=64)
description: str | None = Field(default=None, max_length=500)
group_name: str | None = Field(default=None, max_length=32)
unit: str | None = Field(default=None, max_length=16)
enabled: bool | None = None
@router.get("", response_model=list[ConditionFieldOut], summary="字段库列表")
def list_condition_fields(
repo: ConditionFieldRepoDep, session: DbSession, include_disabled: bool = True
) -> list[ConditionFieldOut]:
"""读字段库;缺失的内置字段当场补齐(幂等,稳态零写入)。
`include_disabled=false` 供条件编辑器使用(只列启用项);字段库管理页用默认值
(列出全部,含停用项,否则用户没法把停用的字段再打开)。
"""
return [ConditionFieldOut.from_item(f) for f in svc.list_fields(repo, session, include_disabled=include_disabled)]
@router.get("/available", response_model=list[FieldOption], summary="可新增的字段")
def list_available_fields(repo: ConditionFieldRepoDep, session: DbSession) -> list[FieldOption]:
return [FieldOption.from_def(d) for d in svc.list_available(repo, session)]
@router.post("", response_model=ConditionFieldOut, summary="新增自定义字段")
def create_condition_field(
body: ConditionFieldCreate, repo: ConditionFieldRepoDep, session: DbSession
) -> ConditionFieldOut:
try:
saved = svc.create_field(
repo,
session,
name=body.name,
label=body.label,
description=body.description,
group_name=body.group_name,
unit=body.unit,
enabled=body.enabled,
)
except ValueError as exc:
# 422:语义是「引擎算不出来 / 已在库里 / 单位不在可选范围」,属于请求内容不可处理
raise HTTPException(status_code=422, detail=str(exc)) from exc
return ConditionFieldOut.from_item(saved)
@router.put("/{name}", response_model=ConditionFieldOut, summary="编辑字段(中文名/含义/单位/启用)")
def update_condition_field(
name: str, body: ConditionFieldUpdate, repo: ConditionFieldRepoDep, session: DbSession
) -> ConditionFieldOut:
try:
saved = svc.update_field(
repo,
session,
name,
label=body.label,
description=body.description,
group_name=body.group_name,
unit=body.unit,
enabled=body.enabled,
)
except KeyError as exc:
raise HTTPException(status_code=404, detail=f"字段 {name} 不存在") from exc
except ValueError as exc:
raise HTTPException(status_code=422, detail=str(exc)) from exc
return ConditionFieldOut.from_item(saved)
@router.delete("/{name}", summary="删除自定义字段(内置只能停用)")
def delete_condition_field(name: str, repo: ConditionFieldRepoDep, session: DbSession) -> dict:
try:
svc.delete_field(repo, session, name)
except KeyError as exc:
raise HTTPException(status_code=404, detail=f"字段 {name} 不存在") from exc
except ValueError as exc:
raise HTTPException(status_code=400, detail=str(exc)) from exc
return {"deleted": name}
+9
View File
@@ -16,6 +16,7 @@ from app.application.services.selection_service import SelectionService
from app.application.services.signal_service import SignalService from app.application.services.signal_service import SignalService
from app.domain.repositories.combo import ComboRepository, GlobalConfigRepository from app.domain.repositories.combo import ComboRepository, GlobalConfigRepository
from app.domain.repositories.composite import CompositeRepository from app.domain.repositories.composite import CompositeRepository
from app.domain.repositories.condition_field import ConditionFieldRepository
from app.domain.repositories.factor import FactorRepository from app.domain.repositories.factor import FactorRepository
from app.domain.repositories.index import IndexConstituentRepository from app.domain.repositories.index import IndexConstituentRepository
from app.domain.repositories.jobs import ExperimentRepository, JobRepository from app.domain.repositories.jobs import ExperimentRepository, JobRepository
@@ -37,6 +38,9 @@ from app.infrastructure.persistence.sqlalchemy.repositories.combo_impl import (
from app.infrastructure.persistence.sqlalchemy.repositories.composite_impl import ( from app.infrastructure.persistence.sqlalchemy.repositories.composite_impl import (
SqlAlchemyCompositeRepository, SqlAlchemyCompositeRepository,
) )
from app.infrastructure.persistence.sqlalchemy.repositories.condition_field_impl import (
SqlAlchemyConditionFieldRepository,
)
from app.infrastructure.persistence.sqlalchemy.repositories.factor_impl import ( from app.infrastructure.persistence.sqlalchemy.repositories.factor_impl import (
SqlAlchemyFactorRepository, SqlAlchemyFactorRepository,
) )
@@ -174,6 +178,10 @@ def _factor_repo_factory(session: DbSession) -> FactorRepository:
return SqlAlchemyFactorRepository(session) return SqlAlchemyFactorRepository(session)
def _condition_field_repo_factory(session: DbSession) -> ConditionFieldRepository:
return SqlAlchemyConditionFieldRepository(session)
def _composite_repo_factory(session: DbSession) -> CompositeRepository: def _composite_repo_factory(session: DbSession) -> CompositeRepository:
return SqlAlchemyCompositeRepository(session) return SqlAlchemyCompositeRepository(session)
@@ -201,6 +209,7 @@ ResearchServiceDep = Annotated[ResearchService, Depends(_service_factory)]
SelectionServiceDep = Annotated[SelectionService, Depends(_selection_service_factory)] SelectionServiceDep = Annotated[SelectionService, Depends(_selection_service_factory)]
SelectionRepoDep = Annotated[SelectionRepository, Depends(_selection_repo_factory)] SelectionRepoDep = Annotated[SelectionRepository, Depends(_selection_repo_factory)]
FactorRepoDep = Annotated[FactorRepository, Depends(_factor_repo_factory)] FactorRepoDep = Annotated[FactorRepository, Depends(_factor_repo_factory)]
ConditionFieldRepoDep = Annotated[ConditionFieldRepository, Depends(_condition_field_repo_factory)]
CompositeRepoDep = Annotated[CompositeRepository, Depends(_composite_repo_factory)] CompositeRepoDep = Annotated[CompositeRepository, Depends(_composite_repo_factory)]
SignalRepoDep = Annotated[SignalRepository, Depends(_signal_repo_factory)] SignalRepoDep = Annotated[SignalRepository, Depends(_signal_repo_factory)]
SignalServiceDep = Annotated[SignalService, Depends(_signal_service_factory)] SignalServiceDep = Annotated[SignalService, Depends(_signal_service_factory)]
+144 -21
View File
@@ -1,34 +1,157 @@
"""因子目录 API:/api/factors(M7.1 起读 DB factor_definition)。 """因子目录 API:/api/factors(M7.1 起读 DB,2026-10 支持参数化实例)。
目录为空时自动从代码注册表 seed(幂等);随后可登记自定义因子元数据。 ## 目录的三条规则(详见 application/services/factor_catalog.py)
响应为 FactorDefinition 实体(含 requires 列表等)。
① 注册表有、库里没有 → 补齐(历史 bug:表非空后新因子永远进不了目录)。
② 能算出来的行,口径字段按代码改回 —— 目录不允许与引擎口径不一致(手改会被纠正)。
③ 库里多出来的行保留(只补不删),标 `resolvable` 告知是否算得出来。
## 参数化(本文件新增的部分)
- **暴露**:每个因子返回 `template` / `params` / `param_specs`(可编辑参数与允许范围)/
`label`(中文名含参数)/ `source` / `resolvable` / `enabled`,界面据此渲染参数表与表单。
- **新建**:`POST /api/factors` 传 `{template, params}` → 生成参数化实例
`momentum(window=90,direction=higher_is_better)`。参数写在名字里,所以它**冻结**了自己的
口径:以后无论谁再改参数,既有策略/归档按各自名字里的参数计算,不会变义。
- **停用**:`PATCH /api/factors` 传 `{name, enabled}`。名字里有括号/等号/逗号,放路径里
会被各种代理折腾,所以放在 body 里。
- 越界参数、未知模板、重复的参数组合 → **422**,并在 detail 里说明允许范围。
""" """
from __future__ import annotations from __future__ import annotations
from fastapi import APIRouter from typing import Any
from fastapi import APIRouter, HTTPException
from pydantic import BaseModel, Field
from app.api.deps import DbSession, FactorRepoDep from app.api.deps import DbSession, FactorRepoDep
from app.application.services.factor_catalog import seed_registry_factors from app.application.services.factor_catalog import (
from app.domain.entities.factor import FactorDefinition create_parameterized_factor,
from app.quant.factors import list_factors set_factor_enabled,
sync_registry_factors,
)
from app.domain.entities.factor import FactorDefinition, FactorParam
from app.quant.factors import FactorError, list_templates, resolve_factor
router = APIRouter(prefix="/factors", tags=["factors"]) router = APIRouter(prefix="/factors", tags=["factors"])
@router.get("", summary="因子目录(含元数据,来自 factor_definition 表)") class FactorTemplateOut(BaseModel):
def list_factor_catalog(factor_repo: FactorRepoDep, session: DbSession) -> list[FactorDefinition]: """模板(算法家族):可编辑参数 + 默认值 + 口径说明,供「新建参数化因子」表单。"""
"""读目录;**缺失的注册表因子当场补齐**(幂等),稳态下零写入。
历史 bug:原先只在「表为空」时 seed,于是表非空后**代码里新增的因子永远进不了目录**—— name: str
实测表内 9 条、注册表 11 条,`dividend_yield` / `dividend_yield_ttm` 长期缺失, label: str
前端因子下拉选不到「股息率」、归档页也查不到它的方向与含义(违反 §7 不静默)。 description: str
现在按「注册表有、库里没有」的差集触发 upsert:既保证目录与可计算因子一致, formula: str
又不会覆盖用户登记的自定义因子元数据(只补不删)。 brief: str
requires: list[str] = Field(default_factory=list)
frequency: str = "daily"
direction_default: str = "higher_is_better"
param_specs: list[FactorParam] = Field(default_factory=list)
defaults: dict[str, Any] = Field(default_factory=dict)
instances: list[str] = Field(default_factory=list) # 该模板已有的内置实例名
class FactorCreate(BaseModel):
"""新建参数化因子:模板 + 参数(缺省项取模板默认值)。"""
template: str
params: dict[str, Any] = Field(default_factory=dict)
class FactorPatch(BaseModel):
"""开关因子(name 放 body:名字里有括号/等号/逗号,不适合放路径)。"""
name: str
enabled: bool
def _template_out(tpl) -> FactorTemplateOut:
return FactorTemplateOut(
name=tpl.name,
label=tpl.label,
description=tpl.description,
formula=tpl.formula,
brief=tpl.brief,
requires=list(tpl.requires),
frequency=tpl.frequency,
direction_default=tpl.direction_default,
param_specs=[FactorParam.from_spec(s) for s in tpl.specs()],
defaults=tpl.defaults(),
instances=[name for name, _params in tpl.instances],
)
def _enrich(row: FactorDefinition) -> FactorDefinition:
"""把库里的行补成「引擎口径的投影」:能算出来的行一律以引擎为准。
- 算得出来 → description/formula/brief/frequency/lookback/direction/requires/
template/params/param_specs/label/source 全部取自引擎(目录永不撒谎);
`enabled` 仍是库里的人配值。
- 算不出来(历史手工登记行)→ 原样返回并标 `resolvable=False`,界面显示为不可用。
""" """
existing = factor_repo.list() try:
missing = {d.name for d in list_factors()} - {f.name for f in existing} defn, _fn = resolve_factor(row.name)
if not existing or missing: except FactorError:
seed_registry_factors(factor_repo, session) return row.model_copy(update={"resolvable": False, "label": row.name})
existing = factor_repo.list() return FactorDefinition.from_factor_def(defn, enabled=row.enabled).model_copy(
return existing update={"created_at": row.created_at, "version": row.version}
)
@router.get("", summary="因子目录(注册表投影 + 参数化实例,含可编辑参数)")
def list_factor_catalog(factor_repo: FactorRepoDep, session: DbSession) -> list[FactorDefinition]:
"""读目录(先做幂等同步,稳态零写入),每行补上参数化视图。
前端据此渲染:因子名 / 中文名含参数 / 窗口 · 方向等真实参数 / 是否可当过滤条件。
"""
sync_registry_factors(factor_repo, session)
return [_enrich(row) for row in factor_repo.list()]
@router.get("/templates", summary="因子模板:可编辑参数与允许范围")
def list_factor_templates() -> list[FactorTemplateOut]:
"""全部模板(动量 / 波动率 / 量比 / 乖离 / 反转 / 接近新高 / 股息率…)。
每个模板给出 `param_specs`(参数名、类型、允许范围/枚举、默认值、说明)——
界面据此渲染受控表单:**参数只在给定范围内选/填**,越界在 API 层就被拒。
"""
return [_template_out(tpl) for tpl in list_templates()]
@router.post("", status_code=201, summary="新建参数化因子(模板 + 参数)")
def create_factor(
payload: FactorCreate,
factor_repo: FactorRepoDep,
session: DbSession,
) -> FactorDefinition:
"""从模板派生一个新的参数化因子实例;参数写进名字,因此口径被冻结。
重复的参数组合不会重复创建(409 语义由 422 承载并给出已有名字,前端直接提示即可)。
"""
try:
row = create_parameterized_factor(
factor_repo, session, template=payload.template, params=payload.params
)
except FactorError as exc:
raise HTTPException(status_code=422, detail=str(exc)) from exc
except ValueError as exc:
raise HTTPException(status_code=422, detail=str(exc)) from exc
return _enrich(row)
@router.patch("", summary="启用/停用因子(只影响能否被选中)")
def patch_factor(
payload: FactorPatch,
factor_repo: FactorRepoDep,
session: DbSession,
) -> FactorDefinition:
"""停用只把因子从下拉里拿掉:既有策略/归档仍按名字解析(历史不变义)。"""
try:
row = set_factor_enabled(factor_repo, session, name=payload.name, enabled=payload.enabled)
except LookupError as exc:
raise HTTPException(status_code=404, detail=str(exc)) from exc
except ValueError as exc:
raise HTTPException(status_code=422, detail=str(exc)) from exc
return _enrich(row)
+3 -1
View File
@@ -42,7 +42,9 @@ def _decode_result(kind: str, result_json: str | None):
return None return None
if kind == "selection": if kind == "selection":
return SelectionResult.model_validate_json(result_json) return SelectionResult.model_validate_json(result_json)
model = BacktestResult if kind == "backtest" else FactorTestReport # combo 回测的归档 kind 记为 "backtest",但 Job.kind 仍是 "combo" ——
# 其结果同样是 BacktestResult,按 backtest 解码(否则会被当成因子测试而校验失败)。
model = BacktestResult if kind in ("backtest", "combo") else FactorTestReport
return model.model_validate_json(result_json) return model.model_validate_json(result_json)
+2
View File
@@ -13,6 +13,7 @@ from app.api import (
charts, charts,
combos, combos,
composites, composites,
condition_fields,
config, config,
experiments, experiments,
factors, factors,
@@ -30,6 +31,7 @@ api_router = APIRouter()
api_router.include_router(health.router) api_router.include_router(health.router)
api_router.include_router(stocks.router) api_router.include_router(stocks.router)
api_router.include_router(factors.router) api_router.include_router(factors.router)
api_router.include_router(condition_fields.router)
api_router.include_router(composites.router) api_router.include_router(composites.router)
api_router.include_router(research.router) api_router.include_router(research.router)
api_router.include_router(selections.router) api_router.include_router(selections.router)
+12 -4
View File
@@ -4,13 +4,20 @@
回测改由「回测组合」(/api/combos)驱动,故旧的 /{id}/expand(→ResearchSpec)已移除。 回测改由「回测组合」(/api/combos)驱动,故旧的 /{id}/expand(→ResearchSpec)已移除。
POST /api/strategies 保存策略(name 唯一;description 为空时自动补全) POST /api/strategies 保存策略(name 唯一;description 为空时自动补全)
POST /api/strategies/describe body: ResearchSpec → StrategyDoc(未保存的策略也能预览) POST /api/strategies/describe body: ResearchSpec → StrategyDoc(见下方说明,非策略库路径)
GET /api/strategies 列表 GET /api/strategies 列表
GET /api/strategies/{id} GET /api/strategies/{id}
PUT /api/strategies/{id} 原地更新(不新建、不刷新 created_at) PUT /api/strategies/{id} 原地更新(不新建、不刷新 created_at)
DELETE /api/strategies/{id} DELETE /api/strategies/{id}
GET /api/strategies/{id}/describe → StrategyDoc GET /api/strategies/{id}/describe → StrategyDoc
关于两个 describe 端点(不是历史遗留,各有明确用途,别合并):
- `GET /{id}/describe` → 入参是已保存的 **SelectionStrategy**(策略库「看说明/公式」用);
- `POST /describe` → 入参是 **ResearchSpec**,**给归档页**用:`/experiments/{id}` 要按当时
归档的旧 ResearchSpec 快照(单策略回测路径,含 selection/rebalance/costs)复述口径。
该路径仍然存在(`POST /api/backtests` 是底层 escape hatch),所以这里必须继续支持。
回测页本身已不再调用它(组合回测走 ComboRunSpec + 归档页组合卡片)。
路由顺序注意:`/describe` 这类**字面量路径**一律声明在 `/{strategy_id}` 之前 —— 路由顺序注意:`/describe` 这类**字面量路径**一律声明在 `/{strategy_id}` 之前 ——
否则会被路径参数吞掉(AGENT.md §17 的既有教训,/api/stocks/names 同源问题)。 否则会被路径参数吞掉(AGENT.md §17 的既有教训,/api/stocks/names 同源问题)。
""" """
@@ -71,10 +78,11 @@ def create_strategy(
@router.post("/describe", response_model=StrategyDoc, summary="按 ResearchSpec 生成策略说明与公式") @router.post("/describe", response_model=StrategyDoc, summary="按 ResearchSpec 生成策略说明与公式")
def describe_research_spec(spec: ResearchSpec) -> StrategyDoc: def describe_research_spec(spec: ResearchSpec) -> StrategyDoc:
"""回测页参数即时预览用:**未保存的策略**(只有 spec)也能生成说明/公式。 """按 **ResearchSpec** 生成说明/公式 —— 服务于归档页,不是策略库路径。
纯函数实现(app.quant.strategy_doc),无 IO/DB,因此不会因保存状态而失败。 调用方是 `/experiments/{id}`:它按归档里冻结的 ResearchSpec 快照(单策略回测)复述
路径与 `POST /api/strategies` 不冲突(字面量 /describe 优先于路径参数声明)。 「选股条件 + 交易执行依据」。纯函数实现(app.quant.strategy_doc),无 IO/DB,因此
不依赖任何保存状态,历史归档随时可复述。策略库自身的说明走 `GET /{id}/describe`。
""" """
return describe_strategy(spec) return describe_strategy(spec)
@@ -78,6 +78,15 @@ class ComboService:
daily=daily, daily=daily,
eligibility_fns=eligibility_fns, eligibility_fns=eligibility_fns,
) )
# 把价格口径写进 config_snapshot(与 ResearchService._annotate_price_basis 同口径),
# 否则归档页/结果头读不到 adjust_mode,会误显示「不复权」(组合实际用的是公共配置的复权)。
# 注意:不能覆盖整个 config_snapshot —— 引擎已把 ComboRunSpec 固化在里面(可复现依据)。
mode = config.price_adjustment
result.config_snapshot["price_basis"] = {
"adjust_mode": mode,
"price_basis": "adjust_factor" if mode != "none" else "raw_close",
"execution_price_basis": "close_adj" if mode != "none" else "close_raw",
}
_stage(on_stage, "analysis") _stage(on_stage, "analysis")
return _fill_names(result, self._last_stocks) return _fill_names(result, self._last_stocks)
@@ -0,0 +1,198 @@
"""字段库用例:目录 seed + 校验 + 增删改(2026-10)。
与 factor_catalog 同一套规矩:**DB 是目录契约源,代码注册表是可用性的唯一事实来源**。
读取时把「注册表有、库里没有」的内置字段补进去(只补不删、不覆盖用户改过的文案)。
单位的两层含义(2026-10 补)见 `quant/condition_fields.py`:``unit`` 存的是**界面单位**
(输入/显示用,可从注册表给的阶梯里选),引擎始终按**基准单位**存储与比较,
换算是提交/回显时按系数做的 —— 所以改单位不会让任何历史策略变义。
"""
from __future__ import annotations
from app.domain.entities.condition_field import ConditionField
from app.domain.repositories.condition_field import ConditionFieldRepository
from app.quant.condition_fields import (
GROUP_ORDER,
FieldDef,
available_fields,
curated_fields,
get_field,
reason_unsupported,
unit_allowed,
unit_options,
)
_SORT_BASE = 100
def _check_unit(name: str, unit: str) -> str:
"""界面单位必须落在注册表给的阶梯里(不允许自由文本 —— 见 AGENT §24)。
空串 = 保持基准单位。给出可用单位清单,用户不用猜为什么被拒。
"""
d = get_field(name)
base = d.unit if d else ""
if not unit.strip():
return base
if not unit_allowed(name, unit.strip()):
allowed = "、".join(u for u, _ in unit_options(name)) or "(无可用单位)"
raise ValueError(
f"字段「{name}」不支持单位「{unit.strip()}」:可选 {allowed}。"
"单位只能从这些里选,因为换算是按固定系数做的(引擎按基准单位比较)。"
)
return unit.strip()
def _sort_order(d: FieldDef) -> int:
"""同分组内保持注册表顺序(分组顺序 × 1000 + 组内序号)。"""
try:
g = GROUP_ORDER.index(d.group_name)
except ValueError:
g = len(GROUP_ORDER)
return g * 1000 + _SORT_BASE
def _from_def(d: FieldDef) -> ConditionField:
return ConditionField(
name=d.name,
label=d.label,
description=d.description,
kind=d.kind,
group_name=d.group_name,
unit=d.unit,
source="builtin",
enabled=True,
sort_order=_sort_order(d),
)
def sync_builtin_fields(repo: ConditionFieldRepository, session) -> int:
"""补齐缺失的**默认内置字段**(幂等;已存在的一律不动,用户改过的文案得以保留)。
只 seed `curated=True` 的那批:`curated=False` 的字段是「引擎支持但默认不进库」的
选项,留给用户在字段库里按需新增(见 `list_available`)—— 否则「新增字段」永远
无字段可选。
"""
existing = {f.name for f in repo.list()}
missing = [_from_def(d) for d in curated_fields() if d.name not in existing]
added = repo.insert_missing(missing)
if added:
session.commit()
return added
def list_fields(repo: ConditionFieldRepository, session, include_disabled: bool = True):
"""字段库列表(首次读取自动 seed)。按分组/注册表顺序排序。"""
items = repo.list()
if not items:
sync_builtin_fields(repo, session)
items = repo.list()
# 代码里新注册的默认字段也要补上:按差集触发,稳态零写入
if {d.name for d in curated_fields()} - {f.name for f in items}:
sync_builtin_fields(repo, session)
items = repo.list()
if not include_disabled:
items = [i for i in items if i.enabled]
return items
def list_available(repo: ConditionFieldRepository, session):
"""引擎支持但尚未进目录的字段(「新增字段」的可选项)。"""
return available_fields({f.name for f in repo.list()})
def create_field(
repo: ConditionFieldRepository,
session,
*,
name: str,
label: str = "",
description: str = "",
group_name: str = "",
unit: str = "",
enabled: bool = True,
) -> ConditionField:
"""新增自定义字段。
校验顺序有意如此:先看引擎能不能算(不能算就 422,绝不放行)→ 再看是否已存在。
这样用户拿到的是「这个字段引擎算不出来」而不是含糊的「已存在」。
"""
key = name.strip()
reason = reason_unsupported(key)
if reason:
raise ValueError(reason)
if repo.get(key) is not None:
raise ValueError(f"字段「{key}」已在字段库中:直接编辑它,或给它改个中文名/含义即可")
d = get_field(key)
assert d is not None # reason_unsupported 为空 ⇒ 注册表必有此字段
chosen = _check_unit(key, unit)
item = ConditionField(
name=key,
label=(label.strip() or d.label),
description=(description.strip() or d.description),
kind=d.kind, # 类型来自引擎,不接受调用方声明
group_name=(group_name.strip() or d.group_name),
unit=chosen,
source="custom",
enabled=enabled,
sort_order=_sort_order(d),
)
saved = repo.save(item)
session.commit()
return saved
def update_field(
repo: ConditionFieldRepository,
session,
name: str,
*,
label: str | None = None,
description: str | None = None,
group_name: str | None = None,
unit: str | None = None,
enabled: bool | None = None,
) -> ConditionField:
"""编辑字段文案/分组/**界面单位**/启用状态。
``name`` / ``kind`` / ``source`` 不可改:前者是引擎字段名,后两者是引擎事实与
条目来历 —— 允许改会让「字段库」与引擎脱节(§24 不做假支持)。
``unit`` 可改但**只能在注册表给的阶梯里选**(如 万元 ⇄ 亿元):它是界面单位,
提交/回显按固定系数换算,引擎始终用基准单位 —— 所以改它不会让历史策略变义。
"""
item = repo.get(name)
if item is None:
raise KeyError(name)
patch = item.model_copy(
update={
k: v
for k, v in {
"label": label,
"description": description,
"group_name": group_name,
"unit": _check_unit(name, unit) if unit is not None else None,
"enabled": enabled,
}.items()
if v is not None
}
)
if not patch.label.strip():
raise ValueError("中文名不能为空(下拉里要显示它)")
saved = repo.save(patch)
session.commit()
return saved
def delete_field(repo: ConditionFieldRepository, session, name: str) -> None:
"""删除自定义字段。内置字段不允许删除 —— 删了下次 seed 又会补回来,只会让人困惑。"""
item = repo.get(name)
if item is None:
raise KeyError(name)
if item.source == "builtin":
raise ValueError(
f"「{name}」是内置字段,不能删除(删掉下次读取也会自动补回)。"
"如果不想在条件里看到它,请改为「停用」。"
)
repo.delete(name)
session.commit()
@@ -1,21 +1,147 @@
"""因子目录用例:把代码注册表因子 seed 进 DB(M7.1)。 """因子目录用例:注册表投影 + 参数化实例的新建/停用(M7.1 起,2026-10 参数化)。
DB 为目录契约源;本服务在 /api/factors 首次读取为空时自动 seed(幂等), ## 目录与代码的分工(这是本模块的核心规矩)
后续代码新增因子也通过同一入口同步,保持目录与可计算因子一致。
- **能不能算** = 代码注册表(`quant/factors.py`)唯一决定。库里的行只要能解析出
`(模板, 参数)` 就算得出来;解析不出来的行(历史手工登记)保留但在目录里标
`resolvable=False`,引用时抛 `FactorError`(不假装支持)。
- **有哪些因子** = 目录(DB)。内置实例由代码投影进来;**参数化实例**(如
`momentum(window=90,direction=higher_is_better)`)由人在目录里创建 —— 参数写在名字里,
所以「同一个模板的多个参数版本」天然并存,且任何一个都冻结了自己的口径。
- **口径文案**(description/formula/brief/frequency/lookback/direction/requires)永远按
代码收敛:只要这行算得出来,它的文案就是引擎的文案。手改会被改回 —— 因为
「文档写一套、代码跑另一套」是本仓库明令禁止的;要改口径就改代码。
- **唯一人配的字段**是 `enabled`(是否出现在因子下拉里):只对参数化实例生效;
内置实例的开关同样由代码收敛(恒 True)。停用不影响已引用它的策略/归档解析 ——
历史口径不能被一个开关改义。
稳态(目录 == 代码投影)零写入,不会每次 GET 都刷库。
""" """
from __future__ import annotations from __future__ import annotations
from collections.abc import Mapping
from typing import Any
from app.domain.entities.factor import FactorDefinition from app.domain.entities.factor import FactorDefinition
from app.domain.repositories.factor import FactorRepository from app.domain.repositories.factor import FactorRepository
from app.quant.factors import list_factors from app.quant.factors import (
FactorError,
build_factor_def,
canonical_key,
get_template,
is_resolvable,
list_factors,
resolve_factor,
)
# 「引擎口径」字段:一旦与代码不一致,目录就在撒谎,必须按代码改回。
# enabled 也在其中 —— 但它只对**注册表实例**收敛(见 sync_registry_factors)。
REGISTRY_FIELDS = (
"description",
"formula",
"brief",
"frequency",
"lookback",
"direction",
"requires",
"enabled",
)
# 参数化实例同样要收敛的字段(不含 enabled:开关是人配的)。
_ENGINE_FIELDS = tuple(f for f in REGISTRY_FIELDS if f != "enabled")
def seed_registry_factors(repo: FactorRepository, session) -> int: def _registry_names() -> set[str]:
"""把 quant/factors 注册表的元数据 upsert 进 factor_definition(幂等)。""" return {d.name for d in list_factors()}
defs = [FactorDefinition.from_registry_def(d) for d in list_factors()]
if not defs:
def _wanted_rows(current: list[FactorDefinition]) -> list[FactorDefinition]:
"""应然状态:内置实例按代码;能解析出来的参数化实例口径也按代码(开关保留)。"""
registry = _registry_names()
wanted = [FactorDefinition.from_registry_def(d) for d in list_factors()]
for row in current:
if row.name in registry:
continue # 内置实例已由上面的投影覆盖
try:
defn, _fn = resolve_factor(row.name)
except FactorError:
continue # 算不出来的手登记行:不碰(只有人写的备注,没有可收敛的引擎口径)
want = FactorDefinition.from_factor_def(defn, enabled=row.enabled)
want = want.model_copy(update={"created_at": row.created_at, "version": row.version})
wanted.append(want)
return wanted
def _differs(current: FactorDefinition | None, want: FactorDefinition, fields) -> bool:
if current is None:
return True
return any(getattr(current, name) != getattr(want, name) for name in fields)
def sync_registry_factors(repo: FactorRepository, session) -> int:
"""把代码投影进目录(幂等);返回本次写入的行数,稳态为 0。"""
current = repo.list()
by_name = {f.name: f for f in current}
registry = _registry_names()
stale: list[FactorDefinition] = []
for want in _wanted_rows(current):
row = by_name.get(want.name)
fields = REGISTRY_FIELDS if want.name in registry else _ENGINE_FIELDS
if row is None or _differs(row, want, fields):
stale.append(want)
if not stale:
return 0 return 0
n = repo.upsert_many(defs) n = repo.upsert_many(stale)
session.commit() session.commit()
return n return n
def create_parameterized_factor(
repo: FactorRepository,
session,
*,
template: str,
params: Mapping[str, Any] | None = None,
) -> FactorDefinition:
"""从模板 + 参数创建一个**新的参数化因子实例**(参数写在名字里,冻结口径)。
参数缺省项取模板默认值;越界/未知参数/重复的参数组合一律报错(不静默纠正)。
"""
tpl = get_template(template) # 未知模板 → FactorError
key = canonical_key(tpl, params or {}) # 越界/未知参数 → FactorError
if repo.get(key) is not None:
raise ValueError(f"该参数组合的因子已存在:{key}(参数相同不会重复创建)")
defn = build_factor_def(tpl, params or {}, name=key, source="custom")
entity = FactorDefinition.from_factor_def(defn, enabled=True)
repo.upsert_many([entity])
session.commit()
return entity
def set_factor_enabled(
repo: FactorRepository,
session,
*,
name: str,
enabled: bool,
) -> FactorDefinition:
"""启用/停用目录里的因子(只影响「能不能被选中」,不影响历史解析)。"""
if name in _registry_names():
raise ValueError(
f"「{name}」是代码注册表里的内置因子,开关由代码决定,不能在目录里停用;"
"如果要一个不同参数的版本,请从模板新建参数化因子。"
)
row = repo.get(name)
if row is None:
raise LookupError(f"因子 {name} 不在目录里")
if not is_resolvable(name):
raise ValueError(f"因子 {name} 引擎算不出来(未注册模板/参数非法),不能启用或停用")
updated = row.model_copy(update={"enabled": enabled})
repo.upsert_many([updated])
session.commit()
return updated
# 兼容旧名(语义相同:把注册表同步进目录)。
seed_registry_factors = sync_registry_factors
+11 -1
View File
@@ -18,7 +18,7 @@ from __future__ import annotations
from datetime import date, datetime from datetime import date, datetime
from pydantic import BaseModel, Field, field_validator, model_validator from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator
from app.domain.entities.research import CostSpec from app.domain.entities.research import CostSpec
@@ -31,8 +31,13 @@ class GlobalConfig(BaseModel):
为什么把复权口径也放这里:一次回测只能有一个复权口径(同一份行情不能既前复权 为什么把复权口径也放这里:一次回测只能有一个复权口径(同一份行情不能既前复权
又后复权),而多个选股策略可能想混用 —— 与其让它们在组合里打架,不如统一为 又后复权),而多个选股策略可能想混用 —— 与其让它们在组合里打架,不如统一为
全局口径,高股息默认 hfq。若将来确需按组合区分,再加字段即可(向前兼容)。 全局口径,高股息默认 hfq。若将来确需按组合区分,再加字段即可(向前兼容)。
`extra="forbid"`:PUT /api/config 若带未知字段(拼错键名、旧版遗留键)直接报错,
避免「以为改了某项、其实被静默忽略」。
""" """
model_config = ConfigDict(extra="forbid")
id: str = Field(default="default", description="单例主键,恒为 default") id: str = Field(default="default", description="单例主键,恒为 default")
commission_rate: float = Field(default=0.0003, ge=0, le=0.01, description="佣金率(如 0.0003 = 万三)") commission_rate: float = Field(default=0.0003, ge=0, le=0.01, description="佣金率(如 0.0003 = 万三)")
stamp_tax_rate: float = Field(default=0.0005, ge=0, le=0.01, description="印花税率(仅卖出)") stamp_tax_rate: float = Field(default=0.0005, ge=0, le=0.01, description="印花税率(仅卖出)")
@@ -72,8 +77,13 @@ class BacktestCombo(BaseModel):
- `hold_max_days` = Tmax:个股**最多**持有天数 —— 超过即强制了结(None = 不限)。 - `hold_max_days` = Tmax:个股**最多**持有天数 —— 超过即强制了结(None = 不限)。
- `rebalance_freq`:多久重新打分排序并调仓一次(日/周/月)。 - `rebalance_freq`:多久重新打分排序并调仓一次(日/周/月)。
⚠️ Tmax 强制卖出**每个交易日**都检查(不只调仓日),否则月频下会远超 Tmax。 ⚠️ Tmax 强制卖出**每个交易日**都检查(不只调仓日),否则月频下会远超 Tmax。
`extra="forbid"`:回测参数写错键名(如 hold_days、capital)时报错而非静默用默认值 ——
静默用默认值会让「我明明设了 30 天」变成「其实没生效」,是本项目明确禁止的降级方式。
""" """
model_config = ConfigDict(extra="forbid")
id: str = "" id: str = ""
name: str = Field(min_length=1, max_length=64) name: str = Field(min_length=1, max_length=64)
description: str = "" description: str = ""
@@ -0,0 +1,54 @@
"""字段库领域实体(2026-10:过滤条件字段目录 DB 化)。
与因子目录(``FactorDefinition``)同一套思路:
- **DB 是字段库的契约源**:中文名、含义、单位、是否启用、自定义条目都入库;
- **引擎是字段可用性的唯一事实来源**:字段能不能算由 ``quant.condition_fields``
对着引擎域校验,登记不出来的字段一律拒绝(防「建出来永远选不出股票」的伪字段)。
可编辑边界(有意为之):
- ``name`` 是引擎字段名,**不可改**(改了就指向另一个字段,等于换字段);
- ``kind`` 由引擎类型决定,**不可改**(字符串字段不能比大小);
- ``label`` / ``description`` / ``group_name`` / ``unit`` / ``enabled`` 可编辑 ——
内置字段也允许改文案(seed「只补不删」,不会覆盖用户的措辞)。
"""
from __future__ import annotations
from datetime import datetime
from pydantic import BaseModel, ConfigDict, Field, computed_field
FIELD_KINDS: tuple[str, ...] = ("num", "str")
FIELD_SOURCES: tuple[str, ...] = ("builtin", "custom")
# 类型 → 可用比较符(单一事实来源)。
# 字符串字段只能等值/集合:引擎 _compare 对字符串的 >/≥/</≤ 一律返回 False,
# 若前端把「行业 > 5」这类选项摆出来,用户点出来的就是永远为假的条件。
OPS_NUM: tuple[str, ...] = ("gt", "gte", "lt", "lte", "eq", "ne")
OPS_STR: tuple[str, ...] = ("eq", "ne", "in", "not_in")
OPS_BY_KIND: dict[str, tuple[str, ...]] = {"num": OPS_NUM, "str": OPS_STR}
class ConditionField(BaseModel):
"""字段库中的一个条件字段。"""
model_config = ConfigDict(extra="forbid")
name: str = Field(min_length=1, max_length=64, description="引擎字段名,如 dv_ratio / static.industry")
label: str = Field(default="", max_length=64, description="中文名(下拉里展示)")
description: str = Field(default="", max_length=500, description="含义 / 口径(含单位)")
kind: str = Field(default="num", pattern="^(num|str)$")
group_name: str = Field(default="行情", max_length=32)
unit: str = Field(default="", max_length=16)
source: str = Field(default="builtin", pattern="^(builtin|custom)$")
enabled: bool = Field(default=True, description="False=从选择器隐藏(内置不可删除,只能停用)")
sort_order: int = 100
created_at: datetime | None = None
updated_at: datetime | None = None
@computed_field # type: ignore[prop-decorator]
@property
def ops(self) -> list[str]:
"""该字段可用的比较符(前端据此收窄下拉,不自己猜)。"""
return list(OPS_BY_KIND.get(self.kind, OPS_NUM))
+58 -6
View File
@@ -1,19 +1,51 @@
"""因子目录领域实体(M7.1:因子元数据 DB 化,v2 §11)。 """因子目录领域实体(M7.1:因子元数据 DB 化,v2 §11;2026-10 参数化)。
DB 是因子目录的契约源:元数据(含自定义因子登记)入库; DB 是因子目录的契约源:**哪些因子存在**(含用户从模板派生的参数化实例)入库;
计算执行仍由代码注册表(quant/factors.py)提供 —— 登记但未注册计算的因子 **能不能算**仍由代码注册表(quant/factors.py)唯一决定 —— 登记但解析不出来的因子
在 score/condition 中引用时仍抛 FactorError(防静默伪因子)。 在 score/condition 里引用时抛 FactorError(不假装支持)。
参数化的读法:参数化实例的名字本身就是身份(`momentum(window=90,direction=...…)`),
所以 template / params / param_specs / label / source / resolvable 都是**由名字解析出来的
投影**,不落库。落库的只有 `enabled`(是否出现在下拉里)—— 这是人做的配置,
不是引擎事实。好处:参数不可能出现「表里一套、键里一套」的分裂。
""" """
from __future__ import annotations from __future__ import annotations
from datetime import datetime from datetime import datetime
from typing import Any
from pydantic import BaseModel, Field from pydantic import BaseModel, Field
class FactorParam(BaseModel):
"""一个可编辑参数的约束(与 quant/factors.ParamSpec 对齐,供界面渲染表单)。"""
name: str
label: str = ""
kind: str = "int" # "int" | "enum"
default: Any = None
minimum: int | None = None
maximum: int | None = None
choices: list[str] = Field(default_factory=list)
note: str = ""
@classmethod
def from_spec(cls, spec) -> FactorParam:
return cls(
name=spec.name,
label=spec.label,
kind=spec.kind,
default=spec.default,
minimum=spec.minimum,
maximum=spec.maximum,
choices=list(spec.choices),
note=spec.note,
)
class FactorDefinition(BaseModel): class FactorDefinition(BaseModel):
name: str = Field(min_length=1, max_length=64) name: str = Field(min_length=1, max_length=128)
description: str = "" description: str = ""
formula: str = "" formula: str = ""
brief: str = "" brief: str = ""
@@ -23,10 +55,23 @@ class FactorDefinition(BaseModel):
requires: list[str] = Field(default_factory=list) requires: list[str] = Field(default_factory=list)
version: str = "1" version: str = "1"
created_at: datetime | None = None created_at: datetime | None = None
# ---- 参数化(落库的只有 enabled;其余由 name 解析投影而来)----
enabled: bool = True
template: str = "" # 模板名,如 "momentum"
params: dict[str, Any] = Field(default_factory=dict) # 冻结的参数取值
param_specs: list[FactorParam] = Field(default_factory=list) # 可编辑参数与约束
label: str = "" # 中文显示名(含参数)
source: str = "builtin" # builtin(代码注册表实例)| custom(目录里的参数化实例)
resolvable: bool = True # False = 登记了但引擎算不出来(历史手工登记行)
@classmethod @classmethod
def from_registry_def(cls, d) -> FactorDefinition: def from_registry_def(cls, d) -> FactorDefinition:
"""由 quant/factors.FactorDef(dataclass)构造目录实体(seed 用)。""" """由 quant/factors.FactorDef(dataclass)构造目录实体(seed 用)。"""
return cls.from_factor_def(d, enabled=True)
@classmethod
def from_factor_def(cls, d, *, enabled: bool = True) -> FactorDefinition:
"""由因子实例(内置或参数化)构造目录实体(seed / 新建用)。"""
return cls( return cls(
name=d.name, name=d.name,
description=d.description, description=d.description,
@@ -36,4 +81,11 @@ class FactorDefinition(BaseModel):
lookback=d.lookback, lookback=d.lookback,
direction=d.direction, direction=d.direction,
requires=list(d.requires), requires=list(d.requires),
) enabled=enabled,
template=d.template,
params=dict(d.params),
param_specs=[FactorParam.from_spec(s) for s in d.param_specs],
label=d.label,
source=d.source,
resolvable=True,
)
+4
View File
@@ -54,6 +54,10 @@ class ConditionSpec(BaseModel):
右操作数取 value(字面量)或 ref(另一字段名),二者二选一。 右操作数取 value(字面量)或 ref(另一字段名),二者二选一。
字段域的事实来源:`quant/condition_fields.py`(字段库注册表,含中文名与口径)。
`/api/condition-fields`(前端下拉)、该注册表与引擎求值共用同一份定义,
避免「前端列一个、引擎算另一个」的漂移。
定义位置说明:本模型被 ResearchSpec(回测)与 SelectionQuery(选股)共用, 定义位置说明:本模型被 ResearchSpec(回测)与 SelectionQuery(选股)共用,
故落在 research.py(被 selection.py 依赖的低层模块),selection.py 再 re-export, 故落在 research.py(被 selection.py 依赖的低层模块),selection.py 再 re-export,
避免循环导入。 避免循环导入。
+15 -4
View File
@@ -12,18 +12,29 @@ from __future__ import annotations
from datetime import datetime from datetime import datetime
from pydantic import BaseModel, Field, model_validator from pydantic import BaseModel, ConfigDict, Field, model_validator
from app.domain.entities.research import ConditionSpec, FactorSpec, UniverseSpec from app.domain.entities.research import ConditionSpec, FactorSpec, UniverseSpec
class SelectionStrategy(BaseModel): class SelectionStrategy(BaseModel):
"""一个选股策略 = 选股条件组合(不含任何回测执行参数)。""" """一个选股策略 = 选股条件组合(不含任何回测执行参数)。
`extra="forbid"`(2026-09 收尾):请求里若混入旧版的回测执行参数(selection /
rebalance / costs / portfolio / initial_capital / period …),一律**报错**而不是
静默丢弃 —— 否则调用方会以为「在策略上设了费率/调仓」,实际服务端根本没存
(AGENT.md 禁止静默降级与假装支持)。回测参数的正确位置是 BacktestCombo + GlobalConfig。
兼容性:历史行残留的旧键由仓储 `_to_entity` 在**读出前**剔除,因此不受 forbid 影响。
"""
model_config = ConfigDict(extra="forbid")
id: str = "" id: str = ""
name: str = Field(min_length=1, max_length=64) name: str = Field(min_length=1, max_length=64)
description: str = "" description: str = ""
spec_type: str = Field(default="selection", pattern="^(selection|backtest)$") # 取值域收敛为 selection:本实体就是「选股策略」。历史 DB 列里的 "backtest"
# 不会被读出(仓储 _to_entity 丢弃该键并回落到默认值),故收紧不会破坏旧数据。
spec_type: str = Field(default="selection", pattern="^selection$")
universe: UniverseSpec = UniverseSpec() universe: UniverseSpec = UniverseSpec()
factors: list[FactorSpec] = Field(min_length=1, description="打分因子(至少 1 个)") factors: list[FactorSpec] = Field(min_length=1, description="打分因子(至少 1 个)")
conditions: list[ConditionSpec] = Field( conditions: list[ConditionSpec] = Field(
@@ -33,7 +44,7 @@ class SelectionStrategy(BaseModel):
created_at: datetime | None = None created_at: datetime | None = None
@model_validator(mode="after") @model_validator(mode="after")
def _no_duplicate_factors(self) -> "SelectionStrategy": def _no_duplicate_factors(self) -> SelectionStrategy:
names = [f.name for f in self.factors] names = [f.name for f in self.factors]
if len(set(names)) != len(names): if len(set(names)) != len(names):
raise ValueError("factors 存在重复因子名") raise ValueError("factors 存在重复因子名")
@@ -0,0 +1,29 @@
"""字段库 Repository 协议(依赖倒置)。"""
from __future__ import annotations
from typing import Protocol
from app.domain.entities.condition_field import ConditionField
class ConditionFieldRepository(Protocol):
def list(self) -> list[ConditionField]:
"""按 sort_order, name 排序返回全部条目(含停用项)。"""
...
def get(self, name: str) -> ConditionField | None: ...
def insert_missing(self, items: list[ConditionField]) -> int:
"""只插入不存在的条目(seed 内置字段用)。
语义上是「只补不删、不覆盖」:已存在的行**原样保留** —— 用户在字段库里
改过的中文名/含义不会被下次 seed 冲掉(与 factor_definition 同规矩)。
"""
...
def save(self, item: ConditionField) -> ConditionField:
"""新增或整体更新一条(自定义字段增改、内置字段改文案/停用)。"""
...
def delete(self, name: str) -> bool: ...
@@ -0,0 +1,48 @@
"""factor_definition:支持参数化因子实例(2026-10)
两处改动:
1. `name` 64 → 128:参数化实例把参数写进名字
(`momentum(window=90,direction=higher_is_better)`),64 位不够留余量。
2. 新增 `enabled`:唯一由人配置的字段 —— 是否出现在因子下拉/字段库里。
内置实例的开关注仍由代码注册表收敛;停用**不影响**已引用它的策略/归档解析,
历史口径不能被开关改义。
SQLite 不支持直接改列类型,因此用 batch_alter_table(与本仓库既有迁移一致)。
"""
from __future__ import annotations
from collections.abc import Sequence
import sqlalchemy as sa
from alembic import op
revision: str = "a7c1e4b90f21"
down_revision: str | None = "d6e7f8a9b0c1"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
def upgrade() -> None:
with op.batch_alter_table("factor_definition", schema=None) as batch_op:
batch_op.alter_column(
"name",
existing_type=sa.String(length=64),
type_=sa.String(length=128),
existing_nullable=False,
)
op.add_column(
"factor_definition",
sa.Column("enabled", sa.Boolean(), nullable=False, server_default="1"),
)
def downgrade() -> None:
op.drop_column("factor_definition", "enabled")
with op.batch_alter_table("factor_definition", schema=None) as batch_op:
batch_op.alter_column(
"name",
existing_type=sa.String(length=128),
type_=sa.String(length=64),
existing_nullable=False,
)
@@ -0,0 +1,155 @@
"""重算存量选股策略的过时 description(2026-09 重构收尾)
Revision ID: c5d6e7f8a9b0
Revises: b4c5d6e7f8a9
Create Date: 2026-10-01
背景:b4c5d6e7f8a9 把一个策略的 config_json 里回测执行参数剥掉了,但**没有**重算
`strategy.description`。旧描述是重构前由 describe_strategy 从「全套参数」自动生成的,
于是策略库里会出现这种自相矛盾的说明:
「…每 6 个月重新择股、每 6 个月调仓,后复权口径、按调仓日收盘价成交
(含佣金 0.03%/印花税 0.05%/滑点 0.1%)。」
而选股策略现在**不再持有**调仓/成本/复权,这些由「回测组合 + 公共配置」在回测时决定。
本迁移用当前口径的 describe_strategy(纯函数,无 IO/DB)重算这些陈旧说明。
安全性 —— 只改「可证明是旧自动生成」的行,不碰人工撰写的说明:
1. description 为空/纯空白(API 保存契约要求必须有说明,空值必然是历史遗留)→ 补全;
2. description 含旧自动文案独有的回测执行词(佣金/印花税/滑点/调仓/择股/复权口径/
收盘价成交/最低佣金/初始资金)→ 重算。新口径的说明**绝不会**出现这些词
(见 strategy_doc._describe_selection_only),因此命中即旧自动文案。
其余行原样保留(kept)。无法解析/校验失败的行跳过并打印告警,绝不静默改写。
为什么在迁移里 import 应用代码:说明文本的唯一事实来源就是 `describe_strategy`
(AGENT.md §24:不许另写一份近似文案)。自己复制一份文案逻辑才是真正的漂移风险。
代价是该迁移的产物依赖当时的代码版本 —— 对「一次性回填存量说明」这个用途可以接受,
且新库 upgrade 时 strategy 表为空、不受影响。
downgrade 仅回滚结构层面:**不恢复**被重算的旧说明(原文未备份),因此不可逆。
"""
from __future__ import annotations
import json
import re
from collections.abc import Sequence
import sqlalchemy as sa
from alembic import op
revision: str = "c5d6e7f8a9b0"
down_revision: str | None = "b4c5d6e7f8a9"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
# 与 strategy.description 列宽一致(StrategyModel.description = String(300))
_DESCRIPTION_MAX_CHARS = 300
# 选股策略不再承载的键(与 b4c5d6e7f8a9 一致;历史行可能仍残留)
_LEGACY_KEYS = (
"selection",
"rebalance",
"costs",
"portfolio",
"price_adjustment",
"selection_interval_months",
"rebalance_interval_months",
)
# 由 DB 列承载、不应从 config_json 再喂给实体的键
_COLUMN_KEYS = ("id", "name", "description", "spec_type", "version")
# 旧「全套参数」自动文案独有的回测执行词 —— 新口径说明不会出现(命中即认定陈旧)
_LEGACY_MARKERS = (
"佣金",
"印花税",
"滑点",
"调仓",
"择股",
"复权口径",
"收盘价成交",
"最低佣金",
"初始资金",
)
_LEGACY_MARKER_RE = re.compile("|".join(_LEGACY_MARKERS))
def _truncate(text: str) -> str:
"""与 API 的说明补全同口径:超列宽按字符截断并显式加省略号。"""
if len(text) <= _DESCRIPTION_MAX_CHARS:
return text
return text[: _DESCRIPTION_MAX_CHARS - 1] + "…"
def _derive_summary(name: str, description: str, data: dict) -> str | None:
"""按当前口径重算一句话说明;无法构造实体时返回 None(调用方跳过并告警)。"""
# 延迟 import:保持迁移模块导入轻量,且让 alembic env 先完成自身引导。
from app.domain.entities.strategy import SelectionStrategy
from app.quant.strategy_doc import describe_strategy
payload = dict(data)
for key in _COLUMN_KEYS + _LEGACY_KEYS:
payload.pop(key, None)
try:
st = SelectionStrategy(name=name, description=description, **payload)
except Exception as exc: # noqa: BLE001 —— 逐行容错:坏行跳过并告警,不阻断整次迁移
print(f"[refresh-strategy-docs] 跳过无法解析的策略 {name!r}: {exc}", flush=True)
return None
return describe_strategy(st).summary
def _is_stale(name: str, description: str) -> bool:
return (not (description or "").strip()) or bool(_LEGACY_MARKER_RE.search(description or ""))
def upgrade() -> None:
conn = op.get_bind()
rows = conn.execute(
sa.text("SELECT id, name, description, config_json FROM strategy")
).fetchall()
rewritten = kept = broken = 0
for row_id, name, description, cfg_text in rows:
try:
data = json.loads(cfg_text) if cfg_text else {}
except json.JSONDecodeError:
broken += 1
print(f"[refresh-strategy-docs] 跳过 config_json 损坏的策略 {row_id}", flush=True)
continue
if not isinstance(data, dict):
broken += 1
print(f"[refresh-strategy-docs] 跳过 config_json 非对象的策略 {row_id}", flush=True)
continue
if not _is_stale(name, description):
kept += 1 # 人工撰写的说明:不动它
continue
summary = _derive_summary(name, description or "", data)
if summary is None:
broken += 1
continue
summary = _truncate(summary)
if summary == (description or ""):
kept += 1
continue
conn.execute(
sa.text("UPDATE strategy SET description = :desc WHERE id = :id"),
{"desc": summary, "id": row_id},
)
rewritten += 1
print(
f"[refresh-strategy-docs] 重算 {rewritten} 条陈旧/空说明,"
f"保留 {kept} 条,跳过 {broken} 条异常行(共 {len(rows)} 条)",
flush=True,
)
def downgrade() -> None:
# 旧说明原文未备份,无法还原:回滚只表示「结构层面无事可做」。
# 显式空实现(而非 pass 无说明),避免读者误以为会恢复文案。
print(
"[refresh-strategy-docs] downgrade:被重算的说明不可还原(原文未备份),不执行任何写操作",
flush=True,
)
@@ -0,0 +1,45 @@
"""condition_field 表(2026-10 字段库:过滤条件字段目录入库)
Revision ID: d6e7f8a9b0c1
Revises: c5d6e7f8a9b0
Create Date: 2026-10-01
背景:策略库的过滤条件此前只能手填字段名(dv_ratio / static.industry …),
用户看不到含义、写错也不报错(未知字段求值恒为 None,条件永远不通过)。
本表存放字段库目录:内置字段由 quant/condition_fields.py 注册表在 API 首次读取时
seed(只补不删,不覆盖用户改过的文案),自定义字段与停用状态也落在本表。
"""
from __future__ import annotations
from collections.abc import Sequence
import sqlalchemy as sa
from alembic import op
revision: str = "d6e7f8a9b0c1"
down_revision: str | None = "c5d6e7f8a9b0"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
def upgrade() -> None:
op.create_table(
"condition_field",
sa.Column("name", sa.String(length=64), nullable=False),
sa.Column("label", sa.String(length=64), nullable=False),
sa.Column("description", sa.String(length=500), nullable=False),
sa.Column("kind", sa.String(length=8), nullable=False),
sa.Column("group_name", sa.String(length=32), nullable=False),
sa.Column("unit", sa.String(length=16), nullable=False),
sa.Column("source", sa.String(length=8), nullable=False),
sa.Column("enabled", sa.Boolean(), nullable=False),
sa.Column("sort_order", sa.Integer(), nullable=False),
sa.Column("created_at", sa.DateTime(), nullable=False),
sa.Column("updated_at", sa.DateTime(), nullable=False),
sa.PrimaryKeyConstraint("name"),
)
def downgrade() -> None:
op.drop_table("condition_field")
@@ -4,13 +4,16 @@
模型统一继承 infra.persistence.sqlalchemy.base.Base。 模型统一继承 infra.persistence.sqlalchemy.base.Base。
""" """
from app.infrastructure.persistence.sqlalchemy.models.composite import ( # noqa: F401
FactorCompositeModel,
)
from app.infrastructure.persistence.sqlalchemy.models.combo import ( # noqa: F401 from app.infrastructure.persistence.sqlalchemy.models.combo import ( # noqa: F401
BacktestComboModel, BacktestComboModel,
GlobalConfigModel, GlobalConfigModel,
) )
from app.infrastructure.persistence.sqlalchemy.models.composite import ( # noqa: F401
FactorCompositeModel,
)
from app.infrastructure.persistence.sqlalchemy.models.condition_field import ( # noqa: F401
ConditionFieldModel,
)
from app.infrastructure.persistence.sqlalchemy.models.factor import ( # noqa: F401 from app.infrastructure.persistence.sqlalchemy.models.factor import ( # noqa: F401
FactorDefinitionModel, FactorDefinitionModel,
) )
@@ -0,0 +1,31 @@
"""字段库表(2026-10)。
condition_field:过滤条件字段的目录契约源(name 主键幂等)。
内置字段由 quant/condition_fields.py 的注册表在 API 首次读取时 seed(只补不删),
自定义字段与用户改过的文案都落在这张表里。
"""
from __future__ import annotations
from datetime import datetime
from sqlalchemy import Boolean, DateTime, Integer, String
from sqlalchemy.orm import Mapped, mapped_column
from app.infrastructure.persistence.sqlalchemy.base import Base
class ConditionFieldModel(Base):
__tablename__ = "condition_field"
name: Mapped[str] = mapped_column(String(64), primary_key=True)
label: Mapped[str] = mapped_column(String(64), default="")
description: Mapped[str] = mapped_column(String(500), default="")
kind: Mapped[str] = mapped_column(String(8), default="num")
group_name: Mapped[str] = mapped_column(String(32), default="行情")
unit: Mapped[str] = mapped_column(String(16), default="")
source: Mapped[str] = mapped_column(String(8), default="builtin")
enabled: Mapped[bool] = mapped_column(Boolean, default=True)
sort_order: Mapped[int] = mapped_column(Integer, default=100)
created_at: Mapped[datetime] = mapped_column(DateTime)
updated_at: Mapped[datetime] = mapped_column(DateTime)
@@ -1,13 +1,16 @@
"""因子目录表(M7.1)。 """因子目录表(M7.1;2026-10 支持参数化实例)。
factor_definition:因子元数据契约源(name 主键幂等);requires 以 JSON 存。 factor_definition:因子名(主键,参数化实例的参数就写在名字里)→ 元数据;requires 以 JSON 存。
`enabled` 是唯一由人配置的字段(是否出现在因子下拉里);内置实例的开关注由代码注册表
收敛(见 application/services/factor_catalog.py),停用只影响「能不能被选中」,
不影响已引用它的策略/归档解析 —— 历史不能被开关改义。
""" """
from __future__ import annotations from __future__ import annotations
from datetime import datetime from datetime import datetime
from sqlalchemy import DateTime, Integer, String, Text from sqlalchemy import Boolean, DateTime, Integer, String, Text
from sqlalchemy.orm import Mapped, mapped_column from sqlalchemy.orm import Mapped, mapped_column
from app.infrastructure.persistence.sqlalchemy.base import Base from app.infrastructure.persistence.sqlalchemy.base import Base
@@ -16,7 +19,8 @@ from app.infrastructure.persistence.sqlalchemy.base import Base
class FactorDefinitionModel(Base): class FactorDefinitionModel(Base):
__tablename__ = "factor_definition" __tablename__ = "factor_definition"
name: Mapped[str] = mapped_column(String(64), primary_key=True) # 128:参数化实例的名字把参数写全(如 momentum(window=90,direction=lower_is_better))
name: Mapped[str] = mapped_column(String(128), primary_key=True)
description: Mapped[str] = mapped_column(String(500), default="") description: Mapped[str] = mapped_column(String(500), default="")
formula: Mapped[str] = mapped_column(String(500), default="") formula: Mapped[str] = mapped_column(String(500), default="")
brief: Mapped[str] = mapped_column(String(500), default="") brief: Mapped[str] = mapped_column(String(500), default="")
@@ -25,4 +29,5 @@ class FactorDefinitionModel(Base):
direction: Mapped[str] = mapped_column(String(32), default="higher_is_better") direction: Mapped[str] = mapped_column(String(32), default="higher_is_better")
requires_json: Mapped[str] = mapped_column(Text, default="[]") requires_json: Mapped[str] = mapped_column(Text, default="[]")
version: Mapped[str] = mapped_column(String(16), default="1") version: Mapped[str] = mapped_column(String(16), default="1")
enabled: Mapped[bool] = mapped_column(Boolean, default=True, server_default="1")
created_at: Mapped[datetime] = mapped_column(DateTime) created_at: Mapped[datetime] = mapped_column(DateTime)
@@ -16,7 +16,8 @@ class StrategyModel(Base):
id: Mapped[str] = mapped_column(String(32), primary_key=True) id: Mapped[str] = mapped_column(String(32), primary_key=True)
name: Mapped[str] = mapped_column(String(64), unique=True) name: Mapped[str] = mapped_column(String(64), unique=True)
description: Mapped[str] = mapped_column(String(300), default="") description: Mapped[str] = mapped_column(String(300), default="")
spec_type: Mapped[str] = mapped_column(String(16), default="backtest") # 2026-09 重构后策略库只存选股策略,新行一律 selection(历史行的 backtest 由数据迁移收敛)
spec_type: Mapped[str] = mapped_column(String(16), default="selection")
config_json: Mapped[str] = mapped_column(Text) config_json: Mapped[str] = mapped_column(Text)
version: Mapped[str] = mapped_column(String(16), default="1") version: Mapped[str] = mapped_column(String(16), default="1")
created_at: Mapped[datetime] = mapped_column(DateTime) created_at: Mapped[datetime] = mapped_column(DateTime)
@@ -0,0 +1,102 @@
"""字段库 Repository 的 SQLAlchemy 实现(2026-10)。"""
from __future__ import annotations
from datetime import datetime
from sqlalchemy import select
from sqlalchemy.orm import Session
from app.domain.entities.condition_field import ConditionField
from app.infrastructure.persistence.sqlalchemy.models.condition_field import ConditionFieldModel
def _to_entity(row: ConditionFieldModel) -> ConditionField:
return ConditionField(
name=row.name,
label=row.label,
description=row.description,
kind=row.kind,
group_name=row.group_name,
unit=row.unit,
source=row.source,
enabled=row.enabled,
sort_order=row.sort_order,
created_at=row.created_at,
updated_at=row.updated_at,
)
class SqlAlchemyConditionFieldRepository:
def __init__(self, session: Session) -> None:
self._session = session
def list(self) -> list[ConditionField]:
rows = self._session.scalars(
select(ConditionFieldModel).order_by(
ConditionFieldModel.sort_order, ConditionFieldModel.name
)
).all()
return [_to_entity(r) for r in rows]
def get(self, name: str) -> ConditionField | None:
row = self._session.get(ConditionFieldModel, name)
return _to_entity(row) if row else None
def insert_missing(self, items: list[ConditionField]) -> int:
if not items:
return 0
existing = set(
self._session.scalars(
select(ConditionFieldModel.name).where(
ConditionFieldModel.name.in_([i.name for i in items])
)
)
)
now = datetime.now()
added = 0
for item in items:
if item.name in existing:
continue # 只补不删、不覆盖用户改过的文案
self._session.add(
ConditionFieldModel(
name=item.name,
label=item.label,
description=item.description,
kind=item.kind,
group_name=item.group_name,
unit=item.unit,
source=item.source,
enabled=item.enabled,
sort_order=item.sort_order,
created_at=now,
updated_at=now,
)
)
added += 1
return added
def save(self, item: ConditionField) -> ConditionField:
now = datetime.now()
row = self._session.get(ConditionFieldModel, item.name)
if row is None:
row = ConditionFieldModel(name=item.name, created_at=now, updated_at=now)
self._session.add(row)
row.label = item.label
row.description = item.description
row.kind = item.kind
row.group_name = item.group_name
row.unit = item.unit
row.source = item.source
row.enabled = item.enabled
row.sort_order = item.sort_order
row.updated_at = now
self._session.flush()
return _to_entity(row)
def delete(self, name: str) -> bool:
row = self._session.get(ConditionFieldModel, name)
if row is None:
return False
self._session.delete(row)
return True
@@ -1,4 +1,8 @@
"""因子目录 Repository 的 SQLAlchemy 实现(M7.1)。""" """因子目录 Repository 的 SQLAlchemy 实现(M7.1;2026-10 加 enabled)。
只写**真列**:实体上还有 template/params/label/param_specs 等由名字解析出来的投影字段,
它们不是列 —— 若照 model_dump 全量 setattr,会出现「看着写进去了、其实没落库」的假象。
"""
from __future__ import annotations from __future__ import annotations
@@ -9,7 +13,23 @@ from sqlalchemy import select
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
from app.domain.entities.factor import FactorDefinition from app.domain.entities.factor import FactorDefinition
from app.infrastructure.persistence.sqlalchemy.models.factor import FactorDefinitionModel from app.infrastructure.persistence.sqlalchemy.models.factor import (
FactorDefinitionModel,
)
# 真正落库的列(其余实体字段是解析投影)
_STORED_COLUMNS = (
"name",
"description",
"formula",
"brief",
"frequency",
"lookback",
"direction",
"requires",
"version",
"enabled",
)
def _to_entity(row: FactorDefinitionModel) -> FactorDefinition: def _to_entity(row: FactorDefinitionModel) -> FactorDefinition:
@@ -23,6 +43,7 @@ def _to_entity(row: FactorDefinitionModel) -> FactorDefinition:
direction=row.direction, direction=row.direction,
requires=json.loads(row.requires_json or "[]"), requires=json.loads(row.requires_json or "[]"),
version=row.version, version=row.version,
enabled=bool(row.enabled),
created_at=row.created_at, created_at=row.created_at,
) )
@@ -57,11 +78,13 @@ class SqlAlchemyFactorRepository:
direction=d.direction, direction=d.direction,
requires_json=json.dumps(d.requires), requires_json=json.dumps(d.requires),
version=d.version, version=d.version,
enabled=d.enabled,
created_at=now, created_at=now,
) )
) )
else: else:
for k, v in d.model_dump(exclude={"created_at"}).items(): for k in _STORED_COLUMNS:
v = getattr(d, k)
if k == "requires": if k == "requires":
v = json.dumps(v) v = json.dumps(v)
setattr(row, k, v) setattr(row, k, v)
+375
View File
@@ -0,0 +1,375 @@
"""条件字段注册表 —— 「字段库」的唯一事实来源(2026-10)。
要解决的问题
------------
策略库的「过滤条件」此前是**手填字段名**的输入框:用户必须知道 dv_ratio /
static.industry / fundamental.roe 这类内部标识,既看不到含义,写错了也不报错 ——
引擎对未知字段求值一律返回 None,条件**永远不通过**,策略会安静地选出 0 只股票。
这正是 AGENT.md 禁止的「静默失败 / 假装支持」。
本模块的职责
------------
把**引擎真正支持的字段域**集中声明一次(中文名 + 含义 + 单位 + 分组 + 类型 + 排序),
供三方共用:
1. ``/api/condition-fields`` 据此 seed 目录、据此校验用户新增的自定义字段;
2. 前端据此渲染分组下拉、含义提示、并按类型收窄可选比较符;
3. :func:`is_supported_field` 直接查 ``Stock`` / ``FinancialIndicator`` 的字段定义与
因子注册表(``quant.factors``),**不另写一套近似规则** —— 避免注册表与引擎漂移。
诚实性约束(AGENT.md §24)
--------------------------
只登记真能算的字段。日期字段(如 ``static.list_date``)无法比较大小,:func:`reason_unsupported`
会明确拒绝并说明理由,而不是放行让用户建出一条「永远选不出股票」的条件。
单位一律照抄数据源落库口径(见 ``data_sources/tushare.py`` 的换算注释),不凭印象写。
"""
from __future__ import annotations
from dataclasses import dataclass
from app.domain.entities.condition_field import OPS_BY_KIND
from app.domain.entities.market import (
DAILY_BAR_NUMERIC_FIELDS,
DAILY_BASIC_NUMERIC_FIELDS,
FinancialIndicator,
Stock,
)
from app.quant.factors import FactorError, get_factor, list_factors, resolve_factor
# ---------- 分组(下拉的 optgroup 顺序即此顺序) ----------
GROUP_STOCK = "股票基础"
GROUP_QUOTE = "行情"
GROUP_TECH = "技术指标"
GROUP_DAILY = "每日指标"
GROUP_FUNDAMENTAL = "财务指标"
GROUP_FACTOR = "因子"
GROUP_ORDER: tuple[str, ...] = (
GROUP_STOCK,
GROUP_QUOTE,
GROUP_TECH,
GROUP_DAILY,
GROUP_FUNDAMENTAL,
GROUP_FACTOR,
)
# 比较符规则(哪种类型能比大小)定义在领域实体 domain/entities/condition_field.py,
# 此处转发 —— 保证「字段库 API 返回的 ops」与「注册表里 FieldDef.ops」出自同一处。
# ---------- 单位阶梯(2026-10) ----------
#
# 单位分两层,避免「改个显示单位把历史策略的数值偷偷换义」:
# · **基准单位**(FieldDef.unit):引擎存储与比较用的单位,写死在数据源落库口径里,
# 不可改。归档里的 ConditionSpec 存的永远是基准单位值 —— 复现不受界面设置影响。
# · **界面单位**(本阶梯里的备选项):只在「输入/显示」这一层做换算,factor 表示
# 「该单位 → 基准单位的系数」(即提交前 ×factor,回显时 ÷factor)。
# 所以用户在字段库把总市值选成「亿元」,输入 5 会存成 50000(万元)——引擎比较的仍是
# 基准单位,而界面上看到的始终是 5 亿元。改单位不会让任何历史策略变义。
#
# 只登记换算无歧义、且实际会用到的单位:金额(元/万元/亿元)、股数(股/手/万手)、
# 股本(万股/亿股)。百分数(%)与倍数(倍)不提供备选 —— 换成小数只会制造误读。
U_MONEY_YUAN: tuple[tuple[str, float], ...] = (("元", 1.0), ("万元", 1e4), ("亿元", 1e8))
U_MONEY_WAN: tuple[tuple[str, float], ...] = (("万元", 1.0), ("亿元", 1e4))
U_SHARE_GU: tuple[tuple[str, float], ...] = (("股", 1.0), ("手", 100.0), ("万手", 1e6))
U_SHARE_WAN: tuple[tuple[str, float], ...] = (("万股", 1.0), ("亿股", 1e4))
# 哪些字段提供备选界面单位(键 = 引擎字段名;不在表里的字段只能用基准单位)
_UNIT_LADDERS: dict[str, tuple[tuple[str, float], ...]] = {
# 行情原列:volume 入库为股(源为手 ×100),amount 入库为元(源为千元 ×1000)
"volume": U_SHARE_GU,
"amount": U_MONEY_YUAN,
# 每日指标:市值为万元,股本为万股
"total_mv": U_MONEY_WAN,
"circ_mv": U_MONEY_WAN,
"total_share": U_SHARE_WAN,
"float_share": U_SHARE_WAN,
"free_share": U_SHARE_WAN,
# 财务指标:金额入库为元
"fundamental.net_profit": U_MONEY_YUAN,
"fundamental.total_revenue": U_MONEY_YUAN,
}
@dataclass(frozen=True)
class FieldDef:
"""一个条件字段的登记项(引擎真能算的字段)。"""
name: str # 引擎字段名(写进 condition.field)
label: str # 中文名(下拉里给人看的)
description: str # 含义 / 口径(含单位),必须可核对
kind: str # num | str
group_name: str
unit: str = "" # **基准单位**(引擎存储/比较用),不可由界面更改
curated: bool = True # True=默认进字段库;False=仅登记为「可新增」(少用字段)
units: tuple[tuple[str, float], ...] = () # 可选界面单位;(单位, →基准单位系数),首项须是基准单位
@property
def ops(self) -> tuple[str, ...]:
return OPS_BY_KIND.get(self.kind, ())
@property
def unit_options(self) -> tuple[tuple[str, float], ...]:
"""可选界面单位;未登记阶梯的字段只有基准单位一项(界面不给选择)。"""
return self.units or ((self.unit, 1.0),)
# ---------- 股票基础(static.*):来自 Stock 实体,只有字符串字段可比较 ----------
# (attr, 中文名, 含义, curated)
_STATIC_FIELDS: tuple[tuple[str, str, str, bool], ...] = (
("industry", "所属行业", "股票基础信息里的行业名称(字符串),如「银行」「白酒」。等值用「=」,多值用「属于」。", True),
("market", "上市板块", "主板 / 创业板 / 科创板 / 北交所(字符串)。", True),
("area", "注册地域", "公司注册地省份或地区(字符串),如「广东」「北京」。", True),
("exchange", "交易所", "SH 上交所 / SZ 深交所 / BJ 北交所(字符串)。", False),
("status", "上市状态", "L 上市 / D 退市 / P 暂停上市(字符串)。研究池已默认剔除退市股。", False),
("name", "股票名称", "证券简称(字符串)。一般用于核对,不建议拿来做条件。", False),
("symbol", "股票代码", "Tushare 风格代码,如 600519.SH(字符串)。多值用「属于」。", False),
)
# ---------- 行情原列(open/high/low/close/volume/amount) ----------
# 换算口径见 data_sources/tushare.py: volume=vol(手)*100 → 股;amount=amount(千元)*1000 → 元。
_BAR_FIELDS: tuple[tuple[str, str, str, str, bool], ...] = (
("close", "收盘价", "当日收盘价(元)。复权口径由公共配置的 price_adjustment 决定(默认 hfq)。", "元", True),
("open", "开盘价", "当日开盘价(元),复权口径同上。", "元", True),
("high", "最高价", "当日最高价(元),复权口径同上。", "元", True),
("low", "最低价", "当日最低价(元),复权口径同上。", "元", True),
("volume", "成交量", "当日成交股数(股)。数据源原始单位为「手」,入库时已 ×100 换算。", "股", True),
("amount", "成交额", "当日成交金额(元)。数据源原始单位为「千元」,入库时已 ×1000 换算。", "元", True),
)
# ---------- 技术派生(滚动窗口计算,非行情原列) ----------
_TECH_FIELDS: tuple[tuple[str, str, str, str, bool], ...] = (
("ma20", "20 日均线", "收盘价的 20 个交易日简单移动平均(元),在选股日当日取值。", "元", True),
("ma60", "60 日均线", "收盘价的 60 个交易日简单移动平均(元),在选股日当日取值。", "元", True),
)
# ---------- 每日指标(daily_basic) ----------
# 单位照抄 data_sources/tushare.py: 百分数/倍数为原样,股本为万股,市值为万元。
_DAILY_FIELDS: tuple[tuple[str, str, str, str, bool], ...] = (
("dv_ratio", "股息率", "近 12 个月现金分红 / 总市值 × 100(%),逐日时点值 —— 即因子 dividend_yield 的口径。", "%", True),
("dv_ttm", "股息率 TTM", "近 12 个月滚动现金分红 / 总市值 × 100(%),即因子 dividend_yield_ttm 的口径。", "%", True),
("pe", "市盈率 PE", "总市值 / 最新年报净利润(倍,静态口径)。", "倍", True),
("pe_ttm", "市盈率 PE(TTM)", "总市值 / 最近 12 个月净利润(倍)。", "倍", True),
("pb", "市净率 PB", "总市值 / 最新报告期净资产(倍)。", "倍", True),
("turnover_rate", "换手率", "当日成交股数 / 流通股本 × 100(%)。", "%", True),
("volume_ratio", "量比", "当日成交量 / 过去 5 日平均成交量(倍)。", "倍", True),
("total_mv", "总市值", "总股本 × 当日收盘价(万元)。", "万元", True),
("circ_mv", "流通市值", "流通股本 × 当日收盘价(万元)。", "万元", True),
("ps", "市销率 PS", "总市值 / 最新年报营业收入(倍)。", "倍", False),
("ps_ttm", "市销率 PS(TTM)", "总市值 / 最近 12 个月营业收入(倍)。", "倍", False),
("total_share", "总股本", "总股本(万股)。", "万股", False),
("float_share", "流通股本", "流通股本(万股)。", "万股", False),
("free_share", "自由流通股本", "自由流通股本(万股)。", "万股", False),
)
# ---------- 财务指标(fundamental.*) ----------
# 可见性口径:只取 announce_date <= 选股日的最新已公告值(防未来函数,见 selection.run_condition_selection)。
_FUNDAMENTAL_FIELDS: tuple[tuple[str, str, str, str, bool], ...] = (
("roe", "净资产收益率 ROE", "最新已公告报告期的净资产收益率(%)。", "%", True),
("eps", "每股收益 EPS", "最新已公告报告期的每股收益(元)。", "元", True),
("gross_margin", "毛利率", "最新已公告报告期的毛利率(%)。", "%", True),
("net_profit", "归母净利润", "最新已公告报告期的归母净利润(元)。", "元", False),
("total_revenue", "营业总收入", "最新已公告报告期的营业总收入(元)。注意:目前只有新浪兜底源提供该字段,多数行可能为空 —— 缺失时条件视为不通过。", "元", False),
)
# static.* 里明确不支持的字段(存在但不可比较)
_STATIC_UNSUPPORTED: dict[str, str] = {
"list_date": "上市日期是日期,不是可比较的数值/字符串;请改用股票池的「上市天数」设置",
"delist_date": "退市日期是日期,不是可比较的数值/字符串",
}
def _defs() -> list[FieldDef]:
"""构造全部内置字段定义(每次调用重新构造,保证与代码注册表实时一致)。"""
out: list[FieldDef] = []
for attr, label, desc, curated in _STATIC_FIELDS:
out.append(
FieldDef(
name=f"static.{attr}",
label=label,
description=desc,
kind="str",
group_name=GROUP_STOCK,
curated=curated,
)
)
for name, label, desc, unit, curated in _BAR_FIELDS:
out.append(
FieldDef(
name, label, desc, "num", GROUP_QUOTE, unit, curated,
_UNIT_LADDERS.get(name, ()),
)
)
for name, label, desc, unit, curated in _TECH_FIELDS:
out.append(
FieldDef(name, label, desc, "num", GROUP_TECH, unit, curated, _UNIT_LADDERS.get(name, ()))
)
for name, label, desc, unit, curated in _DAILY_FIELDS:
out.append(
FieldDef(name, label, desc, "num", GROUP_DAILY, unit, curated, _UNIT_LADDERS.get(name, ()))
)
for name, label, desc, unit, curated in _FUNDAMENTAL_FIELDS:
full = f"fundamental.{name}"
out.append(
FieldDef(full, label, desc, "num", GROUP_FUNDAMENTAL, unit, curated, _UNIT_LADDERS.get(full, ()))
)
for d in list_factors():
# 因子既能当「打分因子」也能当「过滤条件」:这里复用因子注册表的元数据,
# 不另写描述,避免两处文案漂移。因子的 description 里已写明公式与口径;
# 标签用中文名(含参数),如「动量(窗口 60,越高越好)」——
# 参数化实例不在这里(它们在 /factors 目录里,条件下拉按名并入)。
out.append(
FieldDef(
name=d.name,
label=f"{d.display}(因子)",
description=d.description + (f" 用法:{d.brief}" if d.brief else ""),
kind="num",
group_name=GROUP_FACTOR,
curated=True,
)
)
return out
def builtin_fields() -> list[FieldDef]:
"""全部引擎支持的字段(含 curated=False 的「可新增但不默认展示」项)。
名字唯一性由因子名与各分组前缀保证;一旦重复说明注册表写错,直接抛错而非静默覆盖。
单位阶梯同样自检:首项必须是基准单位且系数为 1.0,系数必须为正 —— 写错了会让
「界面显示 5 亿元、引擎按 5 万元比」这种错静默溜进生产。
"""
defs = _defs()
names = [d.name for d in defs]
dup = {n for n in names if names.count(n) > 1}
if dup:
raise ValueError(f"条件字段注册表存在重名:{sorted(dup)}")
for d in defs:
if not d.units:
continue
base, factor = d.units[0]
if base != d.unit or factor != 1.0:
raise ValueError(
f"字段 {d.name} 的单位阶梯首项必须是基准单位 {d.unit!r}(系数 1.0),实际 {d.units[0]!r}"
)
if any(f <= 0 for _, f in d.units):
raise ValueError(f"字段 {d.name} 的单位换算系数必须为正:{d.units}")
if len({u for u, _ in d.units}) != len(d.units):
raise ValueError(f"字段 {d.name} 的单位阶梯有重复单位:{d.units}")
return defs
def curated_fields() -> list[FieldDef]:
"""默认进「字段库」的字段(下拉里开箱可见的那批)。"""
return [d for d in builtin_fields() if d.curated]
def get_field(name: str) -> FieldDef | None:
"""按字段名取定义:注册表字段,或**参数化因子键**(如 momentum(window=90,direction=…))。
为什么参数化因子也要能取到:它是引擎真认的条件字段(`momentum_60 > 0` 一直合法),
而字段库/说明书的单位后缀、类型判断都走这里。取不到会让人误以为「引擎不支持」,
甚至让「把参数化因子加进字段库」这一步半路 assert 崩掉(500 而不是 422)。
"""
for d in builtin_fields():
if d.name == name:
return d
return _factor_field(name)
def _factor_field(name: str) -> FieldDef | None:
"""参数化因子键 → 字段定义(因子是无量纲量,不带单位)。"""
try:
defn, _fn = resolve_factor(name)
except FactorError:
return None
if defn.name in {d.name for d in builtin_fields()}:
return None # 注册表字段已在上一步返回;这里只处理新增的参数化实例
return FieldDef(
name=defn.name,
label=f"{defn.display}(因子)",
description=defn.description + (f" 用法:{defn.brief}" if defn.brief else ""),
kind="num",
group_name=GROUP_FACTOR,
curated=False, # 不自动进字段库:目录在 /factors 管,条件里按名并进来
)
def _stock_attrs() -> set[str]:
return set(Stock.model_fields)
def _fundamental_attrs() -> set[str]:
"""FinancialIndicator 里可比较的数值字段(排除 symbol/报告期/来源等元数据)。"""
skip = {"symbol", "report_date", "announce_date", "source"}
return {n for n in FinancialIndicator.model_fields if n not in skip}
def reason_unsupported(name: str) -> str:
"""字段不可用的理由(用于 422 文案;可用字段返回空串)。"""
if not name or not name.strip():
return "字段名为空"
if name in DAILY_BAR_NUMERIC_FIELDS:
return ""
if name in ("ma20", "ma60"):
return ""
if name in DAILY_BASIC_NUMERIC_FIELDS:
return ""
if name.startswith("static."):
attr = name[len("static.") :]
if attr in _STATIC_UNSUPPORTED:
return f"{name} 不可用作条件:{_STATIC_UNSUPPORTED[attr]}"
if attr in _stock_attrs():
return ""
return f"{name} 不存在:股票基础信息里没有 {attr} 字段"
if name.startswith("fundamental."):
attr = name[len("fundamental.") :]
if attr in _fundamental_attrs():
return ""
return f"{name} 不存在:财务指标里没有 {attr} 字段"
try:
get_factor(name)
except FactorError:
return (
f"{name} 不是引擎支持的字段。可用:行情列({', '.join(DAILY_BAR_NUMERIC_FIELDS)})、"
"ma20/ma60、每日指标列、static.<股票基础字段>、fundamental.<财务字段>、"
"已注册因子名,或参数化因子键(如 momentum(window=90,direction=higher_is_better))"
)
return ""
def is_supported_field(name: str) -> bool:
"""引擎是否真能算这个字段(False = 条件永远不通过,必须拒绝)。"""
return reason_unsupported(name) == ""
def available_fields(existing: set[str]) -> list[FieldDef]:
"""引擎支持但**尚未进目录**的字段(用户「新增字段」时可选项)。
只从注册表里挑,用户因此不可能加进一个引擎算不出来的字段(§24 不假装支持)。
"""
return [d for d in builtin_fields() if not d.curated and d.name not in existing]
# ---------- 单位换算(界面单位 ⇄ 基准单位) ----------
def unit_options(name: str) -> list[tuple[str, float]]:
"""该字段可选的界面单位(首项为基准单位,系数 = 该单位 → 基准单位)。字段不存在 → 空列表。
**换算在界面层做**(前端按系数换算输入/回显),存储与引擎一律用基准单位:
这样归档里的 ConditionSpec 永远不随界面设置改变含义。
"""
d = get_field(name)
return list(d.unit_options) if d else []
def unit_allowed(name: str, unit: str) -> bool:
"""该单位是否在字段允许的阶梯里(API 据此 422 拒绝自由文本单位)。"""
return unit in {u for u, _ in unit_options(name)}
+553 -155
View File
@@ -1,4 +1,34 @@
"""因子引擎:因子注册表、元数据与计算(Phase 2,低频选股因子)。 """因子引擎:因子**模板**(含可编辑参数)、实例注册表与计算(Phase 2 起,2026-10 参数化)。
## 三层概念(这是本模块的核心约定)
1. **模板(FactorTemplate)**:算法的家族,如 `momentum`(动量)、`volatility`(波动率)。
模板声明「哪些参数可编辑、允许范围、默认值」以及计算函数 `fn(fields, params)`。
2. **参数(params)**:模板的可编辑取值,如 `window=90`、`direction=lower_is_better`。
参数约束是**受控范围**(整数区间 / 枚举),不允许自由值 —— 见 AGENT.md §24:
写不进去就报错,绝不静默接受一个引擎其实不支持的设置。
3. **因子实例(FactorDef)**:`(模板, 参数)` 的具体因子,**名字里带着全部参数**:
momentum_60 ← 内置实例(代码里登记的历史名)
momentum(window=90,direction=higher_is_better) ← 参数化实例(目录里创建)
实例名即身份:参数写进名字,任何地方(策略 JSON、归档 spec、组合组件、条件字段)
存下这个名字,就同时冻结了「用哪个模板 + 哪些参数」——**历史归档不会因为之后
改了什么参数而改变含义**。这也是为什么不把参数放在另一个字段里:那需要改动
所有已经存了因子名的地方(策略/归档/回放/信号/Agent 工具),而且容易漏。
## 参数与默认值:为什么键里总是写全 direction
键里**不省略任何可编辑参数**(哪怕等于模板默认值)。若省略,`momentum(window=90)`
的含义就取决于「模板默认方向」这一代码事实:将来代码把默认方向一改,用户已经存下的
策略/归档会跟着变义 —— 与「单位」那次拒绝的做法同理。写全参数后,键自解释、
不依赖任何默认值,改默认值只影响新建实例。
## 兼容性
`get_factor()` 现在能吃两种名字:注册表里的历史名(内置实例)与参数化键;
`compute_factor()`、`list_factors()` 的行为保持不变,因此下游(选股/回测/组合/
说明书/条件字段/Agent 工具)无需感知参数化的存在,也**不会绕过参数校验**。
数据形态:行情长表 DataFrame(列 symbol/trade_date/close/high/low/volume/amount, 数据形态:行情长表 DataFrame(列 symbol/trade_date/close/high/low/volume/amount,
以及经 ResearchService 并入的每日指标列如 dv_ratio/dv_ttm), 以及经 ResearchService 并入的每日指标列如 dv_ratio/dv_ttm),
@@ -10,15 +40,68 @@
from __future__ import annotations from __future__ import annotations
from collections.abc import Callable import inspect
from dataclasses import dataclass import re
from collections.abc import Callable, Mapping
from dataclasses import dataclass, field
from typing import Any
import pandas as pd import pandas as pd
DIRECTION_HIGHER = "higher_is_better"
DIRECTION_LOWER = "lower_is_better"
DIRECTIONS = (DIRECTION_HIGHER, DIRECTION_LOWER)
# 参数名常量(键里的字面量,改它等于改所有已存键的含义,别改)
P_WINDOW = "window"
P_FAST = "fast"
P_SLOW = "slow"
P_DIRECTION = "direction"
# 窗口参数的允许范围:受控区间而非固定档(任意整数都能算,但要有边界)。
# 上限 500 个交易日 ≈ 两年,够长;下限 2 是因为 shift(1)/rolling(1) 的波动率无意义。
WINDOW_MIN = 2
WINDOW_MAX = 500
@dataclass(frozen=True)
class ParamSpec:
"""一个可编辑参数的约束(受控范围,越界一律报错而不是截断/静默忽略)。"""
name: str
label: str
kind: str # "int" | "enum"
default: Any
minimum: int | None = None
maximum: int | None = None
choices: tuple[str, ...] = ()
note: str = ""
def describe(self) -> str:
"""人类可读的约束说明(用于错误文案与目录展示)。"""
if self.kind == "enum":
return "、".join(self.choices)
if self.minimum is not None and self.maximum is not None:
return f"{self.minimum} ~ {self.maximum} 的整数"
return "整数"
DIRECTION_SPEC = ParamSpec(
name=P_DIRECTION,
label="方向",
kind="enum",
default=DIRECTION_HIGHER,
choices=DIRECTIONS,
note="越高越好 / 越低越好:决定复合分里的排序方向(低为好自动取负)。",
)
@dataclass(frozen=True) @dataclass(frozen=True)
class FactorDef: class FactorDef:
"""因子元数据(AGENT.md §22 要求逐项明确)。""" """因子元数据(AGENT.md §22 要求逐项明确)。
新增字段都带默认值:历史代码用位置参数构造 FactorDef 的地方不受影响。
"""
name: str name: str
description: str description: str
@@ -28,9 +111,55 @@ class FactorDef:
lookback: int = 20 lookback: int = 20
direction: str = "higher_is_better" # | lower_is_better direction: str = "higher_is_better" # | lower_is_better
requires: tuple[str, ...] = ("close",) requires: tuple[str, ...] = ("close",)
# ---- 参数化(2026-10)----
template: str = "" # 模板名,如 "momentum";空串 = 手工登记的老式因子
params: Mapping[str, Any] = field(default_factory=dict) # 冻结的参数取值
param_specs: tuple[ParamSpec, ...] = () # 可编辑参数与约束(供目录/界面)
source: str = "builtin" # builtin(代码注册表)| custom(目录里创建的参数化实例)
label: str = "" # 中文显示名(含参数),如「动量(窗口 90,越高越好)」
@property
def display(self) -> str:
"""界面用显示名:没有 label 时退回 name(老因子/自定义登记行)。"""
return self.label or self.name
FactorFn = Callable[[dict[str, pd.DataFrame]], pd.DataFrame] FactorFn = Callable[[dict[str, pd.DataFrame], Mapping[str, Any]], pd.DataFrame]
@dataclass(frozen=True)
class FactorTemplate:
"""算法家族 + 可编辑参数声明 + 内置实例(历史名)。"""
name: str
label: str # 中文家族名,如「动量」
description: str # 可含 {window} / {fast} / {slow} 占位
formula: str
brief: str
fn: FactorFn
param_specs: tuple[ParamSpec, ...] = ()
requires: tuple[str, ...] = ("close",)
frequency: str = "daily"
direction_default: str = DIRECTION_HIGHER
lookback_of: Callable[[Mapping[str, Any]], int] | None = None
check: Callable[[Mapping[str, Any]], str | None] | None = None # 跨参数约束
instances: tuple[tuple[str, Mapping[str, Any]], ...] = () # ((历史名, 参数), ...)
label_of: Callable[[Mapping[str, Any]], str] | None = None
def specs(self) -> tuple[ParamSpec, ...]:
"""全部可编辑参数(模板自己的参数 + 方向,方向恒在最后)。"""
direction = ParamSpec(
name=DIRECTION_SPEC.name,
label=DIRECTION_SPEC.label,
kind=DIRECTION_SPEC.kind,
default=self.direction_default,
choices=DIRECTION_SPEC.choices,
note=DIRECTION_SPEC.note,
)
return (*self.param_specs, direction)
def defaults(self) -> dict[str, Any]:
return {s.name: s.default for s in self.specs()}
class FactorError(ValueError): class FactorError(ValueError):
@@ -38,42 +167,318 @@ class FactorError(ValueError):
_REGISTRY: dict[str, tuple[FactorDef, FactorFn]] = {} _REGISTRY: dict[str, tuple[FactorDef, FactorFn]] = {}
_TEMPLATES: dict[str, FactorTemplate] = {}
# 参数化键:template(k=v,k=v)。模板名与参数名限定为标识符,值限定为标识符/数字,
# 避免出现靠运气才能解析的名字(宁可在创建时就被拒)。
_KEY_RE = re.compile(r"^(?P<template>[A-Za-z_][A-Za-z0-9_]*)\((?P<args>[^()]*)\)$")
_ARG_RE = re.compile(r"^(?P<key>[A-Za-z_][A-Za-z0-9_]*)=(?P<value>[A-Za-z_][A-Za-z0-9_]*|-?\d+)$")
def _accepts_params(fn: Callable) -> bool:
"""判断计算函数是否声明了 params 形参(兼容老的一参数写法)。"""
try:
params = inspect.signature(fn).parameters
except (TypeError, ValueError): # 内建/C 实现,按老写法处理
return False
if any(p.kind is inspect.Parameter.VAR_POSITIONAL for p in params.values()):
return True
positional = [
p
for p in params.values()
if p.kind in (inspect.Parameter.POSITIONAL_ONLY, inspect.Parameter.POSITIONAL_OR_KEYWORD)
]
return len(positional) >= 2
def _bind(fn: Callable) -> FactorFn:
"""把计算函数统一成 (fields, params) 两参数调用。"""
if _accepts_params(fn):
def _bound(fields, params, _fn=fn):
return _fn(fields, params)
return _bound
def _legacy(fields, params, _fn=fn):
return _fn(fields)
return _legacy
def register_template(template: FactorTemplate) -> FactorTemplate:
"""注册模板,并把它的内置实例(历史名)登记进注册表。"""
if template.name in _TEMPLATES:
raise FactorError(f"模板 {template.name} 已注册")
_TEMPLATES[template.name] = template
for name, params in template.instances:
defn = build_factor_def(template, params, name=name, source="builtin")
_REGISTRY[name] = (defn, _bind(template.fn))
return template
def register(defn: FactorDef) -> Callable[[FactorFn], FactorFn]: def register(defn: FactorDef) -> Callable[[FactorFn], FactorFn]:
"""装饰器:注册自定义因子。""" """装饰器:注册**自定义因子**(老式登记,无参数化;测试与扩展用)。"""
def deco(fn: FactorFn) -> FactorFn: def deco(fn: FactorFn) -> FactorFn:
if defn.name in _REGISTRY: if defn.name in _REGISTRY:
raise FactorError(f"因子 {defn.name} 已注册") raise FactorError(f"因子 {defn.name} 已注册")
_REGISTRY[defn.name] = (defn, fn) _REGISTRY[defn.name] = (defn, _bind(fn))
return fn return fn
return deco return deco
def get_factor(name: str) -> tuple[FactorDef, FactorFn]: def list_templates() -> list[FactorTemplate]:
if name not in _REGISTRY: return [t for _n, t in sorted(_TEMPLATES.items())]
raise FactorError(f"未知因子:{name}(可用:{', '.join(sorted(_REGISTRY))})")
return _REGISTRY[name]
def get_template(name: str) -> FactorTemplate:
if name not in _TEMPLATES:
raise FactorError(f"未知因子模板:{name}(可用:{', '.join(sorted(_TEMPLATES))})")
return _TEMPLATES[name]
def _render(text: str, params: Mapping[str, Any]) -> str:
"""渲染带 {param} 占位的文案;没有占位就原样返回(不做 format,避免误伤花括号)。"""
if "{" not in text:
return text
try:
return text.format(**params)
except KeyError as exc: # 模板写错占位名 —— 宁可当场炸,也不要漏出半成品文案
raise FactorError(f"因子文案占位符缺少参数 {exc}:{text}") from None
def _param_summary(template: FactorTemplate, params: Mapping[str, Any]) -> str:
"""非方向参数的摘要,如「窗口 90」「快线 5、慢线 60」。"""
return "、".join(f"{spec.label} {params[spec.name]}" for spec in template.param_specs)
def _default_label(template: FactorTemplate, params: Mapping[str, Any]) -> str:
summary = _param_summary(template, params)
direction = "越高越好" if params[P_DIRECTION] == DIRECTION_HIGHER else "越低越好"
inner = ",".join([p for p in (summary, direction) if p])
return f"{template.label}({inner})"
def fill_params(template: FactorTemplate, params: Mapping[str, Any]) -> dict[str, Any]:
"""校验参数并补全缺省值(**创建路径**用:内置实例、目录里新建参数化因子)。
只认声明过的参数,类型/范围/枚举全部受控,跨参数约束另查 —— 缺的参数取模板默认值。
"""
specs = {s.name: s for s in template.specs()}
unknown = sorted(set(params) - set(specs))
if unknown:
raise FactorError(
f"因子模板 {template.name} 不支持的参数:{', '.join(unknown)};"
f"可编辑参数只有 {', '.join(specs)}。参数不能随便加 —— 引擎算不了的要当场拒绝。"
)
merged: dict[str, Any] = {**template.defaults(), **params}
return _check_all(template, merged)
def validate_params(template: FactorTemplate, params: Mapping[str, Any]) -> dict[str, Any]:
"""严格校验:**每个参数都必须显式给出**(解析已存的因子键用)。
为什么解析时不许省:省掉的参数只能靠「模板默认值」补,而默认值是会随代码改的
事实 —— 一旦改了,用户早就存下的策略/归档就会跟着变义(模块头详述)。
"""
specs = {s.name: s for s in template.specs()}
unknown = sorted(set(params) - set(specs))
if unknown:
raise FactorError(
f"因子模板 {template.name} 不支持的参数:{', '.join(unknown)};"
f"可编辑参数只有 {', '.join(specs)}。"
)
missing = sorted(set(specs) - set(params))
if missing:
canonical = canonical_key(template, params)
raise FactorError(
f"因子键缺少参数:{', '.join(missing)}。键里必须写全所有参数"
f"(否则含义会取决于模板默认值);规范写法:{canonical}"
)
return _check_all(template, params)
def _check_all(template: FactorTemplate, params: Mapping[str, Any]) -> dict[str, Any]:
out = {spec.name: _check_param(template, spec, params[spec.name]) for spec in template.specs()}
if template.check is not None:
problem = template.check(out)
if problem:
raise FactorError(f"因子模板 {template.name} 参数不合法:{problem}")
return out
def _check_param(template: FactorTemplate, spec: ParamSpec, value: Any) -> Any:
if spec.kind == "enum":
if not isinstance(value, str) or value not in spec.choices:
raise FactorError(
f"因子「{template.name}」的参数 {spec.name}={value!r} 不合法:"
f"只能是 {spec.describe()}。"
)
return value
if isinstance(value, str): # 键里解析出来的是字符串,数字要能转
try:
value = int(value)
except ValueError:
raise FactorError(
f"因子「{template.name}」的参数 {spec.name}={value!r} 不是整数:"
f"应为 {spec.describe()}。"
) from None
if not isinstance(value, int) or isinstance(value, bool):
raise FactorError(
f"因子「{template.name}」的参数 {spec.name}={value!r} 不是整数:"
f"应为 {spec.describe()}。"
)
if spec.minimum is not None and value < spec.minimum:
raise FactorError(
f"因子「{template.name}」的参数 {spec.name}={value} 太小:应为 {spec.describe()}。"
)
if spec.maximum is not None and value > spec.maximum:
raise FactorError(
f"因子「{template.name}」的参数 {spec.name}={value} 太大:应为 {spec.describe()}。"
)
return value
def build_factor_def(
template: FactorTemplate,
params: Mapping[str, Any],
*,
name: str,
source: str = "custom",
) -> FactorDef:
"""由模板 + 参数构造因子实例(未知/越界参数在此处被拒)。"""
checked = fill_params(template, params)
lookback = template.lookback_of(checked) if template.lookback_of else 0
label_of = template.label_of or (lambda p: _default_label(template, p))
return FactorDef(
name=name,
description=_render(template.description, checked),
formula=_render(template.formula, checked),
brief=_render(template.brief, checked),
frequency=template.frequency,
lookback=lookback,
direction=checked[P_DIRECTION],
requires=template.requires,
template=template.name,
params=dict(checked),
param_specs=template.specs(),
source=source,
label=label_of(checked),
)
def canonical_key(template: FactorTemplate | str, params: Mapping[str, Any]) -> str:
"""参数化实例的规范名:参数按模板声明顺序写全(含方向),如
momentum(window=90,direction=higher_is_better)
写全的好处见模块头:键自解释,不依赖任何默认值。
"""
tpl = get_template(template) if isinstance(template, str) else template
checked = fill_params(tpl, params)
args = ",".join(f"{spec.name}={checked[spec.name]}" for spec in tpl.specs())
return f"{tpl.name}({args})"
def parse_factor_key(name: str) -> tuple[FactorTemplate, dict[str, Any]]:
"""解析参数化键 → (模板, 参数);不是参数化键或参数不合法都抛 FactorError。"""
m = _KEY_RE.match(name.strip())
if not m:
raise FactorError(
f"因子键格式不对:{name};内置因子用注册表名(如 momentum_60),"
"参数化因子用 template(k=v,...)(如 momentum(window=90,direction=higher_is_better))"
)
template_name = m.group("template")
try:
template = get_template(template_name)
except FactorError as exc:
if template_name in _REGISTRY:
raise FactorError(
f"{template_name} 是内置因子实例名,不能在它上面再带参数;"
"要参数化请用模板名,例如 momentum(window=20,direction=higher_is_better)"
) from None
raise exc
raw: dict[str, Any] = {}
args = m.group("args").strip()
if args:
for part in args.split(","):
am = _ARG_RE.match(part.strip())
if not am:
raise FactorError(
f"因子键里的参数写法不对:{part.strip()};应为 名=值(值只能是整数或标识符)"
)
key = am.group("key")
if key in raw:
raise FactorError(f"因子键里参数重复:{key}")
raw[key] = am.group("value")
return template, validate_params(template, raw)
def _derived(name: str) -> tuple[FactorDef, FactorFn]:
template, params = parse_factor_key(name)
key = canonical_key(template, params)
if key != name.strip():
raise FactorError(
f"因子键 {name} 不是规范写法:同样参数请写成 {key}"
"(参数顺序固定、值要写全,避免同一个因子出现多个名字)"
)
defn = build_factor_def(template, params, name=key, source="custom")
return defn, _bind(template.fn)
def resolve_factor(name: str) -> tuple[FactorDef, FactorFn]:
"""按名字取因子(注册表历史名 / 参数化键都行),取不到就抛可读错误。"""
if name in _REGISTRY:
return _REGISTRY[name]
return _derived(name)
def is_resolvable(name: str) -> bool:
try:
resolve_factor(name)
except FactorError:
return False
return True
# 兼容旧名:所有既有调用点(选股/回测/组合/说明书/条件字段/Agent 工具)自动支持参数化键。
get_factor = resolve_factor
def list_factors() -> list[FactorDef]: def list_factors() -> list[FactorDef]:
"""注册表里的**内置实例**(历史名),按名字排序(目录 seed 用)。"""
return [d for d, _fn in sorted(_REGISTRY.values(), key=lambda x: x[0].name)] return [d for d, _fn in sorted(_REGISTRY.values(), key=lambda x: x[0].name)]
def compute_factor(name: str, daily: pd.DataFrame) -> tuple[FactorDef, pd.DataFrame]: def compute_factor(name: str, daily: pd.DataFrame) -> tuple[FactorDef, pd.DataFrame]:
"""计算因子:从行情长表提取所需字段的面板后调用因子函数。""" """计算因子:从行情长表提取所需字段的面板后调用因子函数。"""
defn, fn = get_factor(name) defn, fn = resolve_factor(name)
fields: dict[str, pd.DataFrame] = {} fields: dict[str, pd.DataFrame] = {}
for col in defn.requires: for col in defn.requires:
panel = daily.pivot(index="trade_date", columns="symbol", values=col).sort_index() panel = daily.pivot(index="trade_date", columns="symbol", values=col).sort_index()
panel.index = pd.to_datetime(panel.index) panel.index = pd.to_datetime(panel.index)
fields[col] = panel fields[col] = panel
return defn, fn(fields) return defn, fn(fields, defn.params)
# ---------- 内置因子 ---------- # ---------- 内置因子模板 ----------
# 每个模板的 instances 是历史名 + 它的参数:这些名字已经存在于策略/归档/文档/测试里,
# 必须继续可解析,所以它们不是「参数化键」,而是代码登记的实例。
def _window_spec(label: str = "窗口") -> ParamSpec:
return ParamSpec(
name=P_WINDOW,
label=label,
kind="int",
default=20,
minimum=WINDOW_MIN,
maximum=WINDOW_MAX,
note=f"{WINDOW_MIN}~{WINDOW_MAX} 个交易日;改窗口 = 换一个因子身份(新键),"
"旧键仍按旧参数计算。",
)
def _rolling_return(prices: pd.DataFrame, lookback: int) -> pd.DataFrame: def _rolling_return(prices: pd.DataFrame, lookback: int) -> pd.DataFrame:
@@ -84,181 +489,174 @@ def _rolling_vol(prices: pd.DataFrame, lookback: int) -> pd.DataFrame:
return prices.pct_change().rolling(lookback).std() return prices.pct_change().rolling(lookback).std()
@register( register_template(
FactorDef( FactorTemplate(
"momentum_20", name="momentum",
"过去 20 个交易日收益率", label="动量",
"close / close.shift(20) - 1", description="过去 {window} 个交易日收益率",
brief="短期动量:近一个月强势股延续性较强,适合趋势延续环境(牛市中段);震荡市易追高。", formula="close / close.shift({window}) - 1",
lookback=20, brief="动量:强者延续,适合趋势延续环境;窗口越短越敏感、越长越稳。",
fn=lambda fields, params: _rolling_return(fields["close"], params[P_WINDOW]),
param_specs=(_window_spec(),),
lookback_of=lambda params: params[P_WINDOW],
instances=(
("momentum_20", {P_WINDOW: 20}),
("momentum_60", {P_WINDOW: 60}),
("momentum_120", {P_WINDOW: 120}),
),
) )
) )
def _momentum_20(fields: dict[str, pd.DataFrame]) -> pd.DataFrame:
return _rolling_return(fields["close"], 20)
register_template(
@register( FactorTemplate(
FactorDef( name="volatility",
"momentum_60", label="波动率",
"过去 60 个交易日收益率", description="过去 {window} 个交易日收益率波动率",
"close / close.shift(60) - 1", formula="std(pct_change, {window})",
brief="中期动量:A 股常见有效时段(约 1~3 个月),趋势行情首选;需结合市场阶段判断方向。", brief="低波动防御:近段波动小的股票抗跌,弱市/熊市阶段相对占优(方向越低越好)。",
lookback=60, fn=lambda fields, params: _rolling_vol(fields["close"], params[P_WINDOW]),
param_specs=(_window_spec(),),
direction_default=DIRECTION_LOWER,
lookback_of=lambda params: params[P_WINDOW],
instances=(
("volatility_20", {P_WINDOW: 20}),
("volatility_60", {P_WINDOW: 60}),
),
) )
) )
def _momentum_60(fields: dict[str, pd.DataFrame]) -> pd.DataFrame:
return _rolling_return(fields["close"], 60)
register_template(
@register( FactorTemplate(
FactorDef( name="close_to_high",
"momentum_120", label="接近新高",
"过去 120 个交易日收益率", description="收盘价相对 {window} 日最高价的接近程度",
"close / close.shift(120) - 1", formula="close / rolling_max(high, {window})",
brief="长期动量:反映近半年强势,适合大级别趋势;换手慢、回撤修复慢,弱市慎用。", brief="贴近 n 日高点(接近新高):趋势确认型强势股,常与动量互补;需配合市场热度判断。",
lookback=120, fn=lambda fields, params: fields["close"] / fields["high"].rolling(params[P_WINDOW]).max(),
) param_specs=(_window_spec(),),
)
def _momentum_120(fields: dict[str, pd.DataFrame]) -> pd.DataFrame:
return _rolling_return(fields["close"], 120)
@register(
FactorDef(
"volatility_20",
"过去 20 个交易日收益率波动率",
"std(pct_change, 20)",
brief="低波防御(方向 lower_is_better):近月波动小的股票抗跌,弱市/熊市阶段相对占优。",
lookback=20,
direction="lower_is_better",
)
)
def _volatility_20(fields: dict[str, pd.DataFrame]) -> pd.DataFrame:
return _rolling_vol(fields["close"], 20)
@register(
FactorDef(
"volatility_60",
"过去 60 个交易日收益率波动率",
"std(pct_change, 60)",
brief="低波动(方向 lower_is_better):近一季低波组合长期回测常有超额,是防御型核心因子。",
lookback=60,
direction="lower_is_better",
)
)
def _volatility_60(fields: dict[str, pd.DataFrame]) -> pd.DataFrame:
return _rolling_vol(fields["close"], 60)
@register(
FactorDef(
"close_to_high_60",
"收盘价相对 60 日最高价的接近程度",
"close / rolling_max(high, 60)",
brief="贴近 60 日高点(接近新高):趋势确认型强势股,常与动量互补;需配合市场热度判断。",
lookback=60,
requires=("close", "high"), requires=("close", "high"),
lookback_of=lambda params: params[P_WINDOW],
instances=(("close_to_high_60", {P_WINDOW: 60}),),
) )
) )
def _close_to_high_60(fields: dict[str, pd.DataFrame]) -> pd.DataFrame:
high = fields["high"]
return fields["close"] / high.rolling(60).max()
@register( def _check_fast_slow(params: Mapping[str, Any]) -> str | None:
FactorDef( if params[P_FAST] >= params[P_SLOW]:
"volume_ratio_5_60", return f"快线窗口({params[P_FAST]}) 必须小于慢线窗口({params[P_SLOW]})"
"量比:5 日均量 / 60 日均量", return None
"mean(volume, 5) / mean(volume, 60)",
def _volume_ratio_label(params: Mapping[str, Any]) -> str:
direction = "越高越好" if params[P_DIRECTION] == DIRECTION_HIGHER else "越低越好"
return f"量比({params[P_FAST]}/{params[P_SLOW]} 日,{direction})"
register_template(
FactorTemplate(
name="volume_ratio",
label="量比",
description="量比:{fast} 日均量 / {slow} 日均量",
formula="mean(volume, {fast}) / mean(volume, {slow})",
brief="量比放大提示资金关注(短线活跃型);高换手也伴随更高波动,注意与波动因子搭配。", brief="量比放大提示资金关注(短线活跃型);高换手也伴随更高波动,注意与波动因子搭配。",
lookback=60, fn=lambda fields, params: (
fields["volume"].rolling(params[P_FAST]).mean()
/ fields["volume"].rolling(params[P_SLOW]).mean()
),
param_specs=(
ParamSpec(
name=P_FAST,
label="快线",
kind="int",
default=5,
minimum=WINDOW_MIN,
maximum=WINDOW_MAX,
note="短窗口天数,必须小于慢线。",
),
ParamSpec(
name=P_SLOW,
label="慢线",
kind="int",
default=60,
minimum=WINDOW_MIN,
maximum=WINDOW_MAX,
note="长窗口天数,决定回看长度。",
),
),
requires=("volume",), requires=("volume",),
lookback_of=lambda params: params[P_SLOW],
check=_check_fast_slow,
instances=(("volume_ratio_5_60", {P_FAST: 5, P_SLOW: 60}),),
label_of=_volume_ratio_label,
) )
) )
def _volume_ratio_5_60(fields: dict[str, pd.DataFrame]) -> pd.DataFrame:
vol = fields["volume"]
return vol.rolling(5).mean() / vol.rolling(60).mean()
register_template(
@register( FactorTemplate(
FactorDef( name="ma_bias",
"ma_bias_20", label="均线乖离",
"20 日均线乖离率", description="{window} 日均线乖离率",
"(close - ma(close, 20)) / ma(close, 20)", formula="(close - ma(close, {window})) / ma(close, {window})",
brief="20 日均线乖离:上行趋势中正乖离偏强;乖离过大易回落,需警惕过热。", brief="均线乖离:上行趋势中正乖离偏强;乖离过大易回落,需警惕过热。",
lookback=20, fn=lambda fields, params: (
(fields["close"] - fields["close"].rolling(params[P_WINDOW]).mean())
/ fields["close"].rolling(params[P_WINDOW]).mean()
),
param_specs=(_window_spec(),),
lookback_of=lambda params: params[P_WINDOW],
instances=(("ma_bias_20", {P_WINDOW: 20}),),
) )
) )
def _ma_bias_20(fields: dict[str, pd.DataFrame]) -> pd.DataFrame:
close = fields["close"]
ma = close.rolling(20).mean()
return (close - ma) / ma
register_template(
@register( FactorTemplate(
FactorDef( name="reversal",
"reversal_5", label="短期反转",
"短期反转:过去 5 日收益率取负(越低越接近超跌)", description="短期反转:过去 {window} 日收益率取负",
"-1 * (close / close.shift(5) - 1)", formula="-1 * (close / close.shift({window}) - 1)",
brief="短期反转(方向 higher_is_better):前期跌幅大的超跌反弹机会,适合震荡/修复行情。", brief="短期反转:前期跌幅大的超跌反弹机会,适合震荡/修复行情。",
lookback=5, fn=lambda fields, params: -1.0 * _rolling_return(fields["close"], params[P_WINDOW]),
param_specs=(_window_spec(),),
lookback_of=lambda params: params[P_WINDOW],
instances=(("reversal_5", {P_WINDOW: 5}),),
) )
) )
def _reversal_5(fields: dict[str, pd.DataFrame]) -> pd.DataFrame:
return -1.0 * _rolling_return(fields["close"], 5)
# ---------- 每日指标(daily_basic)因子 ----------
# 数据来源:daily_basic 表(Tushare daily_basic 接口),由 ResearchService / SelectionService
# 装配后并入 daily 长表(见 quant/service.load_basic_df)。requires 里的列名即
# domain.entities.market.DAILY_BASIC_NUMERIC_FIELDS 中的列。
# 特别分红导致的股息率畸高阈值(%):dv_ratio 会因一次性特别分红冲到 30%+, # 特别分红导致的股息率畸高阈值(%):dv_ratio 会因一次性特别分红冲到 30%+,
# 直接用「最高股息率」排序会被这类非经常性事件占满头部(实测 600738 在 2020-01-02 # 直接用「最高股息率」排序会被这类非经常性事件占满头部(实测 600738 在 2020-01-02
# 为 37.2%)。本因子不隐式截断(截断属选股条件,应由用户在 conditions 里显式配置), # 为 37.2%)。本因子不隐式截断(截断属选股条件,应由用户在 conditions 里显式配置)。
# 但把阈值作为常量暴露,供前端/条件模板引用。 # 注意:该常量目前**没有**任何代码引用(曾计划供条件模板引用);要按此上限过滤,
# 请在策略条件里显式配置 dv_ratio <= 30,而不是指望因子内部截断。
DIVIDEND_YIELD_SPECIAL_CAP_PCT = 30.0 DIVIDEND_YIELD_SPECIAL_CAP_PCT = 30.0
register_template(
@register( FactorTemplate(
FactorDef( name="dividend_yield",
"dividend_yield", label="股息率",
"股息率(近 12 个月现金分红 / 总市值 × 100,%)", description="股息率(近 12 个月现金分红 / 总市值 × 100,%)",
"dv_ratio(Tushare daily_basic,逐日时点值)", formula="dv_ratio(Tushare daily_basic,逐日时点值)",
brief=( brief=(
"高股息:熊市/震荡市防御性较强,分红提供现金回报底;" "高股息:熊市/震荡市防御性较强,分红提供现金回报底;"
"需警惕「高股息陷阱」——股息率高常因股价下跌或一次性特别分红," "需警惕「高股息陷阱」——股息率高常因股价下跌或一次性特别分红,"
"建议配合 dv_ratio 上限过滤与盈利质量条件使用。" "建议配合 dv_ratio 上限过滤与盈利质量条件使用。"
), ),
frequency="daily", fn=lambda fields, params: fields["dv_ratio"],
lookback=0, # 时点截面值,无滚动窗口
direction="higher_is_better",
requires=("dv_ratio",), requires=("dv_ratio",),
lookback_of=lambda params: 0, # 时点截面值,无滚动窗口
instances=(("dividend_yield", {P_DIRECTION: DIRECTION_HIGHER}),),
) )
) )
def _dividend_yield(fields: dict[str, pd.DataFrame]) -> pd.DataFrame:
"""股息率面板(index=trade_date, columns=symbol)。
直接取当日 dv_ratio 时点值:该值由数据源按「过去 12 个月现金分红 / 当日总市值」 register_template(
逐日重算,只含已发生事件,按 trade_date <= as_of 取值即无未来函数。 FactorTemplate(
缺失值保持 NaN(由复合分/排序统一 dropna 处理),不做 0 填充 —— 0 会被误读成 name="dividend_yield_ttm",
「股息率为 0 的合格标的」,从而污染横截面排序。 label="股息率 TTM",
""" description="股息率 TTM(近 12 个月滚动现金分红 / 总市值 × 100,%)",
return fields["dv_ratio"] formula="dv_ttm(Tushare daily_basic,逐日时点值)",
brief="同股息率,但口径为 TTM;与 dv_ratio 多数日期取值一致,可作交叉验证。",
fn=lambda fields, params: fields["dv_ttm"],
@register(
FactorDef(
"dividend_yield_ttm",
"股息率 TTM(近 12 个月滚动现金分红 / 总市值 × 100,%)",
"dv_ttm(Tushare daily_basic,逐日时点值)",
brief="同 dividend_yield,但口径为 TTM;与 dv_ratio 多数日期取值一致,可作交叉验证。",
frequency="daily",
lookback=0,
direction="higher_is_better",
requires=("dv_ttm",), requires=("dv_ttm",),
lookback_of=lambda params: 0,
instances=(("dividend_yield_ttm", {P_DIRECTION: DIRECTION_HIGHER}),),
) )
) )
def _dividend_yield_ttm(fields: dict[str, pd.DataFrame]) -> pd.DataFrame:
return fields["dv_ttm"]
+17 -6
View File
@@ -207,6 +207,7 @@ def condition_needed_columns(query) -> set[str]:
except FactorError: except FactorError:
raise ValueError( raise ValueError(
f"条件字段未知:{f}(可用: 行情列/ma20/ma60/已注册因子/" f"条件字段未知:{f}(可用: 行情列/ma20/ma60/已注册因子/"
"参数化因子键(如 momentum(window=90,direction=higher_is_better))/"
"每日指标列(dv_ratio 等)/static.*/fundamental.*)" "每日指标列(dv_ratio 等)/static.*/fundamental.*)"
) from None ) from None
needed.update(defn.requires) needed.update(defn.requires)
@@ -390,7 +391,13 @@ def _field_value(field, sym, statics, tech, financial):
def _compare(left, right, op: str) -> bool: def _compare(left, right, op: str) -> bool:
"""混合比较:None 视为不可用 → 除 ne 外不通过;数值/字符串分别处理。""" """混合比较:None 视为不可用 → 除 ne 外不通过;数值/字符串分别处理。
任何类型不匹配(拿日期字段去比大小、in 的右侧不是列表…)一律返回 False,
**绝不抛异常**:条件是用户可以随手改的输入,一个手滑的字段名不该把选股/回测
打成 500。字段库(quant.condition_fields)会在源头上拒绝不可比较的字段,
这里只是最后一道防线。
"""
if op == "ne": if op == "ne":
return left != right return left != right
if left is None or right is None: if left is None or right is None:
@@ -403,12 +410,16 @@ def _compare(left, right, op: str) -> bool:
# 字符串/其它:支持 eq/ne/in/not_in # 字符串/其它:支持 eq/ne/in/not_in
if op == "eq": if op == "eq":
return left == right return left == right
if op == "in": if op in ("in", "not_in"):
return left in right try:
if op == "not_in": return left in right if op == "in" else left not in right
return left not in right except TypeError: # 右侧不是容器 → 该条件无法求值
return False
if op in ("gt", "gte", "lt", "lte"): if op in ("gt", "gte", "lt", "lte"):
return _num_cmp(left, right, op) # 尝试数值,字符串会 ValueError → False try:
return _num_cmp(left, right, op) # 尝试数值,字符串会 ValueError → False
except (TypeError, ValueError):
return False
return False return False
+23 -5
View File
@@ -30,6 +30,7 @@ from app.domain.entities.market import (
) )
from app.domain.entities.research import ResearchSpec from app.domain.entities.research import ResearchSpec
from app.domain.entities.strategy import SelectionStrategy, StrategyDefinition from app.domain.entities.strategy import SelectionStrategy, StrategyDefinition
from app.quant.condition_fields import get_field
from app.quant.factors import FactorDef, FactorError, get_factor from app.quant.factors import FactorDef, FactorError, get_factor
# 选股策略(SelectionStrategy)不含回测参数,走 _describe_selection_only 专用分支; # 选股策略(SelectionStrategy)不含回测参数,走 _describe_selection_only 专用分支;
@@ -137,10 +138,12 @@ def _describe_selection_only(st, factor_meta) -> StrategyDoc:
" ⚠ 持仓数量 / 持仓天数区间 / 调仓时机 / 起始资金 / 费率 / 复权口径 / 回测区间" " ⚠ 持仓数量 / 持仓天数区间 / 调仓时机 / 起始资金 / 费率 / 复权口径 / 回测区间"
"均不在本策略内 —— 它们在「回测组合」中指定,运行时与公共配置合并。" "均不在本策略内 —— 它们在「回测组合」中指定,运行时与公共配置合并。"
) )
# 步骤文案**不带序号**:前端把它渲染进 <ol>(StrategyDocCard),编号由列表提供;
# 后端再写一遍「1. 2. 3.」会渲染成「1. 1. …」双重编号(2026-10 修正)。
steps = [ steps = [
"1. 按股票池口径筛出候选 universe(市场 / 剔 ST / 上市天数 / 指数成分)。", "按股票池口径筛出候选 universe(市场 / 剔 ST / 上市天数 / 指数成分)。",
"2." + (" 逐条求值过滤条件(AND),剔除不满足者。" if st.conditions else " (未设过滤条件,候选 = universe。)"), "逐条求值过滤条件(AND),剔除不满足者。" if st.conditions else "未设过滤条件,候选 = universe。",
"3. 对剩余股票按上述因子打分并降序排列 → 得到候选排名(TopN 在回测组合里截取)。", "对剩余股票按上述因子打分并降序排列 → 得到候选排名(TopN 在回测组合里截取)。",
] ]
warnings.append( warnings.append(
"本说明只覆盖选股口径;回测的资金/持仓/调仓/成本/区间由「回测组合」+「公共配置」决定," "本说明只覆盖选股口径;回测的资金/持仓/调仓/成本/区间由「回测组合」+「公共配置」决定,"
@@ -187,13 +190,25 @@ def _describe_factors_from_specs(factor_specs, factor_meta, warnings) -> list[di
return out return out
def _unit_suffix(field: str, cond) -> str:
"""字面量条件的**基准单位**后缀(引擎就是按它比较的)。
字段间比较(ref)不加:两侧同一单位,写出来只会误导。字段不在注册表里 → 不加,
宁可少写也不猜单位(猜错就是静默的口径错误)。
"""
if getattr(cond, "ref", None) is not None:
return ""
d = get_field(field)
return f" {d.unit}" if d is not None and d.unit else ""
def _condition_lines(conditions, warnings) -> list[str]: def _condition_lines(conditions, warnings) -> list[str]:
"""把 ConditionSpec 列表渲染成可读行(复用既有字段域校验逻辑)。""" """把 ConditionSpec 列表渲染成可读行(复用既有字段域校验逻辑)。"""
lines: list[str] = [] lines: list[str] = []
for c in conditions: for c in conditions:
op = _OP_TEXT.get(c.op, c.op) op = _OP_TEXT.get(c.op, c.op)
right = f"字段 {c.ref}" if c.ref else f"{c.value}" right = f"字段 {c.ref}" if c.ref else f"{c.value}"
lines.append(f"{c.field} {op} {right}") lines.append(f"{c.field} {op} {right}{_unit_suffix(c.field, c)}")
return lines return lines
@@ -583,12 +598,15 @@ def _render_condition(cond) -> str:
op = _OP_TEXT.get(cond.op, cond.op) op = _OP_TEXT.get(cond.op, cond.op)
if cond.ref is not None: if cond.ref is not None:
right = cond.ref right = cond.ref
suffix = ""
elif cond.op in ("in", "not_in"): elif cond.op in ("in", "not_in"):
items = cond.value if isinstance(cond.value, Sequence) else [cond.value] items = cond.value if isinstance(cond.value, Sequence) else [cond.value]
right = "[" + ", ".join(_fmt_value(v) for v in items) + "]" right = "[" + ", ".join(_fmt_value(v) for v in items) + "]"
suffix = _unit_suffix(cond.field, cond)
else: else:
right = _fmt_value(cond.value) right = _fmt_value(cond.value)
return f"{cond.field} {op} {right}" suffix = _unit_suffix(cond.field, cond)
return f"{cond.field} {op} {right}{suffix}"
def _is_known_field(field: str) -> bool: def _is_known_field(field: str) -> bool:
+22 -1
View File
@@ -105,7 +105,18 @@ class TestV3Tools:
out = _invoke(tools, "inspect_factor", {"name": "momentum_60"}) out = _invoke(tools, "inspect_factor", {"name": "momentum_60"})
assert "momentum_60" in out and "公式" in out and "lookback" in out assert "momentum_60" in out and "公式" in out and "lookback" in out
miss = _invoke(tools, "inspect_factor", {"name": "no_such"}) 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: def test_create_composite_factor(self, tools) -> None:
out = _invoke( out = _invoke(
@@ -117,6 +128,16 @@ class TestV3Tools:
{"name": "x", "factors": "no_such:1"}) {"name": "x", "factors": "no_such:1"})
assert "无法创建" in bad 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: def test_get_backtest_result(self, tools) -> None:
out = _invoke(tools, "get_backtest_result", {"experiment_id": "EXP-D4"}) out = _invoke(tools, "get_backtest_result", {"experiment_id": "EXP-D4"})
assert "总收益" in out and "成交" in out assert "总收益" in out and "成交" in out
+21
View File
@@ -53,6 +53,14 @@ class TestConfigApi:
assert got["price_adjustment"] == "qfq" assert got["price_adjustment"] == "qfq"
assert got["benchmark"] == "000905.SH" 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: class TestCombosApi:
def _make_strategy(self, client: TestClient, name: str) -> str: def _make_strategy(self, client: TestClient, name: str) -> str:
@@ -107,3 +115,16 @@ class TestCombosApi:
resp = client.post("/api/combos/run", json=body) resp = client.post("/api/combos/run", json=body)
assert resp.status_code == 400 assert resp.status_code == 400
assert "不存在" in resp.json()["detail"] 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 pandas as pd
import pytest import pytest
from app.domain.entities.combo import BacktestCombo from app.domain.entities.combo import BacktestCombo
from app.domain.entities.research import CostSpec from app.domain.entities.research import CostSpec
from app.quant.combo_engine import HoldingBandRunner, borda_combine 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()} names2 = {r["name"] for r in client.get("/api/factors").json()}
assert "my_custom_factor" in names2, "只补不删:自定义因子必须保留" assert "my_custom_factor" in names2, "只补不删:自定义因子必须保留"
assert "dividend_yield" 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 from __future__ import annotations
import json
import sqlite3 import sqlite3
from pathlib import Path from pathlib import Path
@@ -38,6 +39,7 @@ def test_upgrade_head_creates_phase1_tables(tmp_path) -> None:
"trading_calendar", "trading_calendar",
"financial_indicator", "financial_indicator",
"sync_log", "sync_log",
"condition_field",
"alembic_version", "alembic_version",
} }
assert expected <= tables assert expected <= tables
@@ -63,3 +65,110 @@ def test_upgrade_head_idempotent(tmp_path) -> None:
cfg = _alembic_config(db_path) cfg = _alembic_config(db_path)
command.upgrade(cfg, "head") command.upgrade(cfg, "head")
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 from __future__ import annotations
import json
from datetime import datetime
import pytest import pytest
from app.api import deps from app.api import deps
from app.domain.entities.strategy import SelectionStrategy from app.domain.entities.strategy import SelectionStrategy
from app.infrastructure.persistence.sqlalchemy.base import Base 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 ( from app.infrastructure.persistence.sqlalchemy.repositories.strategy_impl import (
SqlAlchemyStrategyRepository, SqlAlchemyStrategyRepository,
) )
from app.main import app from app.main import app
from fastapi.testclient import TestClient from fastapi.testclient import TestClient
from pydantic import ValidationError
from sqlalchemy import create_engine from sqlalchemy import create_engine
from sqlalchemy.orm import sessionmaker from sqlalchemy.orm import sessionmaker
@@ -72,6 +77,52 @@ class TestStrategyRepository:
): ):
assert forbidden not in dumped, f"选股策略不应含回测参数字段 {forbidden}" 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() @pytest.fixture()
def client(tmp_path): def client(tmp_path):
@@ -130,3 +181,17 @@ class TestStrategiesApi:
json={"period": ["2024-01-01", "2024-06-01"]}, json={"period": ["2024-01-01", "2024-06-01"]},
) )
assert resp.status_code in (404, 405) 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) 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: def test_universe_step_conditional_and_time_accurate(self) -> None:
"""股票池步骤只写实际生效的过滤,且时点口径必须正确(起始日过滤一次)。 """股票池步骤只写实际生效的过滤,且时点口径必须正确(起始日过滤一次)。