"""M6.3 选股 API 集成测试:POST /api/selections 落库 → GET 读回 → 历史列表。 使用 tmp SQLite + 真实 SQLAlchemy Repository(override get_session), 验证「提交→落库→读回一致」闭环与 v2 §8 历史选股查询。 """ from __future__ import annotations 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.infrastructure.persistence.sqlalchemy.repositories.market_impl import ( SqlAlchemyDailyBarRepository, SqlAlchemyStockRepository, ) from app.main import app from fastapi.testclient import TestClient from sqlalchemy import create_engine from sqlalchemy.orm import sessionmaker from conftest_quant import bars_dataframe_to_daily_bars, synthetic_daily _SYMS = ["60000" + str(i) + ".SH" for i in range(5)] # 600000~600004 _AS_OF = date(2024, 12, 31) @pytest.fixture() def client(tmp_path) -> TestClient: engine = create_engine(f"sqlite:///{tmp_path / 'api.db'}", future=True) Base.metadata.create_all(engine) SessionFactory = sessionmaker(bind=engine, expire_on_commit=False) drifts = {s: 0.006 - 0.0015 * i for i, s in enumerate(_SYMS)} daily_df = synthetic_daily(drifts, n=320) with SessionFactory() as session: stocks = [ Stock(symbol=s, name=f"测试股份{i}", industry="白酒", list_date=date(1999, 1, 1)) for i, s in enumerate(_SYMS) ] SqlAlchemyStockRepository(session).upsert_many(stocks) bars = bars_dataframe_to_daily_bars(daily_df) SqlAlchemyDailyBarRepository(session).upsert_many(bars) session.commit() def _session_override(): with SessionFactory() as session: yield session app.dependency_overrides[deps.get_session] = _session_override with TestClient(app) as c: yield c app.dependency_overrides.clear() _SCORE_BODY = { "universe": {"min_listing_days": 0}, "method": "score", "factors": [{"name": "momentum_60", "weight": 1.0}], "top_n": 3, "as_of": _AS_OF.isoformat(), } class TestSelectionsApi: def test_submit_then_read_back(self, client: TestClient) -> None: resp = client.post("/api/selections", json=_SCORE_BODY) assert resp.status_code == 200 body = resp.json() sel_id = body["selection_id"] assert sel_id.startswith("SEL-") result = body["result"] assert len(result["candidates"]) == 3 assert [c["rank"] for c in result["candidates"]] == [1, 2, 3] assert result["candidates"][0]["factor_values"] # 有因子值 assert result["candidates"][0]["selection_reason"] # 可解释 assert result["as_of_date"] <= _AS_OF.isoformat() # 读回一致 got = client.get(f"/api/selections/{sel_id}") assert got.status_code == 200 g = got.json() assert g["as_of_date"] == result["as_of_date"] assert [c["symbol"] for c in g["candidates"]] == [ c["symbol"] for c in result["candidates"] ] assert g["statistics"]["selected"] == 3 def test_missing_id_404(self, client: TestClient) -> None: assert client.get("/api/selections/SEL-NOPE").status_code == 404 def test_list_and_filter(self, client: TestClient) -> None: client.post("/api/selections", json=_SCORE_BODY) client.post( "/api/selections", json={**_SCORE_BODY, "top_n": 2, "method": "score", "factors": [{"name": "momentum_20", "weight": 1.0}]}, ) rows = client.get("/api/selections").json() assert len(rows) >= 2 assert all(r["id"].startswith("SEL-") for r in rows) assert all(r["method"] == "score" for r in rows) assert all(r["selected"] > 0 for r in rows) # as_of 过滤 by_date = client.get(f"/api/selections?as_of={_AS_OF.isoformat()}").json() assert len(by_date) == len(rows) empty = client.get("/api/selections?as_of=2020-01-01").json() assert empty == [] def test_condition_submit(self, client: TestClient) -> None: resp = client.post( "/api/selections", json={ "universe": {"min_listing_days": 0}, "method": "condition", "conditions": [ {"field": "static.industry", "op": "eq", "value": "白酒"}, {"field": "momentum_60", "op": "gt", "value": 0}, ], "as_of": _AS_OF.isoformat(), }, ) assert resp.status_code == 200 body = resp.json() assert len(body["result"]["candidates"]) == 5 # 全部上涨 → 5 只全过 assert all(c["filter_status"] for c in body["result"]["candidates"])