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:
@@ -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
@@ -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
|
||||||
@@ -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),
|
||||||
|
)
|
||||||
@@ -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: ...
|
||||||
+40
@@ -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
|
||||||
@@ -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()
|
||||||
|
|||||||
@@ -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()}
|
||||||
Reference in New Issue
Block a user