- DailyBarRepository.get_range_many / stream_range_many_columns 增加 adjust 参数 (默认 'none')→ SQL 层过滤口径,消除 stock_daily 混 source/adjust 污染因子的风险 - ResearchSpec / SelectionQuery 增加 price_adjustment(none|qfq),随 config_snapshot 落库可溯源;ResearchService._load_daily 与 SelectionService 装配按口径取数 - tests/test_price_adjustment.py:repo 读取按 adjust 过滤(none/qfq 各自命中)、 spec 默认与字段记录;全量 pytest 通过
157 lines
5.6 KiB
Python
157 lines
5.6 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.infrastructure.persistence.sqlalchemy.base import Base
|
||
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(tmp_path) -> 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, adjust="none"):
|
||
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
|
||
|
||
# /api/factors 自 M7.1 读 DB(factor_definition)→ 提供 tmp sqlite session
|
||
from sqlalchemy import create_engine
|
||
from sqlalchemy.orm import sessionmaker
|
||
|
||
engine = create_engine(f"sqlite:///{tmp_path / 'api.db'}", future=True)
|
||
Base.metadata.create_all(engine)
|
||
SessionLocal = sessionmaker(bind=engine, expire_on_commit=False)
|
||
|
||
def _session_override():
|
||
with SessionLocal() as s:
|
||
yield s
|
||
|
||
app.dependency_overrides[deps.get_session] = _session_override
|
||
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
|