- 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 通过
131 lines
4.8 KiB
Python
131 lines
4.8 KiB
Python
"""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"])
|