diff --git a/sanguo_data/datareader.py b/sanguo_data/datareader.py index 7d10b78..a19facd 100644 --- a/sanguo_data/datareader.py +++ b/sanguo_data/datareader.py @@ -9,7 +9,7 @@ if _VNPY_SRC not in sys.path: sys.path.insert(0, _VNPY_SRC) import pandas as pd -from datetime import datetime +from datetime import datetime, date from vnpy.trader.object import BarData from vnpy.trader.constant import Exchange, Interval from vnpy.trader.setting import SETTINGS @@ -103,3 +103,43 @@ def read_parquet_15min(symbol: str, start: str, end: str, cfg, dir_key: str = "m 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"]) diff --git a/sanguo_data/index_downloader.py b/sanguo_data/index_downloader.py new file mode 100644 index 0000000..691b13c --- /dev/null +++ b/sanguo_data/index_downloader.py @@ -0,0 +1,79 @@ +"""指数数据下载器 - 下载沪深300等指数日线数据(baostock)""" +import os +import time +from pathlib import Path +from datetime import datetime + + +def download_index(symbol: str, start_year: int, end_year: int, out_dir: str) -> None: + """ + 下载指数日线数据(baostock) + + Args: + symbol: 指数代码,如 "sh000300"(沪深300) + start_year: 起始年份 + end_year: 结束年份 + out_dir: 输出目录(会按年切分:{out_dir}/{year}/{symbol}_daily.parquet) + """ + # 入口处清除代理(直连不走代理,避免数据源封IP) + for key in ("http_proxy", "https_proxy", "HTTP_PROXY", "HTTPS_PROXY"): + os.environ.pop(key, None) + + import pandas as pd + import baostock as bs + + # 登录 baostock + lg = bs.login() + if lg.error_code != "success": + raise RuntimeError(f"baostock 登录失败: {lg.error_msg}") + + try: + for year in range(start_year, end_year + 1): + # 按年下载 + start_date = f"{year}-01-01" + end_date = f"{year}-12-31" + + # baostock 用带点格式 "sh.000300" + bs_symbol = f"{symbol[:2]}.{symbol[2:]}" + rs = bs.query_history_k_data_plus( + bs_symbol, + "date,open,high,low,close,volume", + start_date=start_date, + end_date=end_date, + frequency="d", + adjustfield="3" # 不复权 + ) + + if rs.error_code != "success": + print(f"Warning: 下载 {symbol} {year} 失败: {rs.error_msg}") + continue + + # 转换为 DataFrame + data_list = [] + while (rs.error_code == "success") and rs.next(): + data_list.append(rs.get_row_data()) + + if not data_list: + print(f"Warning: {symbol} {year} 无数据") + continue + + df = pd.DataFrame(data_list, columns=rs.fields) + df["date"] = pd.to_datetime(df["date"]) + + # 数值列转换 + for col in ["open", "high", "low", "close", "volume"]: + df[col] = pd.to_numeric(df[col], errors="coerce") + + # 按年分片写入 parquet(与现有股票日线同路径格式) + year_dir = Path(out_dir) / str(year) + year_dir.mkdir(parents=True, exist_ok=True) + output_file = year_dir / f"{symbol}_daily.parquet" + df.to_parquet(output_file, index=False) + print(f"已下载 {symbol} {year}: {len(df)} 行 → {output_file}") + + # 单线程限速(避免数据源封IP) + time.sleep(0.5) + + finally: + # 登出 baostock + bs.logout()