Files
sanguo_vnpy_v2/sanguo_data/datareader.py
T

146 lines
5.5 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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→SSE0/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"
SETTINGS["database.database"] = cfg.data_paths["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 parquetNAS /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: date, end: date, cfg) -> pd.DataFrame:
"""
读指数日线数据(sh000300/sz000905),复用 read_parquet_daily 的年分片 parquet 路径
Args:
code: 指数代码,如 "sh000300"(沪深300)、"sz000905"(中证500
start: 起始日期
end: 结束日期
cfg: 数据配置对象
Returns:
pd.DataFrame: 包含 date/open/high/low/close/volume 列的日线数据
"""
daily_dir = Path(cfg.data_paths["daily_dir"])
start_dt = start if isinstance(start, datetime) else datetime.combine(start, datetime.min.time())
end_dt = end if isinstance(end, datetime) else datetime.combine(end, datetime.max.time())
dfs: list[pd.DataFrame] = []
# 按年分片读取(与 read_parquet_daily 相同路径逻辑)
for year in range(start_dt.year, end_dt.year + 1):
f = daily_dir / str(year) / f"{code}_daily.parquet"
if not f.exists():
continue
df = pd.read_parquet(f)
# 过滤日期范围
df["date"] = pd.to_datetime(df["date"])
mask = (df["date"] >= start_dt) & (df["date"] <= end_dt)
filtered_df = df[mask].copy()
if not filtered_df.empty:
dfs.append(filtered_df)
if dfs:
result = pd.concat(dfs, ignore_index=True)
result = result.sort_values("date")
return result.reset_index(drop=True)
else:
return pd.DataFrame(columns=["date", "open", "high", "low", "close", "volume"])