feat(data): read_parquet_15min + trader data_source(qfq/raw双源)
- datareader: +read_parquet_15min(sh/sz前缀+15min.parquet), get_database lazy(去tzlocal等依赖)
- data_source: iter_bars cross-section yield(date,{symbol:Bar}), raw首版fallback qfq+warning(spec§17)
- 本机 mock _read_fn 测调度逻辑, read_parquet_15min 容器冒烟 4 tests passed.
This commit is contained in:
@@ -12,7 +12,6 @@ import pandas as pd
|
||||
from datetime import datetime
|
||||
from vnpy.trader.object import BarData
|
||||
from vnpy.trader.constant import Exchange, Interval
|
||||
from vnpy.trader.database import get_database
|
||||
from vnpy.trader.setting import SETTINGS
|
||||
|
||||
def read_parquet_daily(symbol: str, start: str, end: str, cfg) -> list[BarData]:
|
||||
@@ -56,6 +55,7 @@ def guess_exchange(symbol: str) -> Exchange:
|
||||
|
||||
|
||||
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"]
|
||||
@@ -70,3 +70,33 @@ def read_db_daily(symbol: str, start: str, end: str, cfg) -> list[BarData]:
|
||||
start=start_dt,
|
||||
end=end_dt,
|
||||
)
|
||||
|
||||
|
||||
def read_parquet_15min(symbol: str, start: str, end: str, cfg) -> list[BarData]:
|
||||
"""读 15min parquet(NAS /volume1/stock/minute_kline/15min/sh{symbol}_15min.parquet)。"""
|
||||
minute_dir = Path(cfg.data_paths["minute_15_dir"])
|
||||
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)
|
||||
time_col = "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
|
||||
|
||||
@@ -0,0 +1,57 @@
|
||||
"""模拟盘行情统一接口(qfq/raw 双源,spec §3.3 / §5)。
|
||||
|
||||
_read_fn 内 lazy import datareader,避免模块级依赖 vnpy 链(tzlocal 等),
|
||||
本机无 vnpy 完整依赖时仍可 import + 单测(mock _read_fn)。
|
||||
"""
|
||||
import logging
|
||||
from collections.abc import Iterator
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _read_fn(interval: str):
|
||||
"""按 interval 返回读取函数(lazy import)。"""
|
||||
if interval == "d":
|
||||
from sanguo_data.datareader import read_parquet_daily
|
||||
return read_parquet_daily
|
||||
if interval == "15m":
|
||||
from sanguo_data.datareader import read_parquet_15min
|
||||
return read_parquet_15min
|
||||
raise ValueError(f"不支持的 interval: {interval}(首版仅 d / 15m)")
|
||||
|
||||
|
||||
def iter_bars(
|
||||
symbols: list[str],
|
||||
start: str,
|
||||
end: str,
|
||||
interval: str,
|
||||
adjust: str = "qfq",
|
||||
cfg=None,
|
||||
) -> Iterator[tuple]:
|
||||
"""按日期 cross-section yield (date, {symbol: BarData})。
|
||||
|
||||
raw 首版 fallback qfq + warning(NAS 暂无 raw parquet,spec §17 开放项)。
|
||||
"""
|
||||
if adjust == "raw":
|
||||
logger.warning(
|
||||
"raw 模式首版 fallback qfq(NAS 暂无 raw parquet,spec §17 开放项)"
|
||||
)
|
||||
read_fn = _read_fn(interval)
|
||||
by_date: dict = {}
|
||||
for sym in symbols:
|
||||
for bar in read_fn(sym, start, end, cfg):
|
||||
dt = bar.datetime
|
||||
key = dt.date() if hasattr(dt, "date") else dt
|
||||
by_date.setdefault(key, {})[sym] = bar
|
||||
for date in sorted(by_date.keys()):
|
||||
yield date, by_date[date]
|
||||
|
||||
|
||||
def fetch_day(symbol: str, date: str, interval: str,
|
||||
adjust: str = "qfq", cfg=None):
|
||||
"""实走模式拉当日 bar(C-S3 用)。"""
|
||||
if adjust == "raw":
|
||||
logger.warning("raw fallback qfq(spec §17)")
|
||||
read_fn = _read_fn(interval)
|
||||
bars = read_fn(symbol, date, date, cfg)
|
||||
return bars[-1] if bars else None
|
||||
@@ -0,0 +1,58 @@
|
||||
"""data_source 测试(mock _read_fn,不依赖 vnpy 链,spec §5/§3.3)。
|
||||
|
||||
read_parquet_15min 的真测试依赖 vnpy BarData + NAS parquet,本机无完整依赖,
|
||||
在容器内冒烟(spec §17);本文件只测 iter_bars 调度逻辑。
|
||||
"""
|
||||
import logging
|
||||
from datetime import datetime
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
|
||||
from sanguo_trader.data_source import iter_bars
|
||||
|
||||
|
||||
def _mock_bar(sym: str, date_str: str, close: float):
|
||||
return SimpleNamespace(
|
||||
symbol=sym,
|
||||
datetime=datetime.strptime(date_str, "%Y-%m-%d"),
|
||||
close_price=close,
|
||||
)
|
||||
|
||||
|
||||
def test_iter_bars_cross_section_multi_symbol(monkeypatch):
|
||||
def mock_read(sym, start, end, cfg):
|
||||
return [_mock_bar(sym, "2024-01-02", 10.5 if sym == "600000" else 15.5)]
|
||||
monkeypatch.setattr("sanguo_trader.data_source._read_fn", lambda iv: mock_read)
|
||||
sections = list(iter_bars(["600000", "000001"], "2024-01-01", "2024-01-31", "d"))
|
||||
assert len(sections) == 1
|
||||
_date, d = sections[0]
|
||||
assert "600000" in d and "000001" in d
|
||||
assert d["600000"].close_price == 10.5
|
||||
|
||||
|
||||
def test_iter_bars_raw_fallback_warning(monkeypatch, caplog):
|
||||
monkeypatch.setattr(
|
||||
"sanguo_trader.data_source._read_fn",
|
||||
lambda iv: lambda s, st, e, c: [],
|
||||
)
|
||||
with caplog.at_level(logging.WARNING):
|
||||
list(iter_bars(["600000"], "2024-01-01", "2024-01-31", "d", adjust="raw"))
|
||||
assert "raw" in caplog.text
|
||||
|
||||
|
||||
def test_unsupported_interval_rejected():
|
||||
with pytest.raises(ValueError):
|
||||
# _read_fn 在生成器首次 next 时才执行,用 list 触发
|
||||
list(iter_bars(["600000"], "2024-01-01", "2024-01-31", "5m"))
|
||||
|
||||
|
||||
def test_fetch_day_returns_last_bar(monkeypatch):
|
||||
from sanguo_trader.data_source import fetch_day
|
||||
monkeypatch.setattr(
|
||||
"sanguo_trader.data_source._read_fn",
|
||||
lambda iv: lambda s, st, e, c: [_mock_bar(s, "2024-01-02", 10.0),
|
||||
_mock_bar(s, "2024-01-02", 10.5)],
|
||||
)
|
||||
bar = fetch_day("600000", "2024-01-02", "d")
|
||||
assert bar.close_price == 10.5 # 取最后一个
|
||||
Reference in New Issue
Block a user