diff --git a/sanguo_data/datareader.py b/sanguo_data/datareader.py index 8abae60..af19278 100644 --- a/sanguo_data/datareader.py +++ b/sanguo_data/datareader.py @@ -110,61 +110,60 @@ def read_parquet_15min(symbol: str, start: str, end: str, cfg, dir_key: str = "m def read_index_daily(code: str, start, end, cfg) -> pd.DataFrame: """ - 读指数日线数据(sh000300/sz399001 等),从 vnpy DB 读取(统一数据源)。 - parquet 仅作原始备份,不再读取。返回类型保持 pd.DataFrame(cta_engine 消费不变)。 + 读指数日线数据(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: 起始日期(str "YYYY-MM-DD" / date / datetime) - end: 结束日期(str "YYYY-MM-DD" / date / datetime) - cfg: 数据配置对象 + 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 列;无数据返回空 DataFrame。 + pd.DataFrame: date/open/high/low/close/volume 列(date 为 Timestamp);无数据返回空 DataFrame。 """ - from vnpy.trader.database import get_database # lazy:避免模块 import 依赖数据库驱动 + import sqlite3 - # 前缀解析交易所(指数不能用 guess_exchange:000300 以 0 开头会被误判成 SZSE, - # 但 000300 实际属于 SSE)。sh → SSE,sz → SZSE。 + # 前缀解析交易所(指数不能用 guess_exchange:000300 以 0 开头会误判 SZSE,但属 SSE)。 symbol = code[2:] - exchange = Exchange.SSE if code.startswith("sh") else Exchange.SZSE + exchange = "SSE" if code.startswith("sh") else "SZSE" - # 日期归一化:str → parse, date → combine, datetime → as-is - def _to_dt(s, is_start: bool) -> datetime: + # 日期归一化 -> 'YYYY-MM-DD' + def _to_str(s) -> str: if isinstance(s, datetime): - return s + return s.strftime("%Y-%m-%d") if isinstance(s, date): - return datetime.combine(s, datetime.min.time() if is_start else datetime.max.time()) - return datetime.strptime(s, "%Y-%m-%d") + return s.strftime("%Y-%m-%d") + return str(s)[:10] - start_dt = _to_dt(start, True) - end_dt = _to_dt(end, False) + start_str = _to_str(start) + end_str = _to_str(end) - # 配置 vnpy DB(与 read_db_daily 同模式) - SETTINGS["database.name"] = "sqlite" - vnpy_db = getattr(cfg, "data_paths", {}).get("vnpy_db") - if not vnpy_db: + db_path = getattr(cfg, "data_paths", {}).get("vnpy_db") + if not db_path: raise RuntimeError("config 缺 data_paths.vnpy_db — 检查 backtest.yaml 初始化") - SETTINGS["database.database"] = vnpy_db - db = get_database() - bars = db.load_bar_data( - symbol=symbol, - exchange=exchange, - interval=Interval.DAILY, - start=start_dt, - end=end_dt, - ) + 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 not bars: + if df.empty: return pd.DataFrame(columns=["date", "open", "high", "low", "close", "volume"]) - df = pd.DataFrame([{ - "date": b.datetime, - "open": b.open_price, - "high": b.high_price, - "low": b.low_price, - "close": b.close_price, - "volume": b.volume, - } for b in bars]) - return df.sort_values("date").reset_index(drop=True) + 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) diff --git a/tests/data/test_index_downloader.py b/tests/data/test_index_downloader.py index e724ffa..949196e 100644 --- a/tests/data/test_index_downloader.py +++ b/tests/data/test_index_downloader.py @@ -131,55 +131,73 @@ def test_download_index_clears_proxy(tmp_path): assert "https_proxy" not in os.environ -def test_read_index_daily_reads_from_vnpy_db(tmp_path): - """read_index_daily 经 vnpy get_database.load_bar_data 读指数日线,返回 DataFrame。 +def test_read_index_daily_reads_from_dbbardata(tmp_path): + """read_index_daily 从 dbbardata 表读指数日线(2026-08-01 切自 vnpy DbBarData)。 - 注:read_index_daily 数据源是 vnpy DbBarData 表(非 parquet);D2 根治拟切 dbbardata, - 届时 mock 数据源随之更新。当前测 vnpy 路径(与 read_db_daily 同 mock 模式)。 + 建 tmp sqlite 库 + dbbardata 表插指数行,测真实 SQL 路径(非 mock)。 + 中证指数用 sh 前缀 → exchange=SSE(点位由 sina_index_eod 灌入)。 """ - from types import SimpleNamespace - from datetime import datetime + import sqlite3 from sanguo_data.datareader import read_index_daily + db = tmp_path / "q.db" + conn = sqlite3.connect(db) + conn.execute( + "CREATE TABLE dbbardata (symbol TEXT, exchange TEXT, datetime TEXT, interval TEXT, " + "volume REAL, turnover REAL, open_interest REAL, open_price REAL, " + "high_price REAL, low_price REAL, close_price REAL)" + ) + conn.executemany( + "INSERT INTO dbbardata (symbol,exchange,datetime,interval,volume,turnover," + "open_interest,open_price,high_price,low_price,close_price) VALUES (?,?,?,?,?,?,?,?,?,?,?)", + [ + ("000300", "SSE", "2024-01-02", "d", 100000, 0, 0, 3495.0, 3505.0, 3490.0, 3500.0), + ("000300", "SSE", "2024-01-03", "d", 120000, 0, 0, 3505.0, 3515.0, 3500.0, 3510.0), + ("000300", "SSE", "2024-01-04", "d", 110000, 0, 0, 3515.0, 3525.0, 3510.0, 3520.0), + ], + ) + conn.commit() + conn.close() + cfg = DataConfig( - data_paths={"vnpy_db": str(tmp_path / "q.db")}, + data_paths={"vnpy_db": str(db)}, data_sources={}, validation={}, performance={}, ) - # vnpy load_bar_data 返回类 BarData 对象(有 datetime/open_price/... 属性) - bars = [ - SimpleNamespace(datetime=datetime(2024, 1, 2), open_price=3495.0, - high_price=3505.0, low_price=3490.0, close_price=3500.0, volume=100000), - SimpleNamespace(datetime=datetime(2024, 1, 3), open_price=3505.0, - high_price=3515.0, low_price=3500.0, close_price=3510.0, volume=120000), - SimpleNamespace(datetime=datetime(2024, 1, 4), open_price=3515.0, - high_price=3525.0, low_price=3510.0, close_price=3520.0, volume=110000), - ] - mock_db = MagicMock() - mock_db.load_bar_data.return_value = bars - # read_index_daily 内 lazy import,patch 源模块属性(同 test_datareader) - with patch("vnpy.trader.database.get_database", return_value=mock_db): - result = read_index_daily("sh000300", date(2024, 1, 1), date(2024, 12, 31), cfg) + result = read_index_daily("sh000300", date(2024, 1, 1), date(2024, 12, 31), cfg) assert len(result) == 3 - assert "close" in result.columns + assert list(result.columns) == ["date", "open", "high", "low", "close", "volume"] assert result["close"].iloc[0] == 3500.0 -def test_read_index_daily_passes_date_range_to_vnpy(tmp_path): - """read_index_daily 把 start/end 传给 vnpy load_bar_data(日期过滤在 vnpy 层)。""" - from datetime import datetime +def test_read_index_daily_filters_date_range(tmp_path): + """read_index_daily 用 substr(datetime,1,10) 比日期,正确过滤 start/end 范围。""" + import sqlite3 from sanguo_data.datareader import read_index_daily + db = tmp_path / "q.db" + conn = sqlite3.connect(db) + conn.execute( + "CREATE TABLE dbbardata (symbol TEXT, exchange TEXT, datetime TEXT, interval TEXT, " + "volume REAL, turnover REAL, open_interest REAL, open_price REAL, " + "high_price REAL, low_price REAL, close_price REAL)" + ) + conn.executemany( + "INSERT INTO dbbardata (symbol,exchange,datetime,interval,volume,turnover," + "open_interest,open_price,high_price,low_price,close_price) VALUES (?,?,?,?,?,?,?,?,?,?,?)", + [ + ("000300", "SSE", "2024-01-02", "d", 100, 0, 0, 1, 1, 1, 3500.0), + ("000300", "SSE", "2024-06-15", "d", 100, 0, 0, 1, 1, 1, 3600.0), + ("000300", "SSE", "2025-01-10", "d", 100, 0, 0, 1, 1, 1, 3700.0), # 范围外 + ], + ) + conn.commit() + conn.close() + cfg = DataConfig( - data_paths={"vnpy_db": str(tmp_path / "q.db")}, + data_paths={"vnpy_db": str(db)}, data_sources={}, validation={}, performance={}, ) - mock_db = MagicMock() - mock_db.load_bar_data.return_value = [] - with patch("vnpy.trader.database.get_database", return_value=mock_db): - read_index_daily("sh000300", date(2024, 1, 1), date(2024, 6, 30), cfg) - - # start=date→combine 00:00:00; end=date→combine datetime.max.time()(23:59:59.999999) - call = mock_db.load_bar_data.call_args - assert call.kwargs["start"] == datetime(2024, 1, 1, 0, 0, 0) - assert call.kwargs["end"] == datetime(2024, 6, 30, 23, 59, 59, 999999) + result = read_index_daily("sh000300", "2024-01-01", "2024-12-31", cfg) + assert len(result) == 2 # 只 2024 两行, 2025 范围外排除 + assert result["close"].iloc[1] == 3600.0