Files
Simon ef09d5b419 feat(quant): M7.3 研究行情口径显式化(默认不复权 none,可切 qfq)
- 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 通过
2026-09-09 00:32:55 +08:00

157 lines
5.6 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""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