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 等),从 dbbardata 读取(2026-08-01 切自 vnpy DbBarData)。 指数点位已由 sina_index_eod 灌入 dbbardata exchange=SSE(16 个中证指数含 000300 沪深300)。 中证指数须用 sh 前缀(点位在 SSE);深证指数 sz 前缀。cta_engine benchmark 用 sh000300。 切 dbbardata 后不再调 get_database() → 消除覆写 vnpy SETTINGS 副作用(cta_engine.py:73/199)。 Args: code: 指数代码,带交易所前缀,如 "sh000300"(沪深300)、"sz399001"(深证成指) start/end: 日期(str "YYYY-MM-DD" / date / datetime) cfg: 数据配置对象(data_paths.vnpy_db 指向含 dbbardata 表的 quant_trading.db) Returns: pd.DataFrame: date/open/high/low/close/volume 列(date 为 Timestamp);无数据返回空 DataFrame。 """ import sqlite3 # 前缀解析交易所(指数不能用 guess_exchange:000300 以 0 开头会误判 SZSE,但属 SSE)。 symbol = code[2:] exchange = "SSE" if code.startswith("sh") else "SZSE" # 日期归一化 -> 'YYYY-MM-DD' def _to_str(s) -> str: if isinstance(s, datetime): return s.strftime("%Y-%m-%d") if isinstance(s, date): return s.strftime("%Y-%m-%d") return str(s)[:10] start_str = _to_str(start) end_str = _to_str(end) db_path = getattr(cfg, "data_paths", {}).get("vnpy_db") if not db_path: raise RuntimeError("config 缺 data_paths.vnpy_db — 检查 backtest.yaml 初始化") conn = sqlite3.connect(db_path, timeout=30) conn.execute("PRAGMA busy_timeout = 30000") try: # substr(datetime,1,10) 比日期规避混合格式(有纯日期有带时间,同 provider 模式) df = pd.read_sql( "SELECT datetime, open_price, high_price, low_price, close_price, volume " "FROM dbbardata WHERE symbol=? AND exchange=? AND interval='d' " "AND substr(datetime,1,10)>=? AND substr(datetime,1,10)<=? ORDER BY datetime", conn, params=(symbol, exchange, start_str, end_str), ) finally: conn.close() if df.empty: return pd.DataFrame(columns=["date", "open", "high", "low", "close", "volume"]) df = df.rename(columns={ "datetime": "date", "open_price": "open", "high_price": "high", "low_price": "low", "close_price": "close", }) df["date"] = pd.to_datetime(df["date"], format="mixed") return df[["date", "open", "high", "low", "close", "volume"]].reset_index(drop=True)