Files
qlib/backend/tests/test_api.py
T
Simon 92627f5b6b feat: Phase 3 Web — 业务 API(stocks/factors/backtests)+ Next.js 前端
- 后端业务 API:GET /api/stocks(搜索/分页)、GET /api/factors(因子目录)、POST /api/factor-tests 与 /api/backtests(Research Spec 驱动同步执行)、GET /api/backtests/last;Annotated 依赖注入 + CORS(dev)
- Repository 批量查询 get_range_many(研究装配一次查询,避免逐只拉取)
- 前端 frontend/web:Next.js 15(TS) + ECharts —— 总览 / 股票池 / 因子研究(IC·RankIC·分层展示) / 回测(净值·回撤·月度·持仓·未建模标注)
- 前端只消费业务 API 与标准化 BacktestResult,无 Qlib/SQL 概念泄漏
- 真实数据:同步 20 只权重股 2023-2024 日线(9680 根)支撑截面研究
- 验证:API 集成测试 8 项(DTO 校验/装配/引擎/标准结果,内存 repo 全链路)+ 全量 pytest 68 passed;前端 tsc + next build 通过;无头浏览器端到端(factors/backtest 页面渲染后端数据)
- ruff clean
2026-09-06 17:14:33 +08:00

142 lines
5.0 KiB
Python

"""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