"""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()}