"""API 集成测试:/api/stocks、/api/factors、/api/backtests、/api/factor-tests。 使用内存 Repository / 合成行情替换真实 DB 依赖(override 装配工厂), 引擎为真实 LocalEngine —— 覆盖「DTO 校验 → 装配 → 引擎 → 标准结果」链路。 """ from __future__ import annotations from datetime import date import pytest from app.api import deps from app.domain.entities.market import Stock from app.main import app from app.quant.engine import LocalEngine from app.quant.service import ResearchService from fastapi.testclient import TestClient from conftest_quant import bars_dataframe_to_daily_bars, synthetic_daily class _MemStockRepo: def __init__(self, stocks: list[Stock]) -> None: self._stocks = stocks def get_by_symbol(self, symbol: str) -> Stock | None: return next((s for s in self._stocks if s.symbol == symbol), None) def list(self) -> list[Stock]: return self._stocks _SYMS = ["60000" + str(i) + ".SH" for i in range(5)] # 600000~600004 def _mem_stocks() -> list[Stock]: return [ Stock(symbol=sym, name=f"测试股份{i}", list_date=date(1999, 11, 10)) for i, sym in enumerate(_SYMS) ] @pytest.fixture() def client() -> TestClient: drifts = {sym: 0.003 - 0.0015 * i for i, sym in enumerate(_SYMS)} daily_df = synthetic_daily(drifts, n=300) bars = bars_dataframe_to_daily_bars(daily_df) class _MemDailyRepo: def get_range_many(self, symbols, start, end): out = [] for b in bars: if b.symbol in symbols and start <= b.trade_date <= end: out.append(b) return out def get_range(self, symbol, start, end): return [b for b in bars if b.symbol == symbol and start <= b.trade_date <= end] service = ResearchService(_MemStockRepo(_mem_stocks()), _MemDailyRepo(), LocalEngine()) app.dependency_overrides[deps._stock_repo_factory] = lambda: _MemStockRepo(_mem_stocks()) # noqa: SLF001 app.dependency_overrides[deps._service_factory] = lambda: service # noqa: SLF001 with TestClient(app) as c: yield c app.dependency_overrides.clear() _BACKTEST_BODY = { "type": "backtest", "universe": {"exclude_st": False, "min_listing_days": 0}, "factors": [{"name": "momentum_20", "weight": 1.0}], "selection": {"top_n": 1}, "rebalance": "monthly", "period": ["2024-03-01", "2024-10-31"], } class TestStocksApi: def test_list(self, client: TestClient) -> None: resp = client.get("/api/stocks?limit=10") assert resp.status_code == 200 body = resp.json() assert len(body) == 5 assert body[0]["symbol"] assert body[0]["name"] def test_list_search(self, client: TestClient) -> None: resp = client.get("/api/stocks?q=600000") assert resp.status_code == 200 assert len(resp.json()) == 1 assert resp.json()[0]["symbol"] == "600000.SH" def test_get_one_and_missing(self, client: TestClient) -> None: assert client.get("/api/stocks/600000.SH").status_code == 200 assert client.get("/api/stocks/999999.SZ").status_code == 404 class TestFactorsApi: def test_catalog(self, client: TestClient) -> None: resp = client.get("/api/factors") assert resp.status_code == 200 names = {f["name"] for f in resp.json()} assert "momentum_20" in names meta = next(f for f in resp.json() if f["name"] == "momentum_60") assert meta["lookback"] == 60 assert meta["direction"] in {"higher_is_better", "lower_is_better"} class TestResearchApi: def test_backtest_roundtrip(self, client: TestClient) -> None: resp = client.post("/api/backtests", json=_BACKTEST_BODY) assert resp.status_code == 200 body = resp.json() assert body["summary"]["total_return_pct"] > 0 assert body["equity_curve"] assert body["unimplemented"] # 最近结果可读 last = client.get("/api/backtests/last") assert last.status_code == 200 assert last.json()["summary"] == body["summary"] def test_factor_test_roundtrip(self, client: TestClient) -> None: body = dict(_BACKTEST_BODY) body["type"] = "factor_test" resp = client.post("/api/factor-tests", json=body) assert resp.status_code == 200 report = resp.json() assert report["factor_name"] == "momentum_20" assert report["sample_days"] > 5 assert report["ic_mean"] > 0 # 合成数据为强趋势 def test_invalid_spec_422(self, client: TestClient) -> None: bad = dict(_BACKTEST_BODY) bad["period"] = ["2024-10-01", "2024-03-01"] # start > end assert client.post("/api/backtests", json=bad).status_code == 422 def test_unknown_factor_400(self, client: TestClient) -> None: bad = dict(_BACKTEST_BODY) bad["factors"] = [{"name": "no_such_factor", "weight": 1.0}] resp = client.post("/api/backtests", json=bad) assert resp.status_code == 400