Files
qlib/backend/tests/test_api.py
T
Simon 8f47b5b603 feat(factor): M7.1 因子定义入库 + /api/factors 读库(目录契约源)
- factor_definition 表(migration b7f2a5e81c33,MySQL 已应用):name 主键 + 元数据
  (formula/brief/frequency/lookback/direction/requires JSON/version)+ FactorDefinition
  entity(from_registry_def 由代码注册表构造)
- FactorRepository Protocol + SQLAlchemy 实现(幂等 upsert/list/get)
- /api/factors 改读 DB;目录为空自动 seed 注册表(幂等)—— 保留自定义因子登记能力
  (计算仍须代码注册,引用未注册因子照常 FactorError,防伪因子)
- tests/test_factor_catalog.py(repo 幂等/roundtrip/registry seed、API seed+字段齐全);
  test_api 的 client fixture 补 tmp sqlite session(factors 读库);全量 pytest 通过
2026-09-09 00:28:16 +08:00

157 lines
5.6 KiB
Python
Raw 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):
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