feat(factor): M7.1 因子定义入库 + /api/factors 读库(目录契约源)

- factor_definition 表(migration b7f2a5e81c33,MySQL 已应用):name 主键 + 元数据
  (formula/brief/frequency/lookback/direction/requires JSON/version)+ FactorDefinition
  entity(from_registry_def 由代码注册表构造)
- FactorRepository Protocol + SQLAlchemy 实现(幂等 upsert/list/get)
- /api/factors 改读 DB;目录为空自动 seed 注册表(幂等)—— 保留自定义因子登记能力
  (计算仍须代码注册,引用未注册因子照常 FactorError,防伪因子)
- tests/test_factor_catalog.py(repo 幂等/roundtrip/registry seed、API seed+字段齐全);
  test_api 的 client fixture 补 tmp sqlite session(factors 读库);全量 pytest 通过
This commit is contained in:
Simon
2026-09-09 00:28:16 +08:00
parent 0ab9038570
commit 8f47b5b603
11 changed files with 360 additions and 18 deletions
+9
View File
@@ -11,6 +11,7 @@ from fastapi import Depends
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
from app.application.services.selection_service import SelectionService from app.application.services.selection_service import SelectionService
from app.domain.repositories.factor import FactorRepository
from app.domain.repositories.jobs import ExperimentRepository, JobRepository from app.domain.repositories.jobs import ExperimentRepository, JobRepository
from app.domain.repositories.market import ( from app.domain.repositories.market import (
DailyBarRepository, DailyBarRepository,
@@ -18,6 +19,9 @@ from app.domain.repositories.market import (
StockRepository, StockRepository,
) )
from app.domain.repositories.selection import SelectionRepository from app.domain.repositories.selection import SelectionRepository
from app.infrastructure.persistence.sqlalchemy.repositories.factor_impl import (
SqlAlchemyFactorRepository,
)
from app.infrastructure.persistence.sqlalchemy.repositories.market_impl import ( from app.infrastructure.persistence.sqlalchemy.repositories.market_impl import (
SqlAlchemyDailyBarRepository, SqlAlchemyDailyBarRepository,
SqlAlchemyFinancialRepository, SqlAlchemyFinancialRepository,
@@ -70,12 +74,17 @@ def _selection_repo_factory(session: DbSession) -> SelectionRepository:
return SqlAlchemySelectionRepository(session) return SqlAlchemySelectionRepository(session)
def _factor_repo_factory(session: DbSession) -> FactorRepository:
return SqlAlchemyFactorRepository(session)
StockRepoDep = Annotated[StockRepository, Depends(_stock_repo_factory)] StockRepoDep = Annotated[StockRepository, Depends(_stock_repo_factory)]
DailyRepoDep = Annotated[DailyBarRepository, Depends(_daily_repo_factory)] DailyRepoDep = Annotated[DailyBarRepository, Depends(_daily_repo_factory)]
EngineDep = Annotated[QuantEngine, Depends(_engine_factory)] EngineDep = Annotated[QuantEngine, Depends(_engine_factory)]
ResearchServiceDep = Annotated[ResearchService, Depends(_service_factory)] 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)]
def _job_repo_factory(session: DbSession): def _job_repo_factory(session: DbSession):
+13 -17
View File
@@ -1,26 +1,22 @@
"""因子目录 API:/api/factors。""" """因子目录 API:/api/factors(M7.1 起读 DB factor_definition)。
目录为空时自动从代码注册表 seed(幂等);随后可登记自定义因子元数据。
响应为 FactorDefinition 实体(含 requires 列表等)。
"""
from __future__ import annotations from __future__ import annotations
from fastapi import APIRouter from fastapi import APIRouter
from app.quant.factors import list_factors from app.api.deps import DbSession, FactorRepoDep
from app.application.services.factor_catalog import seed_registry_factors
from app.domain.entities.factor import FactorDefinition
router = APIRouter(prefix="/factors", tags=["factors"]) router = APIRouter(prefix="/factors", tags=["factors"])
@router.get("", summary="因子目录(含元数据)") @router.get("", summary="因子目录(含元数据,来自 factor_definition 表)")
def list_factor_catalog() -> list[dict]: def list_factor_catalog(factor_repo: FactorRepoDep, session: DbSession) -> list[FactorDefinition]:
return [ if not factor_repo.list():
{ seed_registry_factors(factor_repo, session) # 首次:从代码注册表 seed
"name": d.name, return factor_repo.list()
"description": d.description,
"brief": d.brief,
"formula": d.formula,
"frequency": d.frequency,
"lookback": d.lookback,
"direction": d.direction,
"requires": list(d.requires),
}
for d in list_factors()
]
@@ -0,0 +1,21 @@
"""因子目录用例:把代码注册表因子 seed 进 DB(M7.1)。
DB 为目录契约源;本服务在 /api/factors 首次读取为空时自动 seed(幂等),
后续代码新增因子也通过同一入口同步,保持目录与可计算因子一致。
"""
from __future__ import annotations
from app.domain.entities.factor import FactorDefinition
from app.domain.repositories.factor import FactorRepository
from app.quant.factors import list_factors
def seed_registry_factors(repo: FactorRepository, session) -> int:
"""把 quant/factors 注册表的元数据 upsert 进 factor_definition(幂等)。"""
defs = [FactorDefinition.from_registry_def(d) for d in list_factors()]
if not defs:
return 0
n = repo.upsert_many(defs)
session.commit()
return n
+39
View File
@@ -0,0 +1,39 @@
"""因子目录领域实体(M7.1:因子元数据 DB 化,v2 §11)。
DB 是因子目录的契约源:元数据(含自定义因子登记)入库;
计算执行仍由代码注册表(quant/factors.py)提供 —— 登记但未注册计算的因子
在 score/condition 中引用时仍抛 FactorError(防静默伪因子)。
"""
from __future__ import annotations
from datetime import datetime
from pydantic import BaseModel, Field
class FactorDefinition(BaseModel):
name: str = Field(min_length=1, max_length=64)
description: str = ""
formula: str = ""
brief: str = ""
frequency: str = "daily"
lookback: int = 20
direction: str = Field(default="higher_is_better", pattern="^(higher_is_better|lower_is_better)$")
requires: list[str] = Field(default_factory=list)
version: str = "1"
created_at: datetime | None = None
@classmethod
def from_registry_def(cls, d) -> FactorDefinition:
"""由 quant/factors.FactorDef(dataclass)构造目录实体(seed 用)。"""
return cls(
name=d.name,
description=d.description,
formula=d.formula,
brief=d.brief,
frequency=d.frequency,
lookback=d.lookback,
direction=d.direction,
requires=list(d.requires),
)
+16
View File
@@ -0,0 +1,16 @@
"""因子目录 Repository Protocol(M7.1)。"""
from __future__ import annotations
from typing import Protocol
from app.domain.entities.factor import FactorDefinition
class FactorRepository(Protocol):
def upsert_many(self, definitions: list[FactorDefinition]) -> int:
"""以 name 为幂等键批量写入/更新,返回处理条数。"""
def list(self) -> list[FactorDefinition]: ...
def get(self, name: str) -> FactorDefinition | None: ...
@@ -0,0 +1,40 @@
"""factor_definition 表(M7.1 因子定义入库)
Revision ID: b7f2a5e81c33
Revises: a6c91d4e7f20
Create Date: 2026-09-09
"""
from __future__ import annotations
from collections.abc import Sequence
import sqlalchemy as sa
from alembic import op
revision: str = "b7f2a5e81c33"
down_revision: str | None = "a6c91d4e7f20"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
def upgrade() -> None:
op.create_table(
"factor_definition",
sa.Column("name", sa.String(length=64), nullable=False),
sa.Column("description", sa.String(length=500), nullable=False),
sa.Column("formula", sa.String(length=500), nullable=False),
sa.Column("brief", sa.String(length=500), nullable=False),
sa.Column("frequency", sa.String(length=16), nullable=False),
sa.Column("lookback", sa.Integer(), nullable=False),
sa.Column("direction", sa.String(length=32), nullable=False),
sa.Column("requires_json", sa.Text(), nullable=False),
sa.Column("version", sa.String(length=16), nullable=False),
sa.Column("created_at", sa.DateTime(), nullable=False),
sa.PrimaryKeyConstraint("name"),
)
def downgrade() -> None:
op.drop_table("factor_definition")
@@ -4,6 +4,9 @@
模型统一继承 infra.persistence.sqlalchemy.base.Base。 模型统一继承 infra.persistence.sqlalchemy.base.Base。
""" """
from app.infrastructure.persistence.sqlalchemy.models.factor import ( # noqa: F401
FactorDefinitionModel,
)
from app.infrastructure.persistence.sqlalchemy.models.jobs import ( # noqa: F401 from app.infrastructure.persistence.sqlalchemy.models.jobs import ( # noqa: F401
ExperimentModel, ExperimentModel,
JobModel, JobModel,
@@ -0,0 +1,28 @@
"""因子目录表(M7.1)。
factor_definition:因子元数据契约源(name 主键幂等);requires 以 JSON 存。
"""
from __future__ import annotations
from datetime import datetime
from sqlalchemy import DateTime, Integer, String, Text
from sqlalchemy.orm import Mapped, mapped_column
from app.infrastructure.persistence.sqlalchemy.base import Base
class FactorDefinitionModel(Base):
__tablename__ = "factor_definition"
name: Mapped[str] = mapped_column(String(64), primary_key=True)
description: Mapped[str] = mapped_column(String(500), default="")
formula: Mapped[str] = mapped_column(String(500), default="")
brief: Mapped[str] = mapped_column(String(500), default="")
frequency: Mapped[str] = mapped_column(String(16), default="daily")
lookback: Mapped[int] = mapped_column(Integer, default=20)
direction: Mapped[str] = mapped_column(String(32), default="higher_is_better")
requires_json: Mapped[str] = mapped_column(Text, default="[]")
version: Mapped[str] = mapped_column(String(16), default="1")
created_at: Mapped[datetime] = mapped_column(DateTime)
@@ -0,0 +1,78 @@
"""因子目录 Repository 的 SQLAlchemy 实现(M7.1)。"""
from __future__ import annotations
import json
from datetime import datetime
from sqlalchemy import select
from sqlalchemy.orm import Session
from app.domain.entities.factor import FactorDefinition
from app.infrastructure.persistence.sqlalchemy.models.factor import FactorDefinitionModel
def _to_entity(row: FactorDefinitionModel) -> FactorDefinition:
return FactorDefinition(
name=row.name,
description=row.description,
formula=row.formula,
brief=row.brief,
frequency=row.frequency,
lookback=row.lookback,
direction=row.direction,
requires=json.loads(row.requires_json or "[]"),
version=row.version,
created_at=row.created_at,
)
class SqlAlchemyFactorRepository:
def __init__(self, session: Session) -> None:
self._session = session
def upsert_many(self, definitions: list[FactorDefinition]) -> int:
if not definitions:
return 0
existing = {
r.name: r
for r in self._session.scalars(
select(FactorDefinitionModel).where(
FactorDefinitionModel.name.in_([d.name for d in definitions])
)
)
}
now = datetime.now()
for d in definitions:
row = existing.get(d.name)
if row is None:
self._session.add(
FactorDefinitionModel(
name=d.name,
description=d.description,
formula=d.formula,
brief=d.brief,
frequency=d.frequency,
lookback=d.lookback,
direction=d.direction,
requires_json=json.dumps(d.requires),
version=d.version,
created_at=now,
)
)
else:
for k, v in d.model_dump(exclude={"created_at"}).items():
if k == "requires":
v = json.dumps(v)
setattr(row, k, v)
return len(definitions)
def list(self) -> list[FactorDefinition]:
rows = self._session.scalars(
select(FactorDefinitionModel).order_by(FactorDefinitionModel.name)
).all()
return [_to_entity(r) for r in rows]
def get(self, name: str) -> FactorDefinition | None:
row = self._session.get(FactorDefinitionModel, name)
return _to_entity(row) if row else None
+16 -1
View File
@@ -11,6 +11,7 @@ from datetime import date
import pytest import pytest
from app.api import deps from app.api import deps
from app.domain.entities.market import Stock from app.domain.entities.market import Stock
from app.infrastructure.persistence.sqlalchemy.base import Base
from app.main import app from app.main import app
from app.quant.engine import LocalEngine from app.quant.engine import LocalEngine
from app.quant.service import ResearchService from app.quant.service import ResearchService
@@ -41,7 +42,7 @@ def _mem_stocks() -> list[Stock]:
@pytest.fixture() @pytest.fixture()
def client() -> TestClient: def client(tmp_path) -> TestClient:
drifts = {sym: 0.003 - 0.0015 * i for i, sym in enumerate(_SYMS)} drifts = {sym: 0.003 - 0.0015 * i for i, sym in enumerate(_SYMS)}
daily_df = synthetic_daily(drifts, n=300) daily_df = synthetic_daily(drifts, n=300)
bars = bars_dataframe_to_daily_bars(daily_df) bars = bars_dataframe_to_daily_bars(daily_df)
@@ -60,6 +61,20 @@ def client() -> TestClient:
service = ResearchService(_MemStockRepo(_mem_stocks()), _MemDailyRepo(), LocalEngine()) service = ResearchService(_MemStockRepo(_mem_stocks()), _MemDailyRepo(), LocalEngine())
app.dependency_overrides[deps._stock_repo_factory] = lambda: _MemStockRepo(_mem_stocks()) # noqa: SLF001 app.dependency_overrides[deps._stock_repo_factory] = lambda: _MemStockRepo(_mem_stocks()) # noqa: SLF001
app.dependency_overrides[deps._service_factory] = lambda: service # noqa: SLF001 app.dependency_overrides[deps._service_factory] = lambda: service # noqa: SLF001
# /api/factors 自 M7.1 读 DB(factor_definition)→ 提供 tmp sqlite session
from sqlalchemy import create_engine
from sqlalchemy.orm import sessionmaker
engine = create_engine(f"sqlite:///{tmp_path / 'api.db'}", future=True)
Base.metadata.create_all(engine)
SessionLocal = sessionmaker(bind=engine, expire_on_commit=False)
def _session_override():
with SessionLocal() as s:
yield s
app.dependency_overrides[deps.get_session] = _session_override
with TestClient(app) as c: with TestClient(app) as c:
yield c yield c
app.dependency_overrides.clear() app.dependency_overrides.clear()
+97
View File
@@ -0,0 +1,97 @@
"""M7.1 因子目录测试:factor_definition 落库(幂等 upsert)+ /api/factors 读库 + seed。
repo 测试走 tmp SQLite;API 测试 override get_session 到 tmp sqlite 种子库。
"""
from __future__ import annotations
import pytest
from app.api import deps
from app.domain.entities.factor import FactorDefinition
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.factors import list_factors
from fastapi.testclient import TestClient
from sqlalchemy import create_engine
from sqlalchemy.orm import sessionmaker
@pytest.fixture()
def session(tmp_path):
engine = create_engine(f"sqlite:///{tmp_path / 'factor.db'}", future=True)
Base.metadata.create_all(engine)
Session = sessionmaker(bind=engine, expire_on_commit=False)
with Session() as s:
yield s
class TestFactorRepository:
def test_upsert_idempotent_and_roundtrip(self, session) -> None:
repo = SqlAlchemyFactorRepository(session)
d = FactorDefinition(
name="test_momentum", description="测试", formula="x", brief="b",
lookback=10, requires=["close", "high"],
)
assert repo.upsert_many([d]) == 1
session.commit()
assert len(repo.list()) == 1
got = repo.get("test_momentum")
assert got is not None and got.requires == ["close", "high"]
# 幂等更新
repo.upsert_many([d.model_copy(update={"description": "更新"})])
session.commit()
assert repo.get("test_momentum").description == "更新"
assert len(repo.list()) == 1
def test_seed_from_registry(self, session) -> None:
repo = SqlAlchemyFactorRepository(session)
defs = [FactorDefinition.from_registry_def(d) for d in list_factors()]
assert repo.upsert_many(defs) == len(defs)
session.commit()
names = {f.name for f in repo.list()}
assert len(names) == len(defs)
# 与注册表一致
assert names == {d.name for d in list_factors()}
assert repo.get("momentum_60").direction == "higher_is_better"
@pytest.fixture()
def client(tmp_path):
engine = create_engine(f"sqlite:///{tmp_path / 'api.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 TestFactorsApi:
def test_list_seeds_and_reads_db(self, client) -> None:
resp = client.get("/api/factors")
assert resp.status_code == 200
rows = resp.json()
assert isinstance(rows, list) and len(rows) >= 9
first = next(r for r in rows if r["name"] == "momentum_60")
# DB 契约源字段齐全(与前端 FactorMeta 匹配 + version)
assert set(first.keys()) >= {
"name", "description", "brief", "formula", "frequency",
"lookback", "direction", "requires", "version",
}
assert "close" in first["requires"]
def test_list_matches_registry_after_seed(self, client) -> None:
"""seed 后目录 == 代码注册表集合(无额外未知项)。"""
client.get("/api/factors") # 首次访问触发 seed
client.get("/api/factors") # 幂等:二次访问不报错、不重复
resp = client.get("/api/factors")
names = {r["name"] for r in resp.json()}
assert names == {d.name for d in list_factors()}