251 lines
10 KiB
Python
251 lines
10 KiB
Python
"""
|
|
K线数据同步 — 个股日K + 指数日K
|
|
数据源: QMT xtdata
|
|
增量同步: 只拉 max(trade_date) 之后的增量数据
|
|
线程锁: KlineStockSync / KlineIndexSync 各自内部锁
|
|
"""
|
|
import pandas as pd
|
|
from datetime import date, datetime, timedelta
|
|
from core.scoring.sync.base import BaseSync
|
|
from core.scoring.models import KlineStock, KlineIndex
|
|
from core.scoring.config import TRACKED_INDICES
|
|
from core.database import db
|
|
from core.logger import LogLevel, PrintLog
|
|
|
|
BATCH_SIZE = 50
|
|
DEFAULT_COUNT = 300
|
|
|
|
# 全局同步状态标记(字典引用传递,可被外部轮询)
|
|
_sync_state = {"kline": False, "index": False, "stocks": False,
|
|
"industry": False, "market": False, "sector": False}
|
|
|
|
|
|
def is_syncing(key="kline") -> bool:
|
|
return _sync_state.get(key, False)
|
|
|
|
|
|
def _latest_date(model_cls) -> date | None:
|
|
from peewee import fn
|
|
row = model_cls.select(fn.MAX(model_cls.trade_date)).scalar()
|
|
if isinstance(row, date):
|
|
return row
|
|
return None
|
|
|
|
|
|
def _safe_get(df_dict, code, dt, default=0.0) -> float:
|
|
"""安全获取 DataFrame 值"""
|
|
if df_dict is None:
|
|
return default
|
|
df = df_dict.get(code)
|
|
if df is None or code not in df.index:
|
|
return default
|
|
try:
|
|
val = df.loc[code, dt]
|
|
if pd.isna(val):
|
|
return default
|
|
return float(val)
|
|
except Exception:
|
|
return default
|
|
|
|
|
|
class KlineStockSync(BaseSync):
|
|
"""个股日K线同步 — 增量:只拉 max(trade_date) 之后的增量数据"""
|
|
|
|
def __init__(self, count: int = DEFAULT_COUNT):
|
|
super().__init__()
|
|
self.count = count
|
|
|
|
def _fetch(self, **kwargs):
|
|
from xtquant import xtdata
|
|
|
|
# 增量判断: 以数据库最新一条记录为准
|
|
latest = _latest_date(KlineStock)
|
|
today = date.today()
|
|
|
|
# 增量起点: last_db_date + 1; 截止: 昨天(盘中不能同步当天数据)
|
|
# 注意: 不能用 latest >= today 跳过,因为 latest 可能是错误的未收盘数据
|
|
start_date = (latest + timedelta(days=1)) if latest else None
|
|
end_date = today - timedelta(days=1) # 固定截止到昨天,收盘后同步昨天数据
|
|
start_str = start_date.strftime('%Y%m%d') if start_date else ""
|
|
end_str = end_date.strftime('%Y%m%d')
|
|
PrintLog(LogLevel.INFO,
|
|
f'[sync] KlineStock: 增量同步 {start_str} ~ {end_str}')
|
|
|
|
all_stocks = xtdata.get_stock_list_in_sector("沪深A股")
|
|
PrintLog(LogLevel.INFO, f'[sync] KlineStock: {len(all_stocks)} 只A股')
|
|
field_list = ['open', 'high', 'low', 'close', 'volume']
|
|
total = len(all_stocks)
|
|
inserted = 0
|
|
|
|
for i, code in enumerate(all_stocks):
|
|
if i > 0 and i % 50 == 0:
|
|
PrintLog(LogLevel.INFO, f'[sync] KlineStock: {i}/{total} ({i*100//total}%)')
|
|
|
|
try:
|
|
xtdata.download_history_data(code, period='1d', start_time=start_str, end_time=end_str)
|
|
except Exception:
|
|
self.stats['errors'] += 1
|
|
continue
|
|
|
|
try:
|
|
result = xtdata.get_market_data(
|
|
field_list=field_list, stock_list=[code], period='1d',
|
|
start_time=start_str, end_time=end_str,
|
|
dividend_type='none', fill_data=False)
|
|
inserted += self._upsert_incremental(code, result, start_date, end_date)
|
|
except Exception:
|
|
self.stats['errors'] += 1
|
|
|
|
self.stats['inserted'] = inserted
|
|
PrintLog(LogLevel.INFO,
|
|
f'[sync] KlineStock 完成: 新增={inserted} '
|
|
f'跳过={self.stats["skipped"]} 错误={self.stats["errors"]}')
|
|
return self.stats
|
|
|
|
def _upsert_incremental(self, full_code: str, result: dict, start_date, end_date) -> int:
|
|
if not result:
|
|
return 0
|
|
close_df = result.get('close')
|
|
if close_df is None or close_df.empty:
|
|
return 0
|
|
stock_code = full_code.split('.')[0]
|
|
records = []
|
|
vol_df = result.get('volume')
|
|
for td in close_df.columns:
|
|
# xtdata 返回的列名可能是字符串 'YYYYMMDD' 或 datetime,需统一转成 date
|
|
if isinstance(td, str):
|
|
td_date = datetime.strptime(td, '%Y%m%d').date()
|
|
else:
|
|
td_date = td.date() if hasattr(td, 'date') else td
|
|
# 过滤: 不在增量范围内的跳过 (start_date < td <= end_date)
|
|
if start_date is not None and td_date < start_date:
|
|
self.stats['skipped'] = self.stats.get('skipped', 0) + 1
|
|
continue
|
|
if end_date is not None and td_date > end_date:
|
|
self.stats['skipped'] = self.stats.get('skipped', 0) + 1
|
|
continue
|
|
# 跳过成交量为0的无效数据(盘中未结算数据)
|
|
vol = vol_df.loc[full_code, td] if vol_df is not None else None
|
|
if vol is None or (isinstance(vol, float) and pd.isna(vol)) or vol == 0:
|
|
self.stats['skipped'] = self.stats.get('skipped', 0) + 1
|
|
continue
|
|
close_val = close_df.loc[full_code, td]
|
|
if close_val is None or (isinstance(close_val, float) and pd.isna(close_val)):
|
|
self.stats['skipped'] = self.stats.get('skipped', 0) + 1
|
|
continue
|
|
records.append({
|
|
'stock_code': stock_code,
|
|
'trade_date': td_date,
|
|
'open': _safe_get(result.get('open'), full_code, td),
|
|
'high': _safe_get(result.get('high'), full_code, td),
|
|
'low': _safe_get(result.get('low'), full_code, td),
|
|
'close': float(close_val),
|
|
'volume': float(vol),
|
|
})
|
|
if records:
|
|
with db.atomic():
|
|
for batch in _chunked(records, 500):
|
|
KlineStock.insert_many(batch).on_conflict_replace().execute()
|
|
return len(records)
|
|
return 0
|
|
|
|
def _upsert(self, data):
|
|
pass
|
|
|
|
|
|
class KlineIndexSync(BaseSync):
|
|
"""指数日K线同步 — 增量同步"""
|
|
|
|
def __init__(self, indices: list = None, count: int = DEFAULT_COUNT):
|
|
super().__init__()
|
|
self.indices = indices or TRACKED_INDICES
|
|
self.count = count
|
|
|
|
def _fetch(self, **kwargs):
|
|
from xtquant import xtdata
|
|
|
|
index_codes = []
|
|
for code in self.indices:
|
|
if code.startswith(('000', '001')):
|
|
index_codes.append(f'{code}.SH')
|
|
elif code.startswith('399'):
|
|
index_codes.append(f'{code}.SZ')
|
|
else:
|
|
index_codes.append(f'{code}.SH')
|
|
|
|
latest = _latest_date(KlineIndex)
|
|
today = date.today()
|
|
start_date = (latest + timedelta(days=1)) if latest else None
|
|
end_date = today - timedelta(days=1) # 截止到昨天
|
|
start_str = start_date.strftime('%Y%m%d') if start_date else ""
|
|
end_str = end_date.strftime('%Y%m%d')
|
|
PrintLog(LogLevel.INFO,
|
|
f'[sync] KlineIndex: 增量同步 {start_str} ~ {end_str}')
|
|
|
|
for code in index_codes:
|
|
try:
|
|
xtdata.download_history_data(code, period='1d', start_time=start_str, end_time=end_str)
|
|
except Exception:
|
|
self.stats['errors'] += 1
|
|
|
|
field_list = ['open', 'high', 'low', 'close', 'volume']
|
|
result = xtdata.get_market_data(
|
|
field_list=field_list, stock_list=index_codes, period='1d',
|
|
start_time=start_str, end_time=end_str,
|
|
dividend_type='none', fill_data=False)
|
|
|
|
def _upsert(self, data):
|
|
if not data:
|
|
return
|
|
latest = _latest_date(KlineIndex)
|
|
today = date.today()
|
|
start_date = (latest + timedelta(days=1)) if latest else None
|
|
end_date = today - timedelta(days=1)
|
|
records = []
|
|
close_df = data.get('close')
|
|
vol_df = data.get('volume')
|
|
if close_df is None or close_df.empty:
|
|
return
|
|
for full_code in close_df.index:
|
|
index_code = full_code.split('.')[0]
|
|
for td in close_df.columns:
|
|
if isinstance(td, str):
|
|
td_date = datetime.strptime(td, '%Y%m%d').date()
|
|
else:
|
|
td_date = td.date() if hasattr(td, 'date') else td
|
|
# 增量范围过滤 (start_date < td <= end_date)
|
|
if start_date is not None and td_date < start_date:
|
|
self.stats['skipped'] = self.stats.get('skipped', 0) + 1
|
|
continue
|
|
if end_date is not None and td_date > end_date:
|
|
self.stats['skipped'] = self.stats.get('skipped', 0) + 1
|
|
continue
|
|
# 过滤成交量为0的无效数据
|
|
vol = vol_df.loc[full_code, td] if vol_df is not None else None
|
|
if vol is None or (isinstance(vol, float) and pd.isna(vol)) or vol == 0:
|
|
self.stats['skipped'] = self.stats.get('skipped', 0) + 1
|
|
continue
|
|
close_val = close_df.loc[full_code, td]
|
|
if close_val is None or (isinstance(close_val, float) and pd.isna(close_val)):
|
|
self.stats['skipped'] = self.stats.get('skipped', 0) + 1
|
|
continue
|
|
records.append({
|
|
'index_code': index_code,
|
|
'trade_date': td_date,
|
|
'open': _safe_get(data.get('open'), full_code, td),
|
|
'high': _safe_get(data.get('high'), full_code, td),
|
|
'low': _safe_get(data.get('low'), full_code, td),
|
|
'close': float(close_val),
|
|
'volume': float(vol),
|
|
})
|
|
if records:
|
|
with db.atomic():
|
|
for batch in _chunked(records, 500):
|
|
KlineIndex.insert_many(batch).on_conflict_replace().execute()
|
|
self.stats['inserted'] = self.stats.get('inserted', 0) + len(records)
|
|
|
|
|
|
def _chunked(lst: list, n: int):
|
|
for i in range(0, len(lst), n):
|
|
yield lst[i:i + n]
|