Files
qlib/backend/tests/test_factor_catalog.py
T
Simon 8f47b5b603 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 通过
2026-09-09 00:28:16 +08:00

98 lines
3.6 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""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()}