From 8f47b5b603d9817a1fb3e5daad78febce0264bd7 Mon Sep 17 00:00:00 2001 From: Simon Date: Wed, 9 Sep 2026 00:28:16 +0800 Subject: [PATCH] =?UTF-8?q?feat(factor):=20M7.1=20=E5=9B=A0=E5=AD=90?= =?UTF-8?q?=E5=AE=9A=E4=B9=89=E5=85=A5=E5=BA=93=20+=20/api/factors=20?= =?UTF-8?q?=E8=AF=BB=E5=BA=93=EF=BC=88=E7=9B=AE=E5=BD=95=E5=A5=91=E7=BA=A6?= =?UTF-8?q?=E6=BA=90=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 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 通过 --- backend/app/api/deps.py | 9 ++ backend/app/api/factors.py | 30 +++--- .../application/services/factor_catalog.py | 21 ++++ backend/app/domain/entities/factor.py | 39 ++++++++ backend/app/domain/repositories/factor.py | 16 +++ .../b7f2a5e81c33_factor_definition_table.py | 40 ++++++++ .../persistence/sqlalchemy/models/__init__.py | 3 + .../persistence/sqlalchemy/models/factor.py | 28 ++++++ .../sqlalchemy/repositories/factor_impl.py | 78 +++++++++++++++ backend/tests/test_api.py | 17 +++- backend/tests/test_factor_catalog.py | 97 +++++++++++++++++++ 11 files changed, 360 insertions(+), 18 deletions(-) create mode 100644 backend/app/application/services/factor_catalog.py create mode 100644 backend/app/domain/entities/factor.py create mode 100644 backend/app/domain/repositories/factor.py create mode 100644 backend/app/infrastructure/persistence/migrations/versions/b7f2a5e81c33_factor_definition_table.py create mode 100644 backend/app/infrastructure/persistence/sqlalchemy/models/factor.py create mode 100644 backend/app/infrastructure/persistence/sqlalchemy/repositories/factor_impl.py create mode 100644 backend/tests/test_factor_catalog.py diff --git a/backend/app/api/deps.py b/backend/app/api/deps.py index 8b76268..fd52565 100644 --- a/backend/app/api/deps.py +++ b/backend/app/api/deps.py @@ -11,6 +11,7 @@ from fastapi import Depends from sqlalchemy.orm import Session 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.market import ( DailyBarRepository, @@ -18,6 +19,9 @@ from app.domain.repositories.market import ( StockRepository, ) 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 ( SqlAlchemyDailyBarRepository, SqlAlchemyFinancialRepository, @@ -70,12 +74,17 @@ def _selection_repo_factory(session: DbSession) -> SelectionRepository: return SqlAlchemySelectionRepository(session) +def _factor_repo_factory(session: DbSession) -> FactorRepository: + return SqlAlchemyFactorRepository(session) + + StockRepoDep = Annotated[StockRepository, Depends(_stock_repo_factory)] DailyRepoDep = Annotated[DailyBarRepository, Depends(_daily_repo_factory)] EngineDep = Annotated[QuantEngine, Depends(_engine_factory)] ResearchServiceDep = Annotated[ResearchService, Depends(_service_factory)] SelectionServiceDep = Annotated[SelectionService, Depends(_selection_service_factory)] SelectionRepoDep = Annotated[SelectionRepository, Depends(_selection_repo_factory)] +FactorRepoDep = Annotated[FactorRepository, Depends(_factor_repo_factory)] def _job_repo_factory(session: DbSession): diff --git a/backend/app/api/factors.py b/backend/app/api/factors.py index f5d3c22..a653667 100644 --- a/backend/app/api/factors.py +++ b/backend/app/api/factors.py @@ -1,26 +1,22 @@ -"""因子目录 API:/api/factors。""" +"""因子目录 API:/api/factors(M7.1 起读 DB factor_definition)。 + +目录为空时自动从代码注册表 seed(幂等);随后可登记自定义因子元数据。 +响应为 FactorDefinition 实体(含 requires 列表等)。 +""" from __future__ import annotations 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.get("", summary="因子目录(含元数据)") -def list_factor_catalog() -> list[dict]: - return [ - { - "name": d.name, - "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() - ] +@router.get("", summary="因子目录(含元数据,来自 factor_definition 表)") +def list_factor_catalog(factor_repo: FactorRepoDep, session: DbSession) -> list[FactorDefinition]: + if not factor_repo.list(): + seed_registry_factors(factor_repo, session) # 首次:从代码注册表 seed + return factor_repo.list() diff --git a/backend/app/application/services/factor_catalog.py b/backend/app/application/services/factor_catalog.py new file mode 100644 index 0000000..a71a88f --- /dev/null +++ b/backend/app/application/services/factor_catalog.py @@ -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 diff --git a/backend/app/domain/entities/factor.py b/backend/app/domain/entities/factor.py new file mode 100644 index 0000000..9c760d4 --- /dev/null +++ b/backend/app/domain/entities/factor.py @@ -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), + ) diff --git a/backend/app/domain/repositories/factor.py b/backend/app/domain/repositories/factor.py new file mode 100644 index 0000000..3bdd9f9 --- /dev/null +++ b/backend/app/domain/repositories/factor.py @@ -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: ... diff --git a/backend/app/infrastructure/persistence/migrations/versions/b7f2a5e81c33_factor_definition_table.py b/backend/app/infrastructure/persistence/migrations/versions/b7f2a5e81c33_factor_definition_table.py new file mode 100644 index 0000000..6b0068d --- /dev/null +++ b/backend/app/infrastructure/persistence/migrations/versions/b7f2a5e81c33_factor_definition_table.py @@ -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") diff --git a/backend/app/infrastructure/persistence/sqlalchemy/models/__init__.py b/backend/app/infrastructure/persistence/sqlalchemy/models/__init__.py index 9504790..345db4a 100644 --- a/backend/app/infrastructure/persistence/sqlalchemy/models/__init__.py +++ b/backend/app/infrastructure/persistence/sqlalchemy/models/__init__.py @@ -4,6 +4,9 @@ 模型统一继承 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 ExperimentModel, JobModel, diff --git a/backend/app/infrastructure/persistence/sqlalchemy/models/factor.py b/backend/app/infrastructure/persistence/sqlalchemy/models/factor.py new file mode 100644 index 0000000..792f4e1 --- /dev/null +++ b/backend/app/infrastructure/persistence/sqlalchemy/models/factor.py @@ -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) diff --git a/backend/app/infrastructure/persistence/sqlalchemy/repositories/factor_impl.py b/backend/app/infrastructure/persistence/sqlalchemy/repositories/factor_impl.py new file mode 100644 index 0000000..49afc6b --- /dev/null +++ b/backend/app/infrastructure/persistence/sqlalchemy/repositories/factor_impl.py @@ -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 diff --git a/backend/tests/test_api.py b/backend/tests/test_api.py index b8e799b..b02c8fa 100644 --- a/backend/tests/test_api.py +++ b/backend/tests/test_api.py @@ -11,6 +11,7 @@ from datetime import date import pytest from app.api import deps from app.domain.entities.market import Stock +from app.infrastructure.persistence.sqlalchemy.base import Base from app.main import app from app.quant.engine import LocalEngine from app.quant.service import ResearchService @@ -41,7 +42,7 @@ def _mem_stocks() -> list[Stock]: @pytest.fixture() -def client() -> TestClient: +def client(tmp_path) -> TestClient: drifts = {sym: 0.003 - 0.0015 * i for i, sym in enumerate(_SYMS)} daily_df = synthetic_daily(drifts, n=300) bars = bars_dataframe_to_daily_bars(daily_df) @@ -60,6 +61,20 @@ def client() -> TestClient: service = ResearchService(_MemStockRepo(_mem_stocks()), _MemDailyRepo(), LocalEngine()) app.dependency_overrides[deps._stock_repo_factory] = lambda: _MemStockRepo(_mem_stocks()) # 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: yield c app.dependency_overrides.clear() diff --git a/backend/tests/test_factor_catalog.py b/backend/tests/test_factor_catalog.py new file mode 100644 index 0000000..e5e6e39 --- /dev/null +++ b/backend/tests/test_factor_catalog.py @@ -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()}