""" K线数据同步 — 个股日K + 指数日K 数据源: QMT xtdata 增量同步: 只拉 max(trade_date) 之后的增量数据 线程锁: KlineStockSync / KlineIndexSync 各自内部锁 """ import pandas as pd from datetime import date, 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() if latest is not None and latest >= today: PrintLog(LogLevel.INFO, f'[sync] KlineStock: 已最新 ({latest}),跳过') self.stats['skipped'] = 0 return self.stats # 增量起点 start_date = (latest + timedelta(days=1)) if latest else None start_str = start_date.strftime('%Y%m%d') if start_date else "" PrintLog(LogLevel.INFO, f'[sync] KlineStock: 增量同步,起点={start_str or "全部"}') 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) except Exception: self.stats['errors'] += 1 continue try: result = xtdata.get_market_data( field_list=field_list, stock_list=[code], period='1d', count=self.count, dividend_type='none', fill_data=False) inserted += self._upsert_incremental(code, result, start_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) -> 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 = [] for td in close_df.columns: td_date = td.date() if hasattr(td, 'date') else td if start_date is not None and td_date <= start_date: 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': _safe_get(result.get('volume'), full_code, td), }) 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() if latest is not None and latest >= today: PrintLog(LogLevel.INFO, f'[sync] KlineIndex: 已最新 ({latest}),跳过') return {} start_date = (latest + timedelta(days=1)) if latest else None start_str = start_date.strftime('%Y%m%d') if start_date else "" PrintLog(LogLevel.INFO, f'[sync] KlineIndex: 增量同步,起点={start_str or "全部"}') for code in index_codes: try: xtdata.download_history_data(code, period='1d', start_time=start_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', count=self.count, dividend_type='none', fill_data=False) return result or {} def _upsert(self, data): if not data: return latest = _latest_date(KlineIndex) records = [] close_df = data.get('close') 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: td_date = td.date() if hasattr(td, 'date') else td if latest is not None and td_date <= latest: 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': _safe_get(data.get('volume'), full_code, td), }) 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]