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
This commit is contained in:
@@ -0,0 +1,51 @@
|
||||
"""API 依赖注入:Repository / 研究服务的装配点(composition root 的一部分)。
|
||||
|
||||
路由层统一使用 Annotated 注入(FastAPI 推荐写法,配合 ruff B008 无冲突)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Annotated
|
||||
|
||||
from fastapi import Depends
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.domain.repositories.market import (
|
||||
DailyBarRepository,
|
||||
StockRepository,
|
||||
)
|
||||
from app.infrastructure.persistence.sqlalchemy.repositories.market_impl import (
|
||||
SqlAlchemyDailyBarRepository,
|
||||
SqlAlchemyStockRepository,
|
||||
)
|
||||
from app.infrastructure.persistence.sqlalchemy.session import get_session
|
||||
from app.quant.engine import LocalEngine, QuantEngine
|
||||
from app.quant.service import ResearchService
|
||||
|
||||
DbSession = Annotated[Session, Depends(get_session)]
|
||||
|
||||
|
||||
def _stock_repo_factory(session: DbSession) -> StockRepository:
|
||||
return SqlAlchemyStockRepository(session)
|
||||
|
||||
|
||||
def _daily_repo_factory(session: DbSession) -> DailyBarRepository:
|
||||
return SqlAlchemyDailyBarRepository(session)
|
||||
|
||||
|
||||
def _engine_factory() -> QuantEngine:
|
||||
return LocalEngine()
|
||||
|
||||
|
||||
def _service_factory(
|
||||
stock_repo: Annotated[StockRepository, Depends(_stock_repo_factory)],
|
||||
daily_repo: Annotated[DailyBarRepository, Depends(_daily_repo_factory)],
|
||||
engine: Annotated[QuantEngine, Depends(_engine_factory)],
|
||||
) -> ResearchService:
|
||||
return ResearchService(stock_repo, daily_repo, engine)
|
||||
|
||||
|
||||
StockRepoDep = Annotated[StockRepository, Depends(_stock_repo_factory)]
|
||||
DailyRepoDep = Annotated[DailyBarRepository, Depends(_daily_repo_factory)]
|
||||
EngineDep = Annotated[QuantEngine, Depends(_engine_factory)]
|
||||
ResearchServiceDep = Annotated[ResearchService, Depends(_service_factory)]
|
||||
@@ -0,0 +1,25 @@
|
||||
"""因子目录 API:/api/factors。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from fastapi import APIRouter
|
||||
|
||||
from app.quant.factors import list_factors
|
||||
|
||||
router = APIRouter(prefix="/factors", tags=["factors"])
|
||||
|
||||
|
||||
@router.get("", summary="因子目录(含元数据)")
|
||||
def list_factor_catalog() -> list[dict]:
|
||||
return [
|
||||
{
|
||||
"name": d.name,
|
||||
"description": d.description,
|
||||
"formula": d.formula,
|
||||
"frequency": d.frequency,
|
||||
"lookback": d.lookback,
|
||||
"direction": d.direction,
|
||||
"requires": list(d.requires),
|
||||
}
|
||||
for d in list_factors()
|
||||
]
|
||||
@@ -0,0 +1,56 @@
|
||||
"""研究执行 API:/api/factor-tests 与 /api/backtests。
|
||||
|
||||
Phase 3 为同步执行(样本有限);Phase 4 将改为 Job + SSE 异步(接口契约不变)。
|
||||
最近一次结果在内存中可读,便于前端展示;持久化实验归档在 Phase 4。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from fastapi import APIRouter, HTTPException
|
||||
|
||||
from app.api.deps import ResearchServiceDep
|
||||
from app.domain.entities.research import (
|
||||
BacktestResult,
|
||||
FactorTestReport,
|
||||
ResearchSpec,
|
||||
)
|
||||
from app.quant.factors import FactorError
|
||||
|
||||
router = APIRouter(tags=["research"])
|
||||
|
||||
# 内存中的最近结果(Phase 4 迁移到 Experiment 表)
|
||||
_LAST_BACKTEST: dict[str, BacktestResult] = {}
|
||||
_LAST_FACTOR_TEST: dict[str, FactorTestReport] = {}
|
||||
|
||||
|
||||
@router.post("/factor-tests", response_model=FactorTestReport, summary="运行单因子测试(同步)")
|
||||
def run_factor_test(
|
||||
spec: ResearchSpec,
|
||||
service: ResearchServiceDep,
|
||||
) -> FactorTestReport:
|
||||
try:
|
||||
report = service.run_factor_test(spec)
|
||||
except (ValueError, FactorError) as exc:
|
||||
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
||||
_LAST_FACTOR_TEST["default"] = report
|
||||
return report
|
||||
|
||||
|
||||
@router.post("/backtests", response_model=BacktestResult, summary="运行回测(同步)")
|
||||
def run_backtest(
|
||||
spec: ResearchSpec,
|
||||
service: ResearchServiceDep,
|
||||
) -> BacktestResult:
|
||||
try:
|
||||
result = service.run_backtest(spec)
|
||||
except (ValueError, FactorError) as exc:
|
||||
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
||||
_LAST_BACKTEST["default"] = result
|
||||
return result
|
||||
|
||||
|
||||
@router.get("/backtests/last", response_model=BacktestResult, summary="最近一次回测结果")
|
||||
def last_backtest() -> BacktestResult:
|
||||
if "default" not in _LAST_BACKTEST:
|
||||
raise HTTPException(status_code=404, detail="尚无回测结果,请先 POST /api/backtests")
|
||||
return _LAST_BACKTEST["default"]
|
||||
@@ -1,14 +1,17 @@
|
||||
"""API 路由聚合。
|
||||
|
||||
后续业务路由按 AGENT.md §17 面向业务对象挂载:
|
||||
/api/stocks /api/universes /api/factors /api/strategies /api/backtests /api/experiments /api/jobs
|
||||
业务路由面向业务对象(AGENT.md §17):/api/stocks /api/factors
|
||||
/api/factor-tests /api/backtests /api/experiments(Phase4) /api/jobs(Phase4) /api/agent(Phase5)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from fastapi import APIRouter
|
||||
|
||||
from app.api import health
|
||||
from app.api import factors, health, research, stocks
|
||||
|
||||
api_router = APIRouter()
|
||||
api_router.include_router(health.router)
|
||||
api_router.include_router(stocks.router)
|
||||
api_router.include_router(factors.router)
|
||||
api_router.include_router(research.router)
|
||||
|
||||
@@ -0,0 +1,34 @@
|
||||
"""股票查询 API:/api/stocks。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from fastapi import APIRouter, HTTPException
|
||||
|
||||
from app.api.deps import StockRepoDep
|
||||
from app.domain.entities.market import Stock
|
||||
|
||||
router = APIRouter(prefix="/stocks", tags=["stocks"])
|
||||
|
||||
|
||||
@router.get("", response_model=list[Stock], summary="股票列表")
|
||||
def list_stocks(
|
||||
repo: StockRepoDep,
|
||||
q: str | None = None,
|
||||
limit: int = 50,
|
||||
offset: int = 0,
|
||||
) -> list[Stock]:
|
||||
if limit > 500:
|
||||
limit = 500
|
||||
stocks = repo.list()
|
||||
if q:
|
||||
needle = q.upper()
|
||||
stocks = [s for s in stocks if needle in s.symbol or needle in s.name.upper()]
|
||||
return stocks[offset : offset + limit]
|
||||
|
||||
|
||||
@router.get("/{symbol}", response_model=Stock, summary="按代码查询")
|
||||
def get_stock(symbol: str, repo: StockRepoDep) -> Stock:
|
||||
stock = repo.get_by_symbol(symbol)
|
||||
if stock is None:
|
||||
raise HTTPException(status_code=404, detail=f"未找到股票 {symbol}")
|
||||
return stock
|
||||
@@ -42,6 +42,9 @@ class DailyBarRepository(Protocol):
|
||||
|
||||
def get_range(self, symbol: str, start: date, end: date) -> list[DailyBar]: ...
|
||||
|
||||
def get_range_many(self, symbols: Sequence[str], start: date, end: date) -> list[DailyBar]:
|
||||
"""批量区间查询(研究服务装配面板用,避免逐只查询)。"""
|
||||
|
||||
def latest_date(self, symbol: str) -> date | None:
|
||||
"""断点续传用:该股票本地已有数据的最新交易日。"""
|
||||
|
||||
|
||||
@@ -140,6 +140,18 @@ class SqlAlchemyDailyBarRepository:
|
||||
).all()
|
||||
return [DailyBar.model_validate(r, from_attributes=True) for r in rows]
|
||||
|
||||
def get_range_many(self, symbols: Sequence[str], start: date, end: date) -> list[DailyBar]:
|
||||
rows = self._session.scalars(
|
||||
select(StockDailyModel)
|
||||
.where(
|
||||
StockDailyModel.symbol.in_(list(symbols)),
|
||||
StockDailyModel.trade_date >= start,
|
||||
StockDailyModel.trade_date <= end,
|
||||
)
|
||||
.order_by(StockDailyModel.trade_date)
|
||||
).all()
|
||||
return [DailyBar.model_validate(r, from_attributes=True) for r in rows]
|
||||
|
||||
def latest_date(self, symbol: str) -> date | None:
|
||||
return self._session.scalar(
|
||||
select(StockDailyModel.trade_date)
|
||||
|
||||
+11
-1
@@ -6,6 +6,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from fastapi import FastAPI
|
||||
from fastapi.middleware.cors import CORSMiddleware
|
||||
|
||||
from app.api.router import api_router
|
||||
from app.core.config import get_settings
|
||||
@@ -15,7 +16,16 @@ settings = get_settings()
|
||||
app = FastAPI(
|
||||
title=settings.app_name,
|
||||
version=settings.app_version,
|
||||
description="A股个人量化研究平台 API(Qlib 引擎 / Tushare 数据源)",
|
||||
description="A股个人量化研究平台 API(研究引擎 / Tushare 数据源)",
|
||||
)
|
||||
|
||||
# 开发期允许本地前端跨域(Phase 3 前端 dev server;上线收紧为白名单)
|
||||
app.add_middleware(
|
||||
CORSMiddleware,
|
||||
allow_origins=["http://localhost:3000", "http://127.0.0.1:3000"],
|
||||
allow_credentials=False,
|
||||
allow_methods=["*"],
|
||||
allow_headers=["*"],
|
||||
)
|
||||
|
||||
app.include_router(api_router, prefix=settings.api_prefix)
|
||||
|
||||
@@ -326,6 +326,7 @@ def run_spec_factor_test(
|
||||
panels = build_factor_panels(daily, spec.factors)
|
||||
panel = panels[0][1]
|
||||
close = daily.pivot(index="trade_date", columns="symbol", values="close").sort_index()
|
||||
close.index = pd.to_datetime(close.index)
|
||||
forward = close.shift(-horizon_days) / close - 1.0
|
||||
report = run_factor_test(panel, forward, factor_name=factor_name)
|
||||
return report, {factor_name: panel}
|
||||
|
||||
@@ -87,7 +87,13 @@ class ResearchService:
|
||||
# 回测前预留因子 warmup(lookback≤120 交易日,取 300 自然日余量)
|
||||
data_start = start - timedelta(days=300)
|
||||
stocks = filter_stocks(self._stock_repo.list(), spec.universe, as_of=start)
|
||||
bars: list = []
|
||||
for s in stocks:
|
||||
bars.extend(self._daily_repo.get_range(s.symbol, data_start, end))
|
||||
if not stocks:
|
||||
return pd.DataFrame()
|
||||
get_many = getattr(self._daily_repo, "get_range_many", None)
|
||||
if get_many is not None:
|
||||
bars = list(get_many([s.symbol for s in stocks], data_start, end))
|
||||
else: # 兜底:逐只查询
|
||||
bars = []
|
||||
for s in stocks:
|
||||
bars.extend(self._daily_repo.get_range(s.symbol, data_start, end))
|
||||
return bars_to_daily_df(bars)
|
||||
|
||||
@@ -0,0 +1,141 @@
|
||||
"""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
|
||||
Reference in New Issue
Block a user