- 后端业务 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
142 lines
5.0 KiB
Python
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
|