diff --git a/config/data_platform.yaml b/config/data_platform.yaml index d6ea536..4f1eeda 100644 --- a/config/data_platform.yaml +++ b/config/data_platform.yaml @@ -4,6 +4,8 @@ data_paths: raw_dir: /volume1/stock/A股数据/日线数据/raw qfq_dir: /volume1/stock/A股数据/日线数据/qfq minute_15_dir: /volume1/stock/minute_kline/15min + minute_15_qfq_dir: /volume1/stock/minute_kline/15min_qfq + minute_15_raw_dir: /volume1/stock/minute_kline/15min_raw vnpy_db: /volume1/stock/sanguo_vnpy/data/quant_trading.db stock_list: /volume1/stock/A股数据/stock_info/stock_basic_info_raw_20260326_113530.csv diff --git a/sanguo_data/datareader.py b/sanguo_data/datareader.py index cb30152..7d10b78 100644 --- a/sanguo_data/datareader.py +++ b/sanguo_data/datareader.py @@ -83,7 +83,9 @@ def read_parquet_15min(symbol: str, start: str, end: str, cfg, dir_key: str = "m if not f.exists(): return [] df = pd.read_parquet(f) - time_col = "date" if "date" in df.columns else df.columns[0] + # 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]) diff --git a/sanguo_trader/data_source.py b/sanguo_trader/data_source.py index 01c0b25..8ff870c 100644 --- a/sanguo_trader/data_source.py +++ b/sanguo_trader/data_source.py @@ -26,29 +26,32 @@ def _read_fn(interval: str): def _resolve_dir_key(adjust: str, interval: str) -> str: - """adjust → cfg.data_paths 的目录 key。 - - raw 仅日线(raw 15min 待分期项);qfq/默认按 interval 选。 - """ - if adjust == "raw": - if interval != "d": - raise ValueError( - f"raw 模式暂仅支持日线(interval='d'),got '{interval}'" - "(raw 15min 待分期项)" - ) - return "raw_dir" - if adjust == "qfq": - return "qfq_dir" # 干净 qfq(daily_dir 是 mixed,仅 backtest 兼容) - return "daily_dir" if interval == "d" else "minute_15_dir" + """adjust+interval → cfg.data_paths 目录 key。日线/15min 均支持 raw/qfq 双源。""" + if interval == "d": + if adjust == "raw": + return "raw_dir" + if adjust == "qfq": + return "qfq_dir" # 干净 qfq(daily_dir 是 mixed,仅 backtest 兼容) + return "daily_dir" + if interval in ("15m", "15min"): + if adjust == "raw": + return "minute_15_raw_dir" + if adjust == "qfq": + return "minute_15_qfq_dir" + return "minute_15_dir" + raise ValueError(f"不支持的 interval: {interval}(仅 d / 15m)") -def _check_adjust_cfg(adjust: str, cfg) -> None: +def _check_adjust_cfg(adjust: str, cfg, interval: str = "d") -> None: """raw/qfq 需对应 dir 配置,缺失明确报错(不静默 fallback,避免混源)。""" - need = {"raw": "raw_dir", "qfq": "qfq_dir"}.get(adjust) + need = { + ("raw", "d"): "raw_dir", ("qfq", "d"): "qfq_dir", + ("raw", "15m"): "minute_15_raw_dir", ("qfq", "15m"): "minute_15_qfq_dir", + }.get((adjust, interval)) if need and cfg and need not in getattr(cfg, "data_paths", {}): raise ValueError( f"{adjust} 模式需 cfg.data_paths['{need}'](未配置;" - f"先用 raw_redownload.py --adjust {adjust} 生成 parquet)" + f"先用对应下载脚本生成 parquet)" ) @@ -62,7 +65,7 @@ def iter_bars( ) -> Iterator[tuple]: """按日期 cross-section yield (date, {symbol: BarData})。""" dir_key = _resolve_dir_key(adjust, interval) - _check_adjust_cfg(adjust, cfg) + _check_adjust_cfg(adjust, cfg, interval) read_fn = _read_fn(interval) by_date: dict = {} for sym in symbols: @@ -78,7 +81,7 @@ def fetch_day(symbol: str, date: str, interval: str, adjust: str = "qfq", cfg=None): """实走模式拉当日 bar(C-S3 用)。""" dir_key = _resolve_dir_key(adjust, interval) - _check_adjust_cfg(adjust, cfg) + _check_adjust_cfg(adjust, cfg, interval) read_fn = _read_fn(interval) bars = read_fn(symbol, date, date, cfg, dir_key) return bars[-1] if bars else None diff --git a/tests/trader/test_data_source.py b/tests/trader/test_data_source.py index 1a1e853..be0b4eb 100644 --- a/tests/trader/test_data_source.py +++ b/tests/trader/test_data_source.py @@ -65,15 +65,20 @@ def test_iter_bars_raw_missing_dir_raises(monkeypatch): adjust="raw", cfg=_cfg({}))) -def test_iter_bars_raw_15min_unsupported(monkeypatch): - """raw 仅日线;15min+raw 报错(raw 15min 待分期项)。""" +def test_iter_bars_raw_15min_routes_dedicated_dir(monkeypatch): + """raw 15min 路由 minute_15_raw_dir(baostock 双源,分期项已落地)。""" + seen = [] monkeypatch.setattr( "sanguo_trader.data_source._read_fn", - lambda iv: lambda s, st, e, c, dir_key: [], + lambda iv: lambda s, st, e, c, dir_key: seen.append(dir_key) or [], ) - with pytest.raises(ValueError, match="日线"): + list(iter_bars(["600000"], "2024-01-01", "2024-01-31", "15m", + adjust="raw", cfg=_cfg({"minute_15_raw_dir": "/x"}))) + assert seen == ["minute_15_raw_dir"] + # 缺 minute_15_raw_dir 配置 → 明确报错(不 fallback,防混源) + with pytest.raises(ValueError, match="minute_15_raw_dir"): list(iter_bars(["600000"], "2024-01-01", "2024-01-31", "15m", - adjust="raw", cfg=_cfg({"raw_dir": "/x"}))) + adjust="raw", cfg=_cfg({}))) def test_unsupported_interval_rejected():