Files
qlib/backend/tests/test_selections_api.py
Simon c60dc78c88 feat(selection): M6.3 选股结果落库 + /api/selections(快照 + 逐候选行)
- domain/repositories/selection.py + SelectionMeta:选股 Repository Protocol
- models/selection.py + migration a6c91d4e7f20:selection_snapshot(查询/统计快照)+
  selection_result(逐候选:rank/score/factor_values/filter_status/reason JSON)
- repositories/selection_impl.py:save/get/list_recent(读回重建 SelectionResult)
- api/selections.py:POST /api/selections(同步执行+落库,返回 selection_id+result)、
  GET /{id}、GET 列表(as_of/method 过滤);deps 装配 SelectionService/Repo
- config._build_mysql_url:仅编码破坏 URL 结构的字符(修 Alembic configparser 遇 %21 崩)
- 迁移已在 MySQL 应用(alembic head a6c91d4e7f20);tests/test_selections_api.py 4 例
  提交→读回一致/404/列表过滤/condition;全量 pytest 通过
2026-09-09 00:20:42 +08:00

131 lines
4.8 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.
"""M6.3 选股 API 集成测试:POST /api/selections 落库 → GET 读回 → 历史列表。
使用 tmp SQLite + 真实 SQLAlchemy Repository(override get_session),
验证「提交→落库→读回一致」闭环与 v2 §8 历史选股查询。
"""
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.infrastructure.persistence.sqlalchemy.repositories.market_impl import (
SqlAlchemyDailyBarRepository,
SqlAlchemyStockRepository,
)
from app.main import app
from fastapi.testclient import TestClient
from sqlalchemy import create_engine
from sqlalchemy.orm import sessionmaker
from conftest_quant import bars_dataframe_to_daily_bars, synthetic_daily
_SYMS = ["60000" + str(i) + ".SH" for i in range(5)] # 600000~600004
_AS_OF = date(2024, 12, 31)
@pytest.fixture()
def client(tmp_path) -> TestClient:
engine = create_engine(f"sqlite:///{tmp_path / 'api.db'}", future=True)
Base.metadata.create_all(engine)
SessionFactory = sessionmaker(bind=engine, expire_on_commit=False)
drifts = {s: 0.006 - 0.0015 * i for i, s in enumerate(_SYMS)}
daily_df = synthetic_daily(drifts, n=320)
with SessionFactory() as session:
stocks = [
Stock(symbol=s, name=f"测试股份{i}", industry="白酒",
list_date=date(1999, 1, 1))
for i, s in enumerate(_SYMS)
]
SqlAlchemyStockRepository(session).upsert_many(stocks)
bars = bars_dataframe_to_daily_bars(daily_df)
SqlAlchemyDailyBarRepository(session).upsert_many(bars)
session.commit()
def _session_override():
with SessionFactory() as session:
yield session
app.dependency_overrides[deps.get_session] = _session_override
with TestClient(app) as c:
yield c
app.dependency_overrides.clear()
_SCORE_BODY = {
"universe": {"min_listing_days": 0},
"method": "score",
"factors": [{"name": "momentum_60", "weight": 1.0}],
"top_n": 3,
"as_of": _AS_OF.isoformat(),
}
class TestSelectionsApi:
def test_submit_then_read_back(self, client: TestClient) -> None:
resp = client.post("/api/selections", json=_SCORE_BODY)
assert resp.status_code == 200
body = resp.json()
sel_id = body["selection_id"]
assert sel_id.startswith("SEL-")
result = body["result"]
assert len(result["candidates"]) == 3
assert [c["rank"] for c in result["candidates"]] == [1, 2, 3]
assert result["candidates"][0]["factor_values"] # 有因子值
assert result["candidates"][0]["selection_reason"] # 可解释
assert result["as_of_date"] <= _AS_OF.isoformat()
# 读回一致
got = client.get(f"/api/selections/{sel_id}")
assert got.status_code == 200
g = got.json()
assert g["as_of_date"] == result["as_of_date"]
assert [c["symbol"] for c in g["candidates"]] == [
c["symbol"] for c in result["candidates"]
]
assert g["statistics"]["selected"] == 3
def test_missing_id_404(self, client: TestClient) -> None:
assert client.get("/api/selections/SEL-NOPE").status_code == 404
def test_list_and_filter(self, client: TestClient) -> None:
client.post("/api/selections", json=_SCORE_BODY)
client.post(
"/api/selections",
json={**_SCORE_BODY, "top_n": 2, "method": "score",
"factors": [{"name": "momentum_20", "weight": 1.0}]},
)
rows = client.get("/api/selections").json()
assert len(rows) >= 2
assert all(r["id"].startswith("SEL-") for r in rows)
assert all(r["method"] == "score" for r in rows)
assert all(r["selected"] > 0 for r in rows)
# as_of 过滤
by_date = client.get(f"/api/selections?as_of={_AS_OF.isoformat()}").json()
assert len(by_date) == len(rows)
empty = client.get("/api/selections?as_of=2020-01-01").json()
assert empty == []
def test_condition_submit(self, client: TestClient) -> None:
resp = client.post(
"/api/selections",
json={
"universe": {"min_listing_days": 0},
"method": "condition",
"conditions": [
{"field": "static.industry", "op": "eq", "value": "白酒"},
{"field": "momentum_60", "op": "gt", "value": 0},
],
"as_of": _AS_OF.isoformat(),
},
)
assert resp.status_code == 200
body = resp.json()
assert len(body["result"]["candidates"]) == 5 # 全部上涨 → 5 只全过
assert all(c["filter_status"] for c in body["result"]["candidates"])