80 lines
2.7 KiB
Python
80 lines
2.7 KiB
Python
"""指数数据下载器 - 下载沪深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 not in ("0", "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",
|
||
adjustflag="3" # 不复权
|
||
)
|
||
|
||
if rs.error_code not in ("0", "success"):
|
||
print(f"Warning: 下载 {symbol} {year} 失败: {rs.error_msg}")
|
||
continue
|
||
|
||
# 转换为 DataFrame
|
||
data_list = []
|
||
while (rs.error_code in ("0", "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()
|