110 lines
4.3 KiB
Python
110 lines
4.3 KiB
Python
"""回测参数校验(2026-08-15:此前结束时间选 2029 也能提交,零校验跑垃圾结果)。
|
|
|
|
L1 静态规则零 IO 秒判(格式/先后/未来/区间长度/资金/费率);
|
|
L2 数据可用性查 dbbardata 锚定标的(600000 日线)最新日期,
|
|
模块级缓存 24h——只在提交时服务端调用一次,页面加载零开销;
|
|
查库失败返回 None 自动退化为仅 L1,数据层抖动不挡正常提交。
|
|
"""
|
|
from __future__ import annotations
|
|
|
|
import sqlite3
|
|
import time
|
|
from datetime import date, datetime
|
|
|
|
from fastapi import HTTPException
|
|
|
|
# 锚定标的:600000 浦发银行,上市以来从未长期停牌,是「数据灌到哪天」的可靠探针
|
|
_ANCHOR_SYMBOL = "600000"
|
|
_ANCHOR_EXCHANGE = "SSE"
|
|
_CACHE_TTL_SEC = 24 * 3600
|
|
_MIN_SPAN_DAYS = 30
|
|
_RATE_MAX = 0.01 # 费率上限 1%:填 3 这种是百分数/万分位口径错
|
|
|
|
_latest_cache: tuple[float, str] | None = None # (查询时刻, 最新日期)
|
|
|
|
|
|
def _parse(name: str, value: str) -> date:
|
|
try:
|
|
return datetime.strptime(value, "%Y-%m-%d").date()
|
|
except (ValueError, TypeError):
|
|
raise HTTPException(400, f"{name}格式应为 YYYY-MM-DD: {value!r}")
|
|
|
|
|
|
def validate_backtest_range(start: str, end: str, min_days: int = _MIN_SPAN_DAYS) -> None:
|
|
"""L1+L2 区间校验,非法即抛 400(中文业务提示)。合法返回 None。"""
|
|
s = _parse("开始日期", start)
|
|
e = _parse("结束日期", end)
|
|
if s >= e:
|
|
raise HTTPException(400, f"开始日期({start})必须早于结束日期({end})")
|
|
if (e - s).days < min_days:
|
|
raise HTTPException(400, f"区间过短(不足 {min_days} 天),统计无参考意义,请至少选 {min_days} 天")
|
|
today = date.today().isoformat()
|
|
if end > today:
|
|
raise HTTPException(400, f"结束日期 {end} 在未来,请选择历史日期(今天: {today})")
|
|
latest = get_latest_daily_date()
|
|
if latest and end > latest:
|
|
raise HTTPException(400, f"结束日期 {end} 超出数据范围:日线数据最新到 {latest}")
|
|
|
|
|
|
def validate_capital(name: str, value: float) -> None:
|
|
if value <= 0:
|
|
raise HTTPException(400, f"{name}需大于 0,当前: {value}")
|
|
|
|
|
|
def validate_rate(name: str, value: float) -> None:
|
|
if not (0 <= value <= _RATE_MAX):
|
|
raise HTTPException(400, f"{name}应在 0~{_RATE_MAX:g} 之间(小数,万3=0.0003),当前: {value}")
|
|
|
|
|
|
def validate_cta_request(req) -> None:
|
|
"""CTA 回测(个股/优化共用):区间 + 资金 + 费率。"""
|
|
validate_backtest_range(req.start, req.end)
|
|
validate_capital("初始资金", req.capital)
|
|
validate_rate("佣金率", req.commission_rate)
|
|
validate_rate("印花税率", req.stamp_duty_rate)
|
|
validate_rate("过户费率", req.transfer_fee_rate)
|
|
validate_rate("滑点", req.slippage)
|
|
|
|
|
|
def validate_portfolio_request(req) -> None:
|
|
"""组合回测:区间 + 资金 + 费率(字段名与 CTA 不同:initial_cash/start_date)。"""
|
|
validate_backtest_range(req.start_date, req.end_date)
|
|
validate_capital("初始资金", req.initial_cash)
|
|
validate_rate("佣金率", req.commission_rate)
|
|
validate_rate("印花税率", req.stamp_duty_rate)
|
|
validate_rate("滑点", req.slippage)
|
|
|
|
|
|
def _query_latest_daily_date() -> str | None:
|
|
"""直查 dbbardata 锚定标的日线最大日期;任何异常返回 None(退化为仅 L1)。"""
|
|
try:
|
|
from sanguo_data.config import load_config, find_config_path
|
|
|
|
db_path = load_config(find_config_path()).data_paths.get("vnpy_db")
|
|
if not db_path:
|
|
return None
|
|
conn = sqlite3.connect(db_path, timeout=5)
|
|
try:
|
|
row = conn.execute(
|
|
"SELECT MAX(substr(datetime,1,10)) FROM dbbardata "
|
|
"WHERE symbol=? AND exchange=? AND interval='d'",
|
|
(_ANCHOR_SYMBOL, _ANCHOR_EXCHANGE),
|
|
).fetchone()
|
|
finally:
|
|
conn.close()
|
|
return row[0] if row and row[0] else None
|
|
except Exception:
|
|
return None
|
|
|
|
|
|
def get_latest_daily_date() -> str | None:
|
|
"""数据最新交易日(带 24h 缓存)。日线 18 点后更新,无需每次提交都查。"""
|
|
global _latest_cache
|
|
now = time.time()
|
|
if _latest_cache and now - _latest_cache[0] < _CACHE_TTL_SEC:
|
|
return _latest_cache[1]
|
|
latest = _query_latest_daily_date()
|
|
if latest:
|
|
_latest_cache = (now, latest)
|
|
return latest
|