- 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 通过
98 lines
3.6 KiB
Python
98 lines
3.6 KiB
Python
"""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()}
|