SHA256
342 lines
11 KiB
Python
342 lines
11 KiB
Python
"""分钟K线抓取模块 — 通达信本地客户端。
|
||
|
||
用法:
|
||
python -m src.main --intraday --freq 1 # 1分钟K线
|
||
python -m src.main --intraday --freq 5 # 5分钟K线
|
||
python -m src.main --intraday --start-date 20260508 --end-date 20260509
|
||
python -m src.main --intraday --symbol 000001
|
||
"""
|
||
|
||
import time
|
||
import random
|
||
from datetime import datetime, timedelta
|
||
|
||
import pandas as pd
|
||
from src.config import get_fetch_config
|
||
from src.db import StockMin1, StockMin5, TradingDay, batch_upsert, get_session, get_stock_codes
|
||
from src.fetchers.tdx_client import get_market_data as tdx_get_market_data
|
||
from src.fetchers.tdx_client import format_tdx_code_for_log
|
||
from src.fetchers.tdx_client import normalize_tdx_code
|
||
from src.fetchers.tdx_client import refresh_minute_cache as tdx_refresh_minute_cache
|
||
from src.log import get_logger
|
||
from sqlalchemy import select, func, distinct, text
|
||
|
||
_logger = get_logger("intraday")
|
||
|
||
|
||
VALID_FREQ = ("1", "5")
|
||
|
||
TDX_AUTOCACHE_PERIODS = ("1m", "5m")
|
||
|
||
FREQ_MODEL = {
|
||
"1": StockMin1,
|
||
"5": StockMin5,
|
||
}
|
||
|
||
FREQ_TDX_PERIOD = {
|
||
"1": "1m",
|
||
"5": "5m",
|
||
}
|
||
|
||
|
||
def _clean(val):
|
||
if val is None:
|
||
return None
|
||
if isinstance(val, str) and val.strip() == "":
|
||
return None
|
||
return val
|
||
|
||
|
||
def _parse_datetime(date_str: str, time_str: str) -> str | None:
|
||
"""将分钟线返回的 date + time 解析为 datetime 字符串
|
||
|
||
time 格式: "20260508093500000" (17位) 或 "09:35:00" (8位)
|
||
"""
|
||
if len(time_str) == 17:
|
||
return f"{time_str[:4]}-{time_str[4:6]}-{time_str[6:8]} " \
|
||
f"{time_str[8:10]}:{time_str[10:12]}:{time_str[12:14]}"
|
||
return f"{date_str} {time_str}"
|
||
|
||
|
||
def _parse_tdx_datetime(value) -> datetime:
|
||
"""将通达信返回的时间索引转成 `datetime`。"""
|
||
if isinstance(value, datetime):
|
||
return value
|
||
if hasattr(value, "to_pydatetime"):
|
||
return value.to_pydatetime()
|
||
return pd.to_datetime(value).to_pydatetime()
|
||
|
||
|
||
def _tdx_market_data_to_rows(code: str, market_data: dict) -> list[dict]:
|
||
"""把通达信 `get_market_data` 的返回值整理成数据库行。"""
|
||
if not market_data:
|
||
return []
|
||
|
||
lower_map = {str(key).lower(): value for key, value in market_data.items()}
|
||
field_map = {
|
||
"open": "open",
|
||
"high": "high",
|
||
"low": "low",
|
||
"close": "close",
|
||
"volume": "volume",
|
||
"amount": "amount",
|
||
}
|
||
|
||
base_frame = None
|
||
for field in ("close", "open", "high", "low", "volume", "amount"):
|
||
frame = lower_map.get(field)
|
||
if hasattr(frame, "index") and len(frame.index) > 0:
|
||
base_frame = frame
|
||
break
|
||
|
||
if base_frame is None:
|
||
return []
|
||
|
||
rows: list[dict] = []
|
||
for ts in base_frame.index:
|
||
row = {
|
||
"code": code,
|
||
"datetime": _parse_tdx_datetime(ts),
|
||
}
|
||
valid_value = False
|
||
for field, column in field_map.items():
|
||
frame = lower_map.get(field)
|
||
if frame is None or ts not in frame.index:
|
||
continue
|
||
cell = frame.loc[ts]
|
||
if isinstance(cell, pd.Series):
|
||
if code in cell.index:
|
||
value = cell[code]
|
||
else:
|
||
value = cell.iloc[0]
|
||
else:
|
||
value = cell
|
||
value = _clean(value)
|
||
row[column] = value
|
||
if value is not None:
|
||
valid_value = True
|
||
if valid_value:
|
||
rows.append(row)
|
||
|
||
return rows
|
||
|
||
|
||
def _tdx_fetch_rows_with_autocache(
|
||
code: str,
|
||
tdx_code: str,
|
||
tdx_period: str,
|
||
start_time: str,
|
||
end_time: str,
|
||
) -> tuple[list[dict], bool]:
|
||
"""先读取通达信本地数据,缺缓存时自动刷新一次后重试。"""
|
||
def _fetch_rows(period: str) -> list[dict]:
|
||
market_data = tdx_get_market_data(
|
||
[tdx_code],
|
||
period=period,
|
||
start_time=start_time,
|
||
end_time=end_time,
|
||
count=-1,
|
||
dividend_type="none",
|
||
fill_data=True,
|
||
)
|
||
return _tdx_market_data_to_rows(code, market_data)
|
||
|
||
rows = _fetch_rows(tdx_period)
|
||
if rows:
|
||
return rows, False
|
||
|
||
_logger.info(
|
||
"通达信 %s(%s) 首次没有返回数据,自动刷新本地分钟缓存后重试",
|
||
code, tdx_code,
|
||
)
|
||
tdx_refresh_minute_cache(
|
||
[tdx_code],
|
||
periods=TDX_AUTOCACHE_PERIODS,
|
||
batch_size=1,
|
||
pause_seconds=0,
|
||
)
|
||
rows = _fetch_rows(tdx_period)
|
||
if rows:
|
||
return rows, True
|
||
|
||
return [], True
|
||
|
||
|
||
def _get_intraday_gaps(model, code: str, sd: str, ed: str) -> list[tuple[str, str]]:
|
||
"""分析单只股票在 [sd, ed] 范围内的分钟K线缺口,返回缺失区间列表"""
|
||
session = get_session()
|
||
try:
|
||
# 获取范围内的交易日
|
||
trading_days = session.execute(
|
||
select(TradingDay.date)
|
||
.where(TradingDay.date >= sd)
|
||
.where(TradingDay.date <= ed)
|
||
.order_by(TradingDay.date)
|
||
).scalars().all()
|
||
if not trading_days:
|
||
return [(sd, ed)]
|
||
|
||
# 获取该股票已有数据的日期(按天去重)
|
||
existing = set(session.execute(
|
||
text(f"SELECT DISTINCT DATE(datetime) FROM {model.__tablename__} "
|
||
"WHERE code = :code AND datetime >= :sd AND datetime <= :ed"),
|
||
{"code": code, "sd": sd, "ed": ed + " 23:59:59"},
|
||
).scalars().all())
|
||
|
||
# 找出缺失的交易日
|
||
missing = [d for d in trading_days if d not in existing]
|
||
if not missing:
|
||
return []
|
||
|
||
# 合并为连续区间
|
||
gaps = []
|
||
gap_start = missing[0]
|
||
gap_end = missing[0]
|
||
for d in missing[1:]:
|
||
if (d - gap_end).days <= 3:
|
||
gap_end = d
|
||
else:
|
||
gaps.append((str(gap_start), str(gap_end)))
|
||
gap_start = d
|
||
gap_end = d
|
||
gaps.append((str(gap_start), str(gap_end)))
|
||
return gaps
|
||
finally:
|
||
session.close()
|
||
|
||
|
||
def fetch_intraday(start_date: str | None = None, end_date: str | None = None,
|
||
symbol: str | None = None, freq: str = "5",
|
||
prewarm_cache: bool = True):
|
||
"""抓取分钟K线行情
|
||
|
||
Args:
|
||
start_date: 开始日期 YYYYMMDD,默认30天前
|
||
end_date: 结束日期 YYYYMMDD,默认今天
|
||
symbol: 单只股票代码,默认全部
|
||
freq: K线频率 1 或 5
|
||
prewarm_cache: 是否先批量刷新通达信本地分钟缓存
|
||
"""
|
||
if freq not in VALID_FREQ:
|
||
_logger.error("不支持的频率 %s,可选: %s", freq, ", ".join(VALID_FREQ))
|
||
return
|
||
|
||
_fetch_one_freq_tdx(freq, start_date, end_date, symbol, prewarm_cache=prewarm_cache)
|
||
|
||
|
||
def _fetch_one_freq_tdx(freq: str, start_date: str | None, end_date: str | None,
|
||
symbol: str | None, prewarm_cache: bool = True):
|
||
"""使用通达信本地客户端抓取单个频率的分钟K线。"""
|
||
cfg = get_fetch_config()
|
||
base_delay = float(cfg.get("delay", 0.1))
|
||
delay = max(base_delay, 0.2)
|
||
model = FREQ_MODEL[freq]
|
||
tdx_period = FREQ_TDX_PERIOD[freq]
|
||
|
||
if end_date is None:
|
||
end_date = datetime.now().strftime("%Y%m%d")
|
||
if start_date is None:
|
||
start_date = (datetime.now() - timedelta(days=30)).strftime("%Y%m%d")
|
||
|
||
sd = f"{start_date[:4]}-{start_date[4:6]}-{start_date[6:8]}"
|
||
ed = f"{end_date[:4]}-{end_date[4:6]}-{end_date[6:8]}"
|
||
|
||
if symbol:
|
||
codes = [symbol]
|
||
else:
|
||
codes = get_stock_codes()
|
||
if not codes:
|
||
_logger.error("无股票列表,请先运行 --stock-info")
|
||
return
|
||
|
||
need_fetch = []
|
||
if symbol:
|
||
need_fetch = [(symbol, [(sd, ed)])]
|
||
else:
|
||
skip = 0
|
||
for code in codes:
|
||
gaps = _get_intraday_gaps(model, code, sd, ed)
|
||
if gaps:
|
||
need_fetch.append((code, gaps))
|
||
else:
|
||
skip += 1
|
||
codes = [c for c, _ in need_fetch]
|
||
|
||
if not need_fetch:
|
||
_logger.info("通达信 %s分钟K线 %s ~ %s 数据已完整,跳过", freq, sd, ed)
|
||
return
|
||
|
||
skip_msg = f"(跳过 {skip} 只已完整)" if not symbol and skip else ""
|
||
preview = ", ".join(format_tdx_code_for_log(code) for code in codes[:8])
|
||
if len(codes) > 8:
|
||
preview += f", ...(+{len(codes) - 8})"
|
||
_logger.info(
|
||
"正在使用通达信抓取%s分钟K线 %s ~ %s,需补缺 %d 只%s,目标: [%s]",
|
||
freq, sd, ed, len(need_fetch), skip_msg, preview,
|
||
)
|
||
|
||
# 全市场抓取先批量预热,单只股票则走按需自动缓存,避免重复刷新。
|
||
if not symbol and prewarm_cache:
|
||
try:
|
||
refresh_results = tdx_refresh_minute_cache(codes, periods=TDX_AUTOCACHE_PERIODS)
|
||
refresh_failed = [item for item in refresh_results if item.get("result", {}).get("ErrorId") not in (None, "0", 0)]
|
||
if refresh_failed:
|
||
_logger.warning(
|
||
"通达信分钟缓存刷新后仍有 %d 个批次失败,分钟线可能不完整",
|
||
len(refresh_failed),
|
||
)
|
||
except Exception as exc:
|
||
_logger.warning("通达信分钟缓存刷新失败,继续尝试取数:%s", exc)
|
||
|
||
total = len(need_fetch)
|
||
success = 0
|
||
fail = 0
|
||
t_start = time.time()
|
||
|
||
for i, (code, gaps) in enumerate(need_fetch):
|
||
total_rows = 0
|
||
tdx_code = normalize_tdx_code(code)
|
||
try:
|
||
for gap_sd, gap_ed in gaps:
|
||
start_time = f"{gap_sd.replace('-', '')}000000"
|
||
end_time = f"{gap_ed.replace('-', '')}235959"
|
||
rows, _refreshed = _tdx_fetch_rows_with_autocache(
|
||
code,
|
||
tdx_code,
|
||
tdx_period,
|
||
start_time,
|
||
end_time,
|
||
)
|
||
if rows:
|
||
batch_upsert(model, rows, ["code", "datetime"])
|
||
total_rows += len(rows)
|
||
time.sleep(delay + random.uniform(0, delay * 0.2))
|
||
if total_rows:
|
||
success += 1
|
||
else:
|
||
fail += 1
|
||
_logger.warning(
|
||
"通达信 %s分钟K线 %s(%s) 没有返回数据,区间 %s ~ %s,period=%s;"
|
||
"通常表示本地客户端没有下载该周期的分钟缓存,或该股在此区间无分钟数据",
|
||
freq, code, tdx_code, sd, ed, tdx_period,
|
||
)
|
||
except Exception as exc:
|
||
fail += 1
|
||
_logger.warning(
|
||
"通达信 %s分钟K线 %s(%s) 失败: %s,区间 %s ~ %s,period=%s",
|
||
freq, code, tdx_code, exc, sd, ed, tdx_period,
|
||
)
|
||
|
||
if (i + 1) % 100 == 0:
|
||
elapsed = time.time() - t_start
|
||
_logger.info(
|
||
"[%d/%d] 进度... 成功:%d 失败:%d 已用时:%.0fs",
|
||
i + 1, total, success, fail, elapsed,
|
||
)
|
||
|
||
total_time = time.time() - t_start
|
||
_logger.info(
|
||
"通达信 %s分钟K线抓取完成,成功:%d 失败:%d 总耗时:%.1fs",
|
||
freq, success, fail, total_time,
|
||
)
|