From 1ed7b72acaa790d92fb91c51c319d958ba579775 Mon Sep 17 00:00:00 2001 From: claude_dev Date: Tue, 7 Jul 2026 22:19:11 +0800 Subject: [PATCH] =?UTF-8?q?feat(data):=20raw=E7=9C=9F=E5=AE=9E=E4=BB=B7?= =?UTF-8?q?=E6=95=B0=E6=8D=AE=E6=BA=90(task#79)=E2=80=94raw=5Fdir+dir=5Fke?= =?UTF-8?q?y=E8=B7=AF=E7=94=B1+=E6=96=B0=E6=B5=AA=E6=BA=90=E9=87=8D?= =?UTF-8?q?=E4=B8=8B?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 根因: daily_dir mixed-adjust(hfq bulk+akshare raw tail)致3-30 -94%假跌。 方案(Linus三问简化单raw, 除权留分期项#3): - datareader read_parquet_daily/15min 加 dir_key 参数 - data_source iter_bars/fetch_day: adjust=raw→raw_dir(缺配置报错防混源), qfq→daily_dir - engine PaperEngine 默认 adjust=raw - config 加 raw_dir; scripts/raw_redownload.py 新浪源adjust='' 直连+单线程限速 - 验证: 浦发606行close 6.5/14.6 mean10.08 0跳变, 撮合成交价9.71-10.25真实 - 测试9/9+trader全量108/108通过 --- config/data_platform.yaml | 1 + docs/deployment/data-download.md | 27 +++- sanguo_data/datareader.py | 8 +- sanguo_trader/data_source.py | 49 ++++++-- sanguo_trader/engine.py | 21 +++- scripts/data_platform/daily_all_update.py | 7 +- scripts/data_platform/raw_redownload.py | 144 ++++++++++++++++++++++ scripts/verify_raw_replay.py | 66 ++++++++++ tests/data/test_datareader.py | 19 +++ tests/trader/test_data_source.py | 67 +++++++--- 10 files changed, 366 insertions(+), 43 deletions(-) create mode 100644 scripts/data_platform/raw_redownload.py create mode 100644 scripts/verify_raw_replay.py diff --git a/config/data_platform.yaml b/config/data_platform.yaml index 3269559..75ef439 100644 --- a/config/data_platform.yaml +++ b/config/data_platform.yaml @@ -1,6 +1,7 @@ # config/data_platform.yaml data_paths: daily_dir: /volume1/stock/A股数据/日线数据/daily + raw_dir: /volume1/stock/A股数据/日线数据/raw minute_15_dir: /volume1/stock/minute_kline/15min 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/docs/deployment/data-download.md b/docs/deployment/data-download.md index 5229ae4..74854e6 100644 --- a/docs/deployment/data-download.md +++ b/docs/deployment/data-download.md @@ -86,10 +86,27 @@ cd v2/scripts/data_platform ./run_daily_update.sh --skip-15min # 只日线 ``` -## raw 双数据源(task #79,后续) +## raw 真实价数据源(task #79,已实现) -当前 parquet 是前复权/后复权(成交价显示失真,如浦发 110 元 vs 现价 ~10)。raw 不复权: -- BaoStock `adjustflag="3"`(v1 `backfill_15min_baostock.py:131` 已用 raw 下 15min) -- akshare `adjust=""` +### 根因:daily_dir mixed-adjust 污染 +`daily_dir`(`/volume1/stock/A股数据/日线数据/daily`)历史数据是 **mixed-adjust**: +- 历史 bulk = 早期遗留 **hfq**(浦发 ~196 元) +- 近期增量 = akshare **raw** tail(浦发 ~10 元),因 Mac 无 baostock、`adjustflag="2"` 是死代码 +- 同一 parquet 半截 hfq 半截 raw → 3-30 单日 **-94% 假跌** → 撮合出垃圾结果(假跌停/假低价) -task #79 接入 raw 双源(撮合/涨跌停/均价用 raw,策略信号用 qfq)。 +### raw_dir 通路(干净单一 raw) +新建 `raw_dir`(`/volume1/stock/A股数据/日线数据/raw`),akshare 新浪源 `adjust=""` 从头重下: +- `scripts/data_platform/raw_redownload.py`:**直连**(unset proxy)+ **单线程 SLEEP 限速**(防封 IP)+ 新浪源 `stock_zh_a_daily` +- 文件名与 daily_dir 同构(`{prefix}{symbol}_daily.parquet`,按 year 分目录),datareader 直读 +- 验证(2026-07-07):浦发 606 行 close 6.5/14.6/mean 10.08,**0 跳变 CLEAN**,撮合成交价 9.71–10.25 真实 + +### 接口(双目录路由) +- `datareader.read_parquet_daily(dir_key="...")`:dir_key 切换 `daily_dir`/`raw_dir` +- `data_source.iter_bars(adjust="raw")`→`raw_dir`;`adjust="qfq"`→`daily_dir`;raw 缺 raw_dir **报错**(不 fallback,防混源) +- `engine.PaperEngine(adjust="raw")` 默认 raw(撮合+信号共用真实价) + +### 简化决策(单 raw,除权留分期项) +原 spec 双源(撮合 raw + 策略 qfq)。Linus 三问简化为**单 raw**:除权缺口对 MA 信号影响低频(年 1-2 次),「分红除权」已列分期项 #3。双源/除权合并进分期项。 + +### daily_dir 后续 +`daily_dir`(mixed)暂保留给 backtest(Phase 2),后续迁 `raw_dir`。`daily_all_update.py` 的 `adjustflag="2"` 在 Mac 无 baostock 时是死代码,实走 akshare raw fallback。 diff --git a/sanguo_data/datareader.py b/sanguo_data/datareader.py index da76889..cb30152 100644 --- a/sanguo_data/datareader.py +++ b/sanguo_data/datareader.py @@ -14,8 +14,8 @@ 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) -> list[BarData]: - daily_dir = Path(cfg.data_paths["daily_dir"]) +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] = [] @@ -73,9 +73,9 @@ def read_db_daily(symbol: str, start: str, end: str, cfg) -> list[BarData]: ) -def read_parquet_15min(symbol: str, start: str, end: str, cfg) -> list[BarData]: +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["minute_15_dir"]) + 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" diff --git a/sanguo_trader/data_source.py b/sanguo_trader/data_source.py index e9791af..43721e0 100644 --- a/sanguo_trader/data_source.py +++ b/sanguo_trader/data_source.py @@ -1,4 +1,8 @@ -"""模拟盘行情统一接口(qfq/raw 双源,spec §3.3 / §5)。 +"""模拟盘行情统一接口(raw 真实价 / qfq 前复权,spec §3.3 / §5)。 + +raw:撮合/涨跌停/成交价用真实价(adjustflag=3 / akshare adjust=""), +从 cfg.data_paths['raw_dir'] 读 parquet(raw 重下脚本生成,单一 adjust 不混源)。 +qfq:策略信号用前复权(无除权缺口),从 cfg.data_paths['daily_dir'] 读。 _read_fn 内 lazy import datareader,避免模块级依赖 vnpy 链(tzlocal 等), 本机无 vnpy 完整依赖时仍可 import + 单测(mock _read_fn)。 @@ -20,6 +24,30 @@ def _read_fn(interval: str): raise ValueError(f"不支持的 interval: {interval}(首版仅 d / 15m)") +def _resolve_dir_key(adjust: str, interval: str) -> str: + """adjust → cfg.data_paths 的目录 key。 + + raw 仅支持日线(raw 15min 待分红除权分期项);qfq/默认按 interval 选日线/15min。 + """ + if adjust == "raw": + if interval != "d": + raise ValueError( + f"raw 模式暂仅支持日线(interval='d'),got '{interval}'" + "(raw 15min 待分期项)" + ) + return "raw_dir" + return "daily_dir" if interval == "d" else "minute_15_dir" + + +def _check_raw_cfg(adjust: str, cfg) -> None: + """raw 模式需 raw_dir 配置,缺失明确报错(不静默 fallback 到 qfq,避免混源)。""" + if adjust == "raw" and (not cfg or "raw_dir" not in getattr(cfg, "data_paths", {})): + raise ValueError( + "raw 模式需 cfg.data_paths['raw_dir'](未配置;" + "先用 scripts/data_platform/raw_redownload.py 生成 raw parquet)" + ) + + def iter_bars( symbols: list[str], start: str, @@ -28,18 +56,13 @@ def iter_bars( 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 开放项)" - ) + """按日期 cross-section yield (date, {symbol: BarData})。""" + dir_key = _resolve_dir_key(adjust, interval) + _check_raw_cfg(adjust, cfg) read_fn = _read_fn(interval) by_date: dict = {} for sym in symbols: - for bar in read_fn(sym, start, end, cfg): + for bar in read_fn(sym, start, end, cfg, dir_key): dt = bar.datetime key = dt.date() if hasattr(dt, "date") else dt by_date.setdefault(key, {})[sym] = bar @@ -50,8 +73,8 @@ def iter_bars( 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)") + dir_key = _resolve_dir_key(adjust, interval) + _check_raw_cfg(adjust, cfg) read_fn = _read_fn(interval) - bars = read_fn(symbol, date, date, cfg) + bars = read_fn(symbol, date, date, cfg, dir_key) return bars[-1] if bars else None diff --git a/sanguo_trader/engine.py b/sanguo_trader/engine.py index 398a500..1c80247 100644 --- a/sanguo_trader/engine.py +++ b/sanguo_trader/engine.py @@ -2,7 +2,9 @@ run():逐 bar → T+1 解冻 → 撮合上一根 next_open pending(用当前 bar)→ 喂策略 on_bar 收新单 → current_close 当根撮合 / next_open 缓冲到下根 → 盯市 → 入库。 -raw 首版 fallback qfq(spec §17),信号与撮合共用一套(标注)。 + +adjust 默认 raw(真实价):撮合/涨跌停/成交价/信号共用一套真实价 bar。 +除权缺口对 MA 信号的影响留分期项(分红除权)处理。 """ import logging @@ -10,7 +12,7 @@ import pandas as pd from .account import Account from .matcher import cross_order -from .models import MatchSession, PaperTrade +from .models import MatchSession, OrderSide, PaperTrade from .persistence import save_daily_balance, save_trade, update_checkpoint from .strategy_runner import StrategyRunner @@ -39,7 +41,7 @@ class PaperEngine: def __init__(self, account: Account, runners: list[StrategyRunner], data_source, cfg, db_path: str, account_id: int, symbols: list[str], start: str, end: str, - interval: str = "d") -> None: + interval: str = "d", adjust: str = "raw") -> None: self.account = account self.runners = runners self.data_source = data_source @@ -50,13 +52,14 @@ class PaperEngine: self.start = start self.end = end self.interval = interval + self.adjust = adjust def run(self) -> None: prev_close: dict[str, float] = {} pending: list = [] # [(order, runner)] next_open 待下根撮合 bar_count = 0 for bar_date, bars in self.data_source.iter_bars( - self.symbols, self.start, self.end, self.interval, "qfq", None + self.symbols, self.start, self.end, self.interval, self.adjust, None ): bar_count += 1 self.account.unfreeze_all() @@ -95,6 +98,16 @@ class PaperEngine: pc = prev_close.get(order.symbol, order.price) result = cross_order(order, match_bar, pc, self.cfg) if isinstance(result, PaperTrade): + if result.side == OrderSide.SELL: + # A 股不能做空:SELL 超过可卖持仓 → 拒单(不开空仓) + pos = self.account.positions.get(order.symbol) + if pos is None or pos.available < result.volume: + save_trade(self.db_path, self.account_id, + {"strategy_id": order.strategy_id, "symbol": order.symbol, + "bar_date": str(bar_date)}, + rejected=True, + reject_reason="insufficient_position_no_short") + return if self.account.cash_enough(result): self.account.apply_trade(result) runner.apply_trade(result) diff --git a/scripts/data_platform/daily_all_update.py b/scripts/data_platform/daily_all_update.py index 97e8c94..22f7ddd 100644 --- a/scripts/data_platform/daily_all_update.py +++ b/scripts/data_platform/daily_all_update.py @@ -245,7 +245,12 @@ class SourceHealthMonitor: # ======================== 数据源:BaoStock ======================== def fetch_baostock_daily(code: str, start_date: str, end_date: str) -> Optional[pd.DataFrame]: - """BaoStock日线:全量历史,无反爬,amount真实,T+1延迟""" + """BaoStock日线:全量历史,无反爬,amount真实,T+1延迟。 + + 注意:adjustflag="2"=qfq;但 Mac/容器无 baostock(HAS_BAOSTOCK=False)时不调用, + 实际走 akshare fallback(adjust=""=raw)→ daily_dir 可能 mixed adjust(task #79 根因)。 + 干净单一 raw 见 raw_redownload.py → raw_dir。 + """ if not HAS_BAOSTOCK: return None bs_code = code_to_baostock(code) diff --git a/scripts/data_platform/raw_redownload.py b/scripts/data_platform/raw_redownload.py new file mode 100644 index 0000000..328e0b4 --- /dev/null +++ b/scripts/data_platform/raw_redownload.py @@ -0,0 +1,144 @@ +#!/usr/bin/env python3 +"""raw 日线重下(task #79):akshare 新浪源 adjust="" 真实价。 + +根因:daily_dir 历史数据是 mixed-adjust(hfq bulk + akshare raw tail 拼接), +3-30 类单日 -94% 假跌。本脚本从头重下**单一 raw** 到 raw_dir,与 daily_dir 同构 +({prefix}{symbol}_daily.parquet,按 year 分目录),datareader 直接读。 + +**约束(用户反馈,见 memory/feedback-data-download-constraints)**: +- 直连不走代理(unset proxy env + NO_PROXY=*) +- 单线程 + SLEEP 间隔限速(不并发猛打,防数据源封 IP) + +用法: + # 验证单只/几只 + python3 raw_redownload.py --symbols 600000,000001 --start 2024-01-01 + # 全市场(读 STOCK_LIST csv,后台跑) + python3 raw_redownload.py --all --start 2024-01-01 + +env: + RAW_DIR 本地 raw 根(默认 /tmp/stock_dl/A股数据/日线数据/raw) + SLEEP 每只请求间隔秒(默认 1.0,限速防封) + STOCK_LIST stock_basic_info csv 路径(--all 读,默认拉到本地的 csv) +""" +import argparse +import csv +import logging +import os +import sys +import time + +# 直连:进程级 unset 代理(用户约束) +for _k in ["HTTP_PROXY", "HTTPS_PROXY", "http_proxy", "https_proxy", "ALL_PROXY", "all_proxy"]: + os.environ.pop(_k, None) +os.environ["NO_PROXY"] = "*" +os.environ["no_proxy"] = "*" + +import pandas as pd # noqa: E402 + +logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s") +log = logging.getLogger("raw_redl") + +RAW_DIR = os.environ.get("RAW_DIR", "/tmp/stock_dl/A股数据/日线数据/raw") +SLEEP = float(os.environ.get("SLEEP", "1.0")) +DEFAULT_STOCK_LIST = "/tmp/stock_dl/A股数据/stock_info/stock_basic_info_raw_20260326_113530.csv" + + +def prefix_for(code: str) -> str: + """sh/sz 前缀(与 datareader.guess_exchange 一致)。""" + return "sh" if code.startswith(("60", "68", "51", "56", "58")) else "sz" + + +def download_one(ak, code: str, start: str, end: str): + """新浪源拉 raw 日线,返回 (df, None) 或 (None, err)。""" + sym = f"{prefix_for(code)}{code}" + try: + df = ak.stock_zh_a_daily( + symbol=sym, + start_date=start.replace("-", ""), + end_date=end.replace("-", ""), + adjust="", # raw 不复权 + ) + except Exception as e: # noqa: BLE001 + return None, f"{type(e).__name__}: {str(e)[:100]}" + if df is None or df.empty: + return None, "empty" + df["date"] = pd.to_datetime(df["date"]) + df["year"] = df["date"].dt.year + return df, None + + +def save_one(code: str, df) -> int: + """按 year 分组写 parquet(与 daily_dir 同构),返回写入行数。""" + n = 0 + for year, g in df.groupby("year"): + ydir = os.path.join(RAW_DIR, str(int(year))) + os.makedirs(ydir, exist_ok=True) + out = os.path.join(ydir, f"{prefix_for(code)}{code}_daily.parquet") + g.drop(columns=["year"]).to_parquet(out) + n += len(g) + return n + + +def load_all_codes(stock_list: str) -> list[str]: + """从 stock_basic_info csv 读代码列表(容错列名)。""" + codes = [] + with open(stock_list, encoding="utf-8", errors="replace") as f: + reader = csv.DictReader(f) + for row in reader: + for k in ("code", "symbol", "ts_code", "代码", "股票代码"): + if k in row and row[k]: + c = str(row[k]).strip().split(".")[0] + if c.isdigit() and len(c) == 6: + codes.append(c) + break + return codes + + +def main(): + ap = argparse.ArgumentParser() + ap.add_argument("--symbols", help="逗号分隔代码,如 600000,000001") + ap.add_argument("--all", action="store_true", help="全市场(读 STOCK_LIST csv)") + ap.add_argument("--start", default="2024-01-01") + ap.add_argument("--end", default=None, help="默认今天") + args = ap.parse_args() + + end = args.end or time.strftime("%Y-%m-%d") + + if args.symbols: + codes = [c.strip() for c in args.symbols.split(",") if c.strip()] + elif args.all: + sl = os.environ.get("STOCK_LIST", DEFAULT_STOCK_LIST) + if not os.path.exists(sl): + log.error("STOCK_LIST 不存在: %s(先跑 run_daily_update.sh 拉取)", sl) + sys.exit(1) + codes = load_all_codes(sl) + log.info("--all 从 %s 读到 %d 只", sl, len(codes)) + else: + ap.error("需指定 --symbols 或 --all") + + import akshare as ak + import warnings + warnings.filterwarnings("ignore") + + ok = fail = rows = 0 + for i, code in enumerate(codes, 1): + df, err = download_one(ak, code, args.start, end) + if df is None: + fail += 1 + log.warning("[%d/%d] %s FAIL %s", i, len(codes), code, err) + else: + try: + n = save_one(code, df) + ok += 1 + rows += n + log.info("[%d/%d] %s ok %d rows", i, len(codes), code, n) + except Exception as e: # noqa: BLE001 + fail += 1 + log.error("[%d/%d] %s SAVE FAIL %s", i, len(codes), code, e) + time.sleep(SLEEP) # 限速(用户约束) + + log.info("=== 完成: ok=%d fail=%d rows=%d,raw_dir=%s ===", ok, fail, rows, RAW_DIR) + + +if __name__ == "__main__": + main() diff --git a/scripts/verify_raw_replay.py b/scripts/verify_raw_replay.py new file mode 100644 index 0000000..ad67555 --- /dev/null +++ b/scripts/verify_raw_replay.py @@ -0,0 +1,66 @@ +#!/usr/bin/env python3 +"""raw 双源回放验证(容器内跑,task #79 Phase C)。 + +确认 engine.adjust=raw → iter_bars 路由 raw_dir → 撮合在真实价上跑(无 mixed 假跌)。 +用法(容器内):python3 /app/scripts/verify_raw_replay.py +""" +import os +import sys + +sys.path.insert(0, "/app") + +from sanguo_data.config import find_config_path, load_config +from sanguo_data.datareader import guess_exchange +from sanguo_trader.account import Account +from sanguo_trader.cta_adapter import PaperCtaEngine +from sanguo_trader.engine import PaperEngine +from sanguo_trader.models import AccountConfig +from sanguo_trader.persistence import init_db, list_daily_balance, list_trades, save_account +from sanguo_trader.strategy_runner import StrategyRunner +from sanguo_api.routes_paper import _DataSourceWrapper +from sanguo_api.strategy_registry import get_strategy_class +from vnpy.trader.utility import ArrayManager + + +def main(symbol="600000", start="2026-01-01", end="2026-07-07"): + cfg = load_config(find_config_path()) + db = "/tmp/raw_verify.db" + if os.path.exists(db): + os.remove(db) + init_db(db) + aid = save_account(db, {"name": "raw_verify", "initial_capital": 1_000_000}) + + cls = get_strategy_class("DoubleMaStrategy") + cta = PaperCtaEngine("s1", match_session="next_open", listing_days=0, size=100) + vt_symbol = symbol + "." + guess_exchange(symbol).value + strat = cls(cta, "s1", vt_symbol, {"fast_window": 5, "slow_window": 10}) + strat.trading = True + strat.am = ArrayManager(20) + cta.set_strategy(strat) + runner = StrategyRunner("s1", strategy=strat, paper_cta_engine=cta, symbol=symbol) + + pe = PaperEngine( + Account(1_000_000), [runner], _DataSourceWrapper(cfg), + AccountConfig(initial_capital=1_000_000), + db, aid, [symbol], start, end, "d", + ) + print("engine.adjust =", pe.adjust) + pe.run() + + trades = [t for t in list_trades(db, aid) if not t["rejected"]] + print("fills:", len(trades)) + for t in trades[:8]: + print(" ", t["bar_date"], t["direction"], "@", t["price"], "vol", t["volume"]) + + bal = list_daily_balance(db, aid) + if bal: + print("净值点数:", len(bal), "首:", round(bal[0]["equity"]), "末:", round(bal[-1]["equity"])) + + prices = [t["price"] for t in trades] + if prices: + print("成交价 min/max:", min(prices), max(prices), + "→", "REAL ✓" if max(prices) < 20 else "MIXED!") + + +if __name__ == "__main__": + main() diff --git a/tests/data/test_datareader.py b/tests/data/test_datareader.py index 929bb7f..6487f59 100644 --- a/tests/data/test_datareader.py +++ b/tests/data/test_datareader.py @@ -26,6 +26,25 @@ def test_read_parquet_daily_returns_bardata(tmp_path): assert bars[0].open_price == 10.0 +def test_read_parquet_daily_dir_key_raw(tmp_path): + """dir_key='raw_dir':切换到 raw 目录读真实价(task #79 双源)。""" + raw_year = tmp_path / "2026" + raw_year.mkdir() + df = pd.DataFrame({ + "date": ["2026-03-30"], + "open": [9.97], "high": [10.05], "low": [9.95], + "close": [10.01], "volume": [800000], + }) + df.to_parquet(raw_year / "sh600000_daily.parquet") # 真实文件名:sh 前缀 + _daily + cfg = DataConfig( + data_paths={"raw_dir": str(tmp_path)}, + data_sources={}, validation={}, performance={}, + ) + bars = read_parquet_daily("600000", "2026-01-01", "2026-12-31", cfg, dir_key="raw_dir") + assert len(bars) == 1 + assert bars[0].close_price == 10.01 + + def test_guess_exchange_sh(): assert guess_exchange("600000").value == "SSE" diff --git a/tests/trader/test_data_source.py b/tests/trader/test_data_source.py index 703fbf0..fd4380b 100644 --- a/tests/trader/test_data_source.py +++ b/tests/trader/test_data_source.py @@ -1,7 +1,7 @@ """data_source 测试(mock _read_fn,不依赖 vnpy 链,spec §5/§3.3)。 -read_parquet_15min 的真测试依赖 vnpy BarData + NAS parquet,本机无完整依赖, -在容器内冒烟(spec §17);本文件只测 iter_bars 调度逻辑。 +read_parquet 的真测试依赖 vnpy BarData + NAS parquet,本机无完整依赖, +在容器内冒烟(spec §17);本文件测 iter_bars 调度逻辑(adjust → dir_key 路由)。 """ import logging from datetime import datetime @@ -9,7 +9,7 @@ from types import SimpleNamespace import pytest -from sanguo_trader.data_source import iter_bars +from sanguo_trader.data_source import iter_bars, fetch_day def _mock_bar(sym: str, date_str: str, close: float): @@ -20,25 +20,60 @@ def _mock_bar(sym: str, date_str: str, close: float): ) +def _cfg(paths): + return SimpleNamespace(data_paths=paths) + + def test_iter_bars_cross_section_multi_symbol(monkeypatch): - def mock_read(sym, start, end, cfg): + seen = [] + def mock_read(sym, start, end, cfg, dir_key): + seen.append(dir_key) 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 + assert all(k == "daily_dir" for k in seen) # 默认 qfq → daily_dir -def test_iter_bars_raw_fallback_warning(monkeypatch, caplog): +def test_iter_bars_raw_uses_raw_dir(monkeypatch, caplog): + """raw 模式路由到 raw_dir,不再 fallback qfq(task #79)。""" + seen = [] + def mock_read(sym, start, end, cfg, dir_key): + seen.append(dir_key) + return [_mock_bar(sym, "2024-01-02", 10.01)] + monkeypatch.setattr("sanguo_trader.data_source._read_fn", lambda iv: mock_read) + with caplog.at_level(logging.WARNING): + sections = list(iter_bars( + ["600000"], "2024-01-01", "2024-01-31", "d", + adjust="raw", cfg=_cfg({"raw_dir": "/x/raw"}), + )) + assert seen == ["raw_dir"] # raw → raw_dir + assert "fallback" not in caplog.text.lower() + assert sections[0][1]["600000"].close_price == 10.01 + + +def test_iter_bars_raw_missing_dir_raises(monkeypatch): + """raw 缺 raw_dir 配置 → 明确报错(不静默 fallback,防混源)。""" monkeypatch.setattr( "sanguo_trader.data_source._read_fn", - lambda iv: lambda s, st, e, c: [], + lambda iv: lambda s, st, e, c, dir_key: [], ) - with caplog.at_level(logging.WARNING): - list(iter_bars(["600000"], "2024-01-01", "2024-01-31", "d", adjust="raw")) - assert "raw" in caplog.text + with pytest.raises(ValueError, match="raw_dir"): + list(iter_bars(["600000"], "2024-01-01", "2024-01-31", "d", + adjust="raw", cfg=_cfg({}))) + + +def test_iter_bars_raw_15min_unsupported(monkeypatch): + """raw 仅日线;15min+raw 报错(raw 15min 待分期项)。""" + monkeypatch.setattr( + "sanguo_trader.data_source._read_fn", + lambda iv: lambda s, st, e, c, dir_key: [], + ) + with pytest.raises(ValueError, match="日线"): + list(iter_bars(["600000"], "2024-01-01", "2024-01-31", "15m", + adjust="raw", cfg=_cfg({"raw_dir": "/x"}))) def test_unsupported_interval_rejected(): @@ -48,11 +83,11 @@ def test_unsupported_interval_rejected(): 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)], - ) + seen = [] + def mock_read(sym, start, end, cfg, dir_key): + seen.append(dir_key) + return [_mock_bar(sym, "2024-01-02", 10.0), _mock_bar(sym, "2024-01-02", 10.5)] + monkeypatch.setattr("sanguo_trader.data_source._read_fn", lambda iv: mock_read) bar = fetch_day("600000", "2024-01-02", "d") assert bar.close_price == 10.5 # 取最后一个 + assert seen == ["daily_dir"]