mirror of
https://ghfast.top/https://github.com/aeroxw/tick-stock-panel.git
synced 2026-09-12 17:54:15 +08:00
Merge pull request #280 from kevin9327/fix/backtest-stream-date-400
fix(backtest): SSE 端点非法 start/end 返回 400 而不是 500
This commit is contained in:
+18
-12
@@ -527,10 +527,12 @@ async def strategy_stream(
|
||||
from app.backtest.strategy import StrategyBacktestConfig
|
||||
from app.backtest.worker import make_worker_task, run_worker_task
|
||||
|
||||
end_date = date.fromisoformat(end) if end else date.today()
|
||||
if start:
|
||||
start_date = date.fromisoformat(start)
|
||||
else:
|
||||
try:
|
||||
end_date = date.fromisoformat(end) if end else date.today()
|
||||
start_date = date.fromisoformat(start) if start else None
|
||||
except ValueError as e:
|
||||
raise HTTPException(status_code=400, detail=f"日期格式错误: {e}") from e
|
||||
if start_date is None:
|
||||
# 空 start = 全部历史: 用本地最早日K日期, 查不到再回退到默认窗口
|
||||
earliest = request.app.state.repo.earliest_daily_date()
|
||||
start_date = earliest or (end_date - timedelta(days=FACTOR_DEFAULT_DAYS))
|
||||
@@ -830,10 +832,12 @@ async def optimize_stream(
|
||||
from app.backtest.optimizer import OptimizeConfig
|
||||
from app.backtest.worker import make_worker_task, run_worker_task
|
||||
|
||||
end_date = date.fromisoformat(end) if end else date.today()
|
||||
if start:
|
||||
start_date = date.fromisoformat(start)
|
||||
else:
|
||||
try:
|
||||
end_date = date.fromisoformat(end) if end else date.today()
|
||||
start_date = date.fromisoformat(start) if start else None
|
||||
except ValueError as e:
|
||||
raise HTTPException(status_code=400, detail=f"日期格式错误: {e}") from e
|
||||
if start_date is None:
|
||||
earliest = request.app.state.repo.earliest_daily_date()
|
||||
start_date = earliest or (end_date - timedelta(days=FACTOR_DEFAULT_DAYS))
|
||||
|
||||
@@ -1052,10 +1056,12 @@ async def walkforward_stream(
|
||||
|
||||
direction = direction or None
|
||||
|
||||
end_date = date.fromisoformat(end) if end else date.today()
|
||||
if start:
|
||||
start_date = date.fromisoformat(start)
|
||||
else:
|
||||
try:
|
||||
end_date = date.fromisoformat(end) if end else date.today()
|
||||
start_date = date.fromisoformat(start) if start else None
|
||||
except ValueError as e:
|
||||
raise HTTPException(status_code=400, detail=f"日期格式错误: {e}") from e
|
||||
if start_date is None:
|
||||
earliest = request.app.state.repo.earliest_daily_date()
|
||||
start_date = earliest or (end_date - timedelta(days=STRATEGY_DEFAULT_DAYS))
|
||||
|
||||
|
||||
@@ -0,0 +1,69 @@
|
||||
"""回测 SSE 端点的 start/end 入参校验 — 非法日期返回 400 而不是 500。
|
||||
|
||||
signals.py 的 `/intraday/replay` 与 mining.py 的 `MiningRunRequest` 都把
|
||||
`date.fromisoformat` 包在校验里, 非法日期给出 400/422; backtest 的三个 SSE
|
||||
端点直接裸调 `date.fromisoformat`, 同样的入参会抛未捕获 ValueError。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
from fastapi import FastAPI
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from app.api.backtest import router
|
||||
from app.config import settings
|
||||
|
||||
# (路径, 该端点的必填参数)
|
||||
STREAMS = [
|
||||
("/api/backtest/strategy/stream", {"strategy_id": "ma_cross"}),
|
||||
("/api/backtest/optimize/stream", {"strategy_id": "ma_cross", "param_grid": '{"n": [5]}'}),
|
||||
("/api/backtest/walkforward/stream", {"strategy_id": "ma_cross", "param_grid": '{"n": [5]}'}),
|
||||
]
|
||||
STREAM_IDS = [path.rsplit("/", 2)[1] for path, _ in STREAMS]
|
||||
|
||||
BAD_DATES = ["not-a-date", "2026-13-01", "2026/09/04", ""]
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def client() -> TestClient:
|
||||
app = FastAPI()
|
||||
app.include_router(router)
|
||||
# raise_server_exceptions=False: 未捕获异常表现为 500 响应, 与线上行为一致
|
||||
return TestClient(app, raise_server_exceptions=False)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("path,extra", STREAMS, ids=STREAM_IDS)
|
||||
@pytest.mark.parametrize("bad", BAD_DATES)
|
||||
def test_malformed_end_returns_400(client, path, extra, bad):
|
||||
if not bad:
|
||||
pytest.skip("空串走默认值分支, 不属于非法日期")
|
||||
resp = client.get(path, params={**extra, "end": bad})
|
||||
assert resp.status_code == 400, resp.text
|
||||
assert "日期" in resp.json()["detail"]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("path,extra", STREAMS, ids=STREAM_IDS)
|
||||
@pytest.mark.parametrize("bad", BAD_DATES)
|
||||
def test_malformed_start_returns_400(client, path, extra, bad):
|
||||
if not bad:
|
||||
pytest.skip("空串走默认值分支, 不属于非法日期")
|
||||
resp = client.get(path, params={**extra, "start": bad, "end": "2026-09-04"})
|
||||
assert resp.status_code == 400, resp.text
|
||||
assert "日期" in resp.json()["detail"]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("path,extra", STREAMS, ids=STREAM_IDS)
|
||||
def test_valid_dates_still_accepted(client, monkeypatch, path, extra):
|
||||
"""合法日期必须照旧进入事件流, 证明这不是把入口收窄成"什么都不收"。
|
||||
|
||||
打开服务端范围保护并给一个超阈值的窗口, 事件流会立刻以 error 事件收尾,
|
||||
既不触发真实回测, 又能证明日期解析已经放行。
|
||||
"""
|
||||
monkeypatch.setattr(settings, "backtest_range_guard", True)
|
||||
params = {**extra, "start": "2020-01-01", "end": "2026-09-04"}
|
||||
if "walkforward" in path:
|
||||
# walkforward 的 guard 作用于单折窗口, 不是总区间
|
||||
params.update({"train_days": 400, "test_days": 400})
|
||||
resp = client.get(path, params=params)
|
||||
assert resp.status_code == 200, resp.text
|
||||
assert "event: error" in resp.text
|
||||
Reference in New Issue
Block a user