Files
ashare-data/src/fetchers/intraday.py
T
2026-05-22 09:34:49 +08:00

342 lines
11 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""分钟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 ~ %speriod=%s"
"通常表示本地客户端没有下载该周期的分钟缓存,或该股在此区间无分钟数据",
freq, code, tdx_code, sd, ed, tdx_period,
)
except Exception as exc:
fail += 1
_logger.warning(
"通达信 %s分钟K线 %s(%s) 失败: %s,区间 %s ~ %speriod=%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,
)