384bcc56d7
D1: 删 test_circuit_breaker.py(测已归档 raw_redownload.check_circuit_breaker 死代码,全仓零活跃引用,致 data_platform 套件 collection error) D2: datareader.py read_db_daily/read_index_daily 两处 vnpy_db 硬访问→.get()+清晰报错防崩溃(根治切 dbbardata 读指数列待办) D3: high_limit/low_limit close±10% 兜底是有意设计非 bug(填 NaN 会复活 bullet_trade 误判停牌)—get_price 加 round(.,2) 对齐 get_current_tick 口径;测试期望从 NaN 改兜底估算
171 lines
6.5 KiB
Python
171 lines
6.5 KiB
Python
import sys
|
||
import os
|
||
from pathlib import Path
|
||
|
||
# Add real vnpy source code to sys.path
|
||
_VNPY_SRC = os.path.join(os.path.dirname(__file__), "..", "vnpy_v4.4.0")
|
||
_VNPY_SRC = os.path.abspath(_VNPY_SRC)
|
||
if _VNPY_SRC not in sys.path:
|
||
sys.path.insert(0, _VNPY_SRC)
|
||
|
||
import pandas as pd
|
||
from datetime import datetime, date
|
||
from vnpy.trader.object import BarData
|
||
from vnpy.trader.constant import Exchange, Interval
|
||
from vnpy.trader.setting import SETTINGS
|
||
|
||
def read_parquet_daily(symbol: str, start: str, end: str, cfg, dir_key: str = "daily_dir") -> list[BarData]:
|
||
daily_dir = Path(cfg.data_paths[dir_key])
|
||
start_dt = datetime.strptime(start, "%Y-%m-%d")
|
||
end_dt = datetime.strptime(end, "%Y-%m-%d")
|
||
bars: list[BarData] = []
|
||
prefix = "sh" if guess_exchange(symbol) == Exchange.SSE else "sz"
|
||
for year in range(start_dt.year, end_dt.year + 1):
|
||
f = daily_dir / str(year) / f"{prefix}{symbol}_daily.parquet"
|
||
if not f.exists():
|
||
continue
|
||
df = pd.read_parquet(f)
|
||
for _, row in df.iterrows():
|
||
d = pd.to_datetime(row["date"])
|
||
if start_dt <= d <= end_dt:
|
||
bars.append(_row_to_bar(symbol, row, Interval.DAILY))
|
||
return bars
|
||
|
||
def _row_to_bar(symbol: str, row, interval: Interval) -> BarData:
|
||
return BarData(
|
||
symbol=symbol,
|
||
exchange=guess_exchange(symbol), # Task 4 已改为 guess_exchange
|
||
datetime=pd.to_datetime(row["date"]).to_pydatetime(),
|
||
interval=interval,
|
||
open_price=float(row["open"]),
|
||
high_price=float(row["high"]),
|
||
low_price=float(row["low"]),
|
||
close_price=float(row["close"]),
|
||
volume=float(row["volume"]),
|
||
gateway_name="DATA",
|
||
)
|
||
|
||
|
||
def guess_exchange(symbol: str) -> Exchange:
|
||
"""按代码前缀判断交易所:6/68/5x→SSE,0/3/15x→SZSE"""
|
||
if symbol.startswith(("60", "68", "51", "56", "58")):
|
||
return Exchange.SSE
|
||
if symbol.startswith(("00", "30", "15")):
|
||
return Exchange.SZSE
|
||
return Exchange.SSE
|
||
|
||
|
||
def read_db_daily(symbol: str, start: str, end: str, cfg) -> list[BarData]:
|
||
from vnpy.trader.database import get_database # lazy:避免模块 import 依赖数据库驱动
|
||
# Configure vnpy database SETTINGS before calling get_database()
|
||
SETTINGS["database.name"] = "sqlite"
|
||
vnpy_db = getattr(cfg, "data_paths", {}).get("vnpy_db")
|
||
if not vnpy_db:
|
||
raise RuntimeError("config 缺 data_paths.vnpy_db — 检查 backtest.yaml 初始化")
|
||
SETTINGS["database.database"] = vnpy_db
|
||
|
||
db = get_database()
|
||
start_dt = datetime.strptime(start, "%Y-%m-%d")
|
||
end_dt = datetime.strptime(end, "%Y-%m-%d")
|
||
return db.load_bar_data(
|
||
symbol=symbol,
|
||
exchange=guess_exchange(symbol),
|
||
interval=Interval.DAILY,
|
||
start=start_dt,
|
||
end=end_dt,
|
||
)
|
||
|
||
|
||
def read_parquet_15min(symbol: str, start: str, end: str, cfg, dir_key: str = "minute_15_dir") -> list[BarData]:
|
||
"""读 15min parquet(NAS /volume1/stock/minute_kline/15min/sh{symbol}_15min.parquet)。"""
|
||
minute_dir = Path(cfg.data_paths[dir_key])
|
||
start_dt = datetime.strptime(start, "%Y-%m-%d")
|
||
end_dt = datetime.strptime(end, "%Y-%m-%d")
|
||
prefix = "sh" if guess_exchange(symbol) == Exchange.SSE else "sz"
|
||
f = minute_dir / f"{prefix}{symbol}_15min.parquet"
|
||
if not f.exists():
|
||
return []
|
||
df = pd.read_parquet(f)
|
||
# baostock 15min 有完整时分 datetime 列;旧数据用 date 列
|
||
time_col = ("datetime" if "datetime" in df.columns
|
||
else ("date" if "date" in df.columns else df.columns[0]))
|
||
bars: list[BarData] = []
|
||
for _, row in df.iterrows():
|
||
d = pd.to_datetime(row[time_col])
|
||
if start_dt <= d <= end_dt:
|
||
bars.append(BarData(
|
||
symbol=symbol,
|
||
exchange=guess_exchange(symbol),
|
||
datetime=d.to_pydatetime(),
|
||
interval=Interval.MINUTE,
|
||
open_price=float(row["open"]),
|
||
high_price=float(row["high"]),
|
||
low_price=float(row["low"]),
|
||
close_price=float(row["close"]),
|
||
volume=float(row["volume"]),
|
||
gateway_name="DATA",
|
||
))
|
||
return bars
|
||
|
||
|
||
def read_index_daily(code: str, start, end, cfg) -> pd.DataFrame:
|
||
"""
|
||
读指数日线数据(sh000300/sz399001 等),从 vnpy DB 读取(统一数据源)。
|
||
parquet 仅作原始备份,不再读取。返回类型保持 pd.DataFrame(cta_engine 消费不变)。
|
||
|
||
Args:
|
||
code: 指数代码,带交易所前缀,如 "sh000300"(沪深300)、"sz399001"(深证成指)
|
||
start: 起始日期(str "YYYY-MM-DD" / date / datetime)
|
||
end: 结束日期(str "YYYY-MM-DD" / date / datetime)
|
||
cfg: 数据配置对象
|
||
|
||
Returns:
|
||
pd.DataFrame: date/open/high/low/close/volume 列;无数据返回空 DataFrame。
|
||
"""
|
||
from vnpy.trader.database import get_database # lazy:避免模块 import 依赖数据库驱动
|
||
|
||
# 前缀解析交易所(指数不能用 guess_exchange:000300 以 0 开头会被误判成 SZSE,
|
||
# 但 000300 实际属于 SSE)。sh → SSE,sz → SZSE。
|
||
symbol = code[2:]
|
||
exchange = Exchange.SSE if code.startswith("sh") else Exchange.SZSE
|
||
|
||
# 日期归一化:str → parse, date → combine, datetime → as-is
|
||
def _to_dt(s, is_start: bool) -> datetime:
|
||
if isinstance(s, datetime):
|
||
return s
|
||
if isinstance(s, date):
|
||
return datetime.combine(s, datetime.min.time() if is_start else datetime.max.time())
|
||
return datetime.strptime(s, "%Y-%m-%d")
|
||
|
||
start_dt = _to_dt(start, True)
|
||
end_dt = _to_dt(end, False)
|
||
|
||
# 配置 vnpy DB(与 read_db_daily 同模式)
|
||
SETTINGS["database.name"] = "sqlite"
|
||
vnpy_db = getattr(cfg, "data_paths", {}).get("vnpy_db")
|
||
if not vnpy_db:
|
||
raise RuntimeError("config 缺 data_paths.vnpy_db — 检查 backtest.yaml 初始化")
|
||
SETTINGS["database.database"] = vnpy_db
|
||
|
||
db = get_database()
|
||
bars = db.load_bar_data(
|
||
symbol=symbol,
|
||
exchange=exchange,
|
||
interval=Interval.DAILY,
|
||
start=start_dt,
|
||
end=end_dt,
|
||
)
|
||
|
||
if not bars:
|
||
return pd.DataFrame(columns=["date", "open", "high", "low", "close", "volume"])
|
||
|
||
df = pd.DataFrame([{
|
||
"date": b.datetime,
|
||
"open": b.open_price,
|
||
"high": b.high_price,
|
||
"low": b.low_price,
|
||
"close": b.close_price,
|
||
"volume": b.volume,
|
||
} for b in bars])
|
||
return df.sort_values("date").reset_index(drop=True)
|