From 365f1ccbfba98e16875788a08206cd7506da1d9d Mon Sep 17 00:00:00 2001 From: shy3130 Date: Sun, 12 Jul 2026 14:23:22 +0800 Subject: [PATCH] =?UTF-8?q?feat(minute-k):=20=E5=88=86=E9=92=9FK=E5=90=8C?= =?UTF-8?q?=E6=AD=A5=E6=B5=81=E5=BC=8F=E8=90=BD=E7=9B=98=20+=20=E6=AE=B5?= =?UTF-8?q?=E5=A4=A7=E5=B0=8F=E5=8F=AF=E9=85=8D=20+=20=E5=9B=9E=E6=B5=8B?= =?UTF-8?q?=E6=88=90=E4=BA=A4=E4=BF=AE=E5=A4=8D=20+=20=E5=8D=A1=E6=AD=BB?= =?UTF-8?q?=E8=B6=85=E6=97=B6=E8=B0=83=E6=95=B4?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 1. 分钟K流式落盘 (修复1年全市场卡死) - sync_minute_batch 加 on_segment 回调,每段拉完立即写盘 - 抽出 _write_minute_partition 公共函数,消除重复 - 内存峰值从全量(~25GB)降到单段,避免 OOM 2. 段大小可配 (默认20交易日,范围[5,30]) - 新增 minute_sync_segment_days 偏好项 - 「单次获取」与「获取1年」共用此分段设置 - 前端设置弹窗新增分段大小步进器 3. 回测分钟K精确成交修复 (从未成功跑通过) - _resolve_minute_fill 改接受 DataFrame,避免 to_numpy() 退化为 object 数组后字段名索引抛 IndexError - _load_minute_for_fills 用 partition_by 向量化分组替换逐行循环 - 新增 test_minute_fill.py 回归测试 (6 用例) 4. 卡死 watchdog 超时阈值按任务类型区分 - 普通任务 600s → 1200s;分钟K长任务 → 1800s - create(timeout_s=) 支持按 job 存阈值,reap_stale 读 job 自身值 - 修复分钟K正常拉取被误杀导致僵尸线程继续写盘的问题 5. UI 优化 - 设置弹窗按 自动同步/手动获取/清空 三区块分组 - 「单次获取」改为按分段大小拉一段 (天数=分段设置) - 状态文案: 盘后自动同步 → 自动同步已开启/已关闭 --- backend/app/api/kline.py | 5 +- backend/app/api/settings.py | 17 ++- backend/app/backtest/engine.py | 51 ++++---- backend/app/services/kline_sync.py | 109 ++++++++++++------ backend/app/services/pipeline_jobs.py | 33 ++++-- backend/app/services/preferences.py | 10 ++ backend/tests/backtest/test_minute_fill.py | 104 +++++++++++++++++ .../src/components/data/MinuteSyncConfig.tsx | 80 ++++++++++--- frontend/src/lib/api.ts | 9 +- 9 files changed, 328 insertions(+), 90 deletions(-) create mode 100644 backend/tests/backtest/test_minute_fill.py diff --git a/backend/app/api/kline.py b/backend/app/api/kline.py index 1bc5d49..ba9d96c 100644 --- a/backend/app/api/kline.py +++ b/backend/app/api/kline.py @@ -585,7 +585,7 @@ async def sync_minute(request: Request): """ import asyncio - from app.services.pipeline_jobs import job_store, release_run_slot, try_acquire_run_slot + from app.services.pipeline_jobs import job_store, release_run_slot, try_acquire_run_slot, LONG_JOB_TIMEOUT_S from app.api.data import invalidate_storage_cache from app.services.preferences import get_minute_sync_days from app.tickflow.capabilities import Cap @@ -607,7 +607,8 @@ async def sync_minute(request: Request): override_days = body.get("days") extend_flag = body.get("extend") - job_id, is_new = job_store.create() + # 分钟K全市场同步是长任务(数据量是日K的 ~240 倍),用更宽松的卡死阈值 + job_id, is_new = job_store.create(timeout_s=LONG_JOB_TIMEOUT_S) if not is_new: return {"status": "reused", "job_id": job_id} diff --git a/backend/app/api/settings.py b/backend/app/api/settings.py index 06d715f..5ace21d 100644 --- a/backend/app/api/settings.py +++ b/backend/app/api/settings.py @@ -337,6 +337,8 @@ def _realtime_allowed() -> bool: class MinuteSyncPrefs(BaseModel): minute_sync_enabled: bool minute_sync_days: int = 5 + # 单段大小(交易日),None 表示不修改现有值。范围 [5, 30],默认 20。 + minute_sync_segment_days: int | None = None class DataProvidersIn(BaseModel): @@ -395,6 +397,7 @@ def get_preferences() -> dict: "indices_nav_pinned": preferences.get_indices_nav_pinned(), "minute_sync_enabled": preferences.get_minute_sync_enabled(), "minute_sync_days": preferences.get_minute_sync_days(), + "minute_sync_segment_days": preferences.get_minute_sync_segment_days(), "daily_data_provider": preferences.get_daily_data_provider(), "adj_factor_provider": preferences.get_adj_factor_provider(), "minute_data_provider": preferences.get_minute_data_provider(), @@ -643,16 +646,24 @@ def update_screener_result_columns(req: dict) -> dict: @router.put("/preferences/minute-sync") def update_minute_sync(req: MinuteSyncPrefs) -> dict: - """保存分钟 K 同步偏好。""" + """保存分钟 K 同步偏好。 + + minute_sync_segment_days 为可选:未传(None)时不覆盖现有值,便于开关/天数 + 与段大小各自独立更新。 + """ from app.services import preferences days = max(1, min(30, req.minute_sync_days)) - preferences.save({ + updates: dict = { "minute_sync_enabled": req.minute_sync_enabled, "minute_sync_days": days, - }) + } + if req.minute_sync_segment_days is not None: + updates["minute_sync_segment_days"] = max(5, min(30, req.minute_sync_segment_days)) + preferences.save(updates) return { "minute_sync_enabled": req.minute_sync_enabled, "minute_sync_days": days, + "minute_sync_segment_days": preferences.get_minute_sync_segment_days(), } diff --git a/backend/app/backtest/engine.py b/backend/app/backtest/engine.py index b9e5437..7c61b77 100644 --- a/backend/app/backtest/engine.py +++ b/backend/app/backtest/engine.py @@ -527,7 +527,7 @@ class BacktestEngine: ) for _key, _mdf in _loaded.items(): if not _mdf.is_empty(): - minute_cache[_key] = _mdf.to_numpy() + minute_cache[_key] = _mdf def _refill_price(idx: int, side: str, daily_price: float) -> float: if not config.minute_fill or not minute_cache: @@ -839,29 +839,33 @@ class BacktestEngine: @staticmethod def _resolve_minute_fill( - minute_rows: np.ndarray, + minute_df: pl.DataFrame, ref_price: float | None, side: str, ) -> float | None: """用当日分钟K确定精确成交价。 Args: - minute_rows: structured numpy array, 字段含 open/high/low/close/volume/amount + minute_df: 当日分钟K polars DataFrame, 列含 open/high/low/close/volume/amount ref_price: 信号参考线价格 (如 MA5 值); None 表示无参考线 side: "buy" 或 "sell", 决定穿越方向 Returns: 精确成交价, 或 None (降级到日K口径) + + 注: 原实现用 df.to_numpy() 转 structured array 再按字段名索引, 但当列类型不 + 一致时 (如 datetime 列 + float 列) to_numpy() 退化为 object 二维数组, 字段名 + 索引 arr["open"] 会抛 IndexError。改为直接按列取 Series, 稳定且更快。 """ - if minute_rows is None or len(minute_rows) == 0: + if minute_df is None or minute_df.is_empty(): return None - opens = minute_rows["open"].astype(float) - highs = minute_rows["high"].astype(float) - lows = minute_rows["low"].astype(float) - closes = minute_rows["close"].astype(float) - volumes = minute_rows["volume"].astype(float) if "volume" in minute_rows.dtype.names else None - amounts = minute_rows["amount"].astype(float) if "amount" in minute_rows.dtype.names else None + opens = minute_df["open"].to_numpy().astype(float) + highs = minute_df["high"].to_numpy().astype(float) + lows = minute_df["low"].to_numpy().astype(float) + closes = minute_df["close"].to_numpy().astype(float) + volumes = minute_df["volume"].to_numpy().astype(float) if "volume" in minute_df.columns else None + amounts = minute_df["amount"].to_numpy().astype(float) if "amount" in minute_df.columns else None # 有参考线 → 穿越价成交 (逻辑同止损: 找价格穿越参考线的时刻) if ref_price is not None and ref_price > 0 and np.isfinite(ref_price): @@ -914,22 +918,19 @@ class BacktestEngine: if df.is_empty(): return {} + # 按 (symbol, 日期) 向量化分组, 替代原 iter_rows 逐行 Python 循环。 + # 原实现对每行做 dict 转换再重建 DataFrame, 回测区间内触发股数多时极慢。 + df = df.with_columns( + pl.col("datetime").dt.strftime("%Y-%m-%d").alias("_d_str") + ) cache: dict = {} - for row in df.iter_rows(named=True): - dt = row.get("datetime") - if dt is None: + for sub in df.partition_by(["symbol", "_d_str"], as_dict=False): + if sub.is_empty(): continue - d_str = str(dt)[:10] - sym = row["symbol"] - key = (sym, d_str) - if key not in cache: - cache[key] = [] - cache[key].append(row) - # 转 DataFrame per key - result: dict = {} - for key, rows in cache.items(): - result[key] = pl.DataFrame(rows) - return result + sym = sub["symbol"][0] + d_str = sub["_d_str"][0] + cache[(sym, d_str)] = sub.drop("_d_str") + return cache def simulate_portfolio( self, @@ -1052,7 +1053,7 @@ class BacktestEngine: ) for key, mdf in loaded.items(): if not mdf.is_empty(): - minute_cache[key] = mdf.to_numpy() + minute_cache[key] = mdf def _refill_price(idx: int, side: str, daily_price: float) -> float: """分钟K精确成交价; 无数据则降级为 daily_price。""" diff --git a/backend/app/services/kline_sync.py b/backend/app/services/kline_sync.py index 7bee346..f13dfd1 100644 --- a/backend/app/services/kline_sync.py +++ b/backend/app/services/kline_sync.py @@ -499,6 +499,34 @@ def _datetime_to_ms(dt: datetime) -> int: return int(dt.timestamp() * 1000) +def _write_minute_partition(df: pl.DataFrame, minute_dir) -> int: + """按 _trade_date 分区落盘分钟 K (读旧→concat→unique→原子写)。返回写入行数。 + + 抽自原 sync_and_persist_minute 末尾的循环, 供流式落盘 (每段一次) 与一次性迁移共用。 + """ + if df.is_empty(): + return 0 + df = df.with_columns(pl.col("datetime").dt.date().alias("_trade_date")) + written = 0 + for day_df in df.partition_by("_trade_date"): + trade_date = day_df["_trade_date"][0] + out = minute_dir / f"date={trade_date}" / "part.parquet" + out.parent.mkdir(parents=True, exist_ok=True) + if out.exists(): + existing = pl.read_parquet(out) + if "datetime" in existing.columns: + existing = existing.filter(pl.col("datetime").is_not_null()) + day_df = pl.concat([existing, day_df.drop("_trade_date")]).unique( + subset=["symbol", "datetime"], keep="last", + ) + else: + day_df = day_df.drop("_trade_date") + day_df = day_df.sort("symbol", "datetime") + _atomic_write_parquet(day_df, out) + written += day_df.height + return written + + def sync_minute_batch( symbols: list[str], start_time: datetime | None = None, @@ -507,6 +535,8 @@ def sync_minute_batch( batch_size: int | None = None, rpm: int | None = None, on_chunk_done: Callable[[int, int, str], None] | None = None, + segment_trading_days: int = 20, + on_segment: Callable[[pl.DataFrame], None] | None = None, ) -> pl.DataFrame: """批量拉取多股分钟 K。 @@ -514,8 +544,15 @@ def sync_minute_batch( count 仅作为 fallback 保留。 on_chunk_done(current, total) 每个 chunk 完成后回调。 - TickFlow count 上限 10000 根/股, 1 天 240 根 → 单次最多约 41 天。 - 当区间超过 35 天时自动按月 (30 天) 分段拉取, 拼接结果。 + segment_trading_days: 单段大小 (交易日), 控制每次 SDK 请求覆盖的天数。 + TickFlow count 上限 10000 根/股, 1 天 240 根 → 物理上限 ~41 交易日; + 默认 20 (4800 根, 安全余量足), 范围建议 [5, 30]。 + 段越小: 单次内存峰值越低 (适合小内存机器), 但总批数↑ → 限速 sleep↑ → 更慢。 + 段越大: 速度越快, 内存峰值越高。 + on_segment: 每个时间段拉完后回调 (传入该段拼接后的 DataFrame)。 + 传入时走「流式落盘」: 段内结果累积到 seg_out, 段末 concat 后回调并清空, + 不进入全局 out → 内存峰值从「全量」降到「单段」。适用于 sync_and_persist_minute。 + 不传时 (如 get_minute_batch 的实时补拉) 保持原契约: 累积进 out 末尾一次性返回。 """ # 自定义数据源分流: minute provider provider_name = preferences.get_minute_data_provider() @@ -531,8 +568,9 @@ def sync_minute_batch( tf = get_client() # TickFlow count 上限 10000 根/股, 1 天 240 根 → 单次最多约 41 个交易日。 - # 按 41 交易日 (≈57 自然日) 分段: ≤41 交易日的区间只产生 1 段 (单次拉满)。 - SEG_CHUNK = timedelta(days=57) # 41 交易日 × 7/5 ≈ 57 自然日 + # 按 segment_trading_days 交易日分段 (交易日→自然日 ×7/5 换算, 含节假日余量)。 + seg_calendar_days = max(1, int(segment_trading_days * 7 / 5)) + SEG_CHUNK = timedelta(days=seg_calendar_days) time_segments: list[tuple[datetime, datetime]] = [] if start_time and end_time: seg_start = start_time @@ -545,7 +583,10 @@ def sync_minute_batch( total_steps = len(time_segments) * len(chunked(symbols, batch_size)) step = 0 + # 全局累积 (仅 on_segment=None 时使用, 末尾一次性 concat 返回) out: list[pl.DataFrame] = [] + # 段内累积: 每段拉完即 flush, 避免全量攒内存 (OOM 根因) + seg_out: list[pl.DataFrame] = [] for seg_idx, (seg_start, seg_end) in enumerate(time_segments): # 当前的日期段描述 (供进度展示) @@ -580,13 +621,21 @@ def sync_minute_batch( for sym, sub in raw.items(): if sub is None or len(sub) == 0: continue - out.append(_normalize_minute(sub, default_symbol=sym)) + seg_out.append(_normalize_minute(sub, default_symbol=sym)) elif raw is not None and len(raw) > 0: - out.append(_normalize_minute(raw)) + seg_out.append(_normalize_minute(raw)) if on_chunk_done: on_chunk_done(step, total_steps, seg_label) + # 段末 flush: 流式落盘回调 或 并入全局 out + if seg_out: + if on_segment: + on_segment(pl.concat(seg_out, how="diagonal_relaxed")) + else: + out.extend(seg_out) + seg_out = [] + if not out: return pl.DataFrame() return pl.concat(out, how="diagonal_relaxed") @@ -777,9 +826,8 @@ def sync_and_persist_minute( if extend_backward: # 向前扩展模式: 从本地最早数据往前补, 叠加已有数据避免缺口。 earliest_dt = _earliest_minute_datetime(repo) - # 按交易日换算自然日 (7/5 系数) - # ≤41 交易日: 不加余量, 确保落在单段 (57 自然日) 内 → 单次拉满 - # >41 交易日: +10 天余量覆盖节假日 + # 按交易日换算自然日 (7/5 系数)。>41 交易日时 +10 天余量覆盖节假日。 + # (分段由 sync_minute_batch 的 segment_trading_days 控制, 与此处的区间天数独立。) calendar_days = int(days * 7 / 5) + (10 if days > 41 else 0) if earliest_dt: end_time = earliest_dt @@ -805,33 +853,26 @@ def sync_and_persist_minute( default_rpm_when_unset=False, ) - df = sync_minute_batch(symbols, start_time=start_time, end_time=end_time, - batch_size=limit.batch, rpm=limit.rpm, - on_chunk_done=on_chunk_done) - if df.is_empty(): - return 0 + # 流式落盘: 每段拉完立即写盘, 内存峰值 = 单段 (而非全量)。 + # 全量攒内存曾导致 1 年全市场分钟 K OOM 卡死 (3 亿行 / 数十 GB)。 + minute_dir = repo.store.data_dir / "kline_minute" + written_box = [0] # list 闭包, 绕过 Python 闭包外层赋值 - # 按日期分区写: data/kline_minute/date={YYYY-MM-DD}/part.parquet - df = df.with_columns( - pl.col("datetime").dt.date().alias("_trade_date") + def _persist(seg_df: pl.DataFrame) -> None: + written_box[0] += _write_minute_partition(seg_df, minute_dir) + + segment_days = preferences.get_minute_sync_segment_days() + sync_minute_batch( + symbols, start_time=start_time, end_time=end_time, + batch_size=limit.batch, rpm=limit.rpm, + on_chunk_done=on_chunk_done, + segment_trading_days=segment_days, + on_segment=_persist, ) - written = 0 - for day_df in df.partition_by("_trade_date"): - trade_date = day_df["_trade_date"][0] - out = repo.store.data_dir / "kline_minute" / f"date={trade_date}" / "part.parquet" - out.parent.mkdir(parents=True, exist_ok=True) - if out.exists(): - existing = pl.read_parquet(out) - if "datetime" in existing.columns: - existing = existing.filter(pl.col("datetime").is_not_null()) - day_df = pl.concat([existing, day_df.drop("_trade_date")]).unique( - subset=["symbol", "datetime"], keep="last", - ) - else: - day_df = day_df.drop("_trade_date") - day_df = day_df.sort("symbol", "datetime") - _atomic_write_parquet(day_df, out) - written += day_df.height + + if written_box[0] == 0: + return 0 + written = written_box[0] # 刷新视图 try: diff --git a/backend/app/services/pipeline_jobs.py b/backend/app/services/pipeline_jobs.py index a7dc694..53f3397 100644 --- a/backend/app/services/pipeline_jobs.py +++ b/backend/app/services/pipeline_jobs.py @@ -26,7 +26,16 @@ JobStatus = Literal["pending", "running", "succeeded", "failed"] # 运行超过此秒数视为卡死(reload 后孤儿 task / 网络读无限阻塞等)。 # 由 reap_stale() 在 /run 和 /jobs/{id} 轮询端点检查 — 保证卡死后能自愈, # 无需用户再次点击「同步」。 -STALE_JOB_TIMEOUT_S = 600 +# +# 超时阈值按任务类型区分: +# - 普通任务(日K管道/扩展/修正/重算): 1200s (20 分钟) +# - 长任务(分钟K全市场同步,数据量是日K的 ~240 倍): 1800s (30 分钟) +# 分钟K即使流式落盘后仍可能跑十几到数十分钟(限速 sleep 是主因), +# 用 600s 会误杀正常任务并留下写盘僵尸线程。 +DEFAULT_JOB_TIMEOUT_S = 1200 +LONG_JOB_TIMEOUT_S = 1800 +# 向后兼容: 旧调用方引用 STALE_JOB_TIMEOUT_S +STALE_JOB_TIMEOUT_S = DEFAULT_JOB_TIMEOUT_S def _default_store_dir() -> Path: @@ -96,7 +105,7 @@ class JobStore: # ===== lifecycle ===== - def create(self) -> tuple[str, bool]: + def create(self, timeout_s: int = DEFAULT_JOB_TIMEOUT_S) -> tuple[str, bool]: """单飞创建任务。返回 (job_id, is_new)。 去重条件为 **pending ∨ running**(而非仅 running):`/run` 先 create() 再在 @@ -105,6 +114,9 @@ class JobStore: _active_id,导致两条全市场拉取同时读改写同一 parquet。纳入 pending 后该窗口关闭。 is_new=False 表示复用了已有活跃任务,调用方**不得**再调度新的后台任务。 + + timeout_s: reap_stale 判定卡死的阈值。普通任务默认 1200s; + 分钟K全市场同步等长任务传 LONG_JOB_TIMEOUT_S (1800s)。 """ with self._lock: if self._active_id: @@ -125,6 +137,7 @@ class JobStore: "duration_s": None, "result": None, "error": None, + "timeout_s": timeout_s, } self._active_id = job_id return job_id, True @@ -224,12 +237,16 @@ class JobStore: def active_id(self) -> str | None: return self._active_id - def reap_stale(self, timeout_s: int = STALE_JOB_TIMEOUT_S) -> None: - """回收运行超过 timeout_s 的卡死 running job(标记为 failed)。 + def reap_stale(self, timeout_s: int | None = None) -> None: + """回收运行超过阈值(卡死)的 running job(标记为 failed)。 在 /run 和 /jobs/{id} 轮询端点都会调用 — 保证卡死后任意轮询都能自愈, 无需用户再次手动触发同步。reload 后的孤儿 task(内存里已无 job 记录) 不在此处理:它们没有 active_id,只能靠 executor 线程自然结束或进程重启。 + + timeout_s: 显式覆盖。None 时用 job 自身 create() 时存的 timeout_s, + 缺失则回退 DEFAULT_JOB_TIMEOUT_S。分钟K长任务在 create 时存了更大阈值, + 不被普通任务的 1200s 误杀。 """ with self._lock: jid = self._active_id @@ -241,6 +258,8 @@ class JobStore: started = j.get("started_at") if not started: return + # 优先用显式传入, 其次 job 自身阈值, 最后默认值 + effective_timeout = timeout_s if timeout_s is not None else j.get("timeout_s", DEFAULT_JOB_TIMEOUT_S) # 时间计算放到锁外(避免 datetime 解析持锁)。 # started_at 形如 "2026-07-04T12:00:00Z"(start() 用 datetime.utcnow 存)。 # 两端都用 timezone-aware UTC 比较,避免 naive/aware 混用导致 TypeError。 @@ -249,9 +268,9 @@ class JobStore: elapsed = (datetime.now(start_dt.tzinfo) - start_dt).total_seconds() except Exception: # noqa: BLE001 return - if elapsed > timeout_s: - logger.warning("reap_stale: 强制取消卡死 job %s (已运行 %.0fs)", - jid, elapsed) + if elapsed > effective_timeout: + logger.warning("reap_stale: 强制取消卡死 job %s (已运行 %.0fs, 阈值 %ss)", + jid, elapsed, effective_timeout) self.fail(jid, f"超时自动取消 (运行 {int(elapsed)}s, 疑似卡死)") # 强制释放重任务锁: 卡死的线程无法被中断, 锁永远不会自然释放。 # job 已标记 failed, 即使僵尸线程后续写入 parquet, 下次拉取会覆盖, 安全。 diff --git a/backend/app/services/preferences.py b/backend/app/services/preferences.py index 84cdbea..e929eba 100644 --- a/backend/app/services/preferences.py +++ b/backend/app/services/preferences.py @@ -99,6 +99,16 @@ def get_minute_sync_days() -> int: return max(1, min(30, load().get("minute_sync_days", 5))) +def get_minute_sync_segment_days() -> int: + """分钟 K 拉取的单段大小(交易日)。默认 20,范围 [5, 30]。 + + 每段拉完后立即落盘(流式),避免全量攒内存导致 OOM。 + 段越小内存峰值越低但总耗时越长(限速 sleep 随段数线性增加); + 物理上限 ~41 交易日(TickFlow 单次 10000 根 / 一天 241 根 ≈ 41 天),max=30 留出余量。 + """ + return max(5, min(30, load().get("minute_sync_segment_days", 20))) + + # ===== 数据源选择 (默认 TickFlow;第一阶段仅日K切换入口) ===== _ALLOWED_DATA_PROVIDERS = {"tickflow"} diff --git a/backend/tests/backtest/test_minute_fill.py b/backend/tests/backtest/test_minute_fill.py new file mode 100644 index 0000000..a71c369 --- /dev/null +++ b/backend/tests/backtest/test_minute_fill.py @@ -0,0 +1,104 @@ +"""分钟K精确成交 (_resolve_minute_fill / _load_minute_for_fills) 回归测试。 + +背景: 原实现用 df.to_numpy() 转 structured array 再按字段名索引 (arr["open"])。 +当 DataFrame 含 datetime 列 + float 列时, to_numpy() 退化为 dtype=object 的二维 +数组, 字段名索引抛 IndexError: "only integers, slices... are valid indices"。 +开启 minute_fill 的回测从未成功跑通过。此测试锁定该 bug 不再复发。 +""" +from __future__ import annotations + +from datetime import date, datetime + +import polars as pl + +from app.backtest.engine import BacktestEngine + + +def _sample_minute_df(symbol: str = "000001.SZ") -> pl.DataFrame: + """构造一份带 datetime 列 + float 列的分钟K (复现 to_numpy 退化的场景)。""" + base = datetime(2024, 1, 2, 9, 31) + return pl.DataFrame({ + "symbol": [symbol] * 4, + "datetime": [base.replace(hour=h, minute=m) for h, m in + [(9, 31), (10, 0), (14, 0), (14, 57)]], + "open": [10.0, 10.5, 10.8, 10.6], + "high": [10.6, 10.7, 10.9, 10.7], + "low": [9.9, 10.4, 10.7, 10.5], + "close": [10.2, 10.6, 10.85, 10.65], + "volume": [100, 200, 150, 120], + "amount": [1020.0, 2120.0, 1627.0, 1278.0], + }) + + +def test_resolve_minute_fill_with_mixed_columns_no_index_error(): + """混合列类型 (datetime + float) 不再抛 IndexError。 + + 这是原 bug 的精确复现点: 旧实现 _resolve_minute_fill 接收 ndarray, + arr["open"] 在 object 数组上会炸。现在接受 DataFrame, 按列取值。 + """ + mdf = _sample_minute_df() + # 三种分支都应正常返回, 不抛 IndexError + assert BacktestEngine._resolve_minute_fill(mdf, ref_price=10.5, side="buy") is not None + assert BacktestEngine._resolve_minute_fill(mdf, ref_price=10.5, side="sell") is not None + # 无参考线 → VWAP 分支 + vwap = BacktestEngine._resolve_minute_fill(mdf, ref_price=None, side="buy") + assert vwap is not None and vwap > 0 + + +def test_resolve_minute_fill_buy_cross_above_ref(): + """买入: 价格涨破参考线 → 开盘已高于则按开盘。""" + mdf = _sample_minute_df() + # ref=9.5, 开盘 10.0 已高于 → 按开盘 + assert BacktestEngine._resolve_minute_fill(mdf, 9.5, "buy") == 10.0 + + +def test_resolve_minute_fill_sell_cross_below_ref(): + """卖出: 价格跌破参考线 → 开盘已低于则按开盘。""" + mdf = _sample_minute_df() + # ref=10.5, 开盘 10.0 已低于 → 按开盘 + assert BacktestEngine._resolve_minute_fill(mdf, 10.5, "sell") == 10.0 + + +def test_resolve_minute_fill_vwap(): + """无参考线 → VWAP = 总成交额 / 总成交量。""" + mdf = _sample_minute_df() + total_amt = 1020.0 + 2120.0 + 1627.0 + 1278.0 + total_vol = 100 + 200 + 150 + 120 + expected = total_amt / total_vol + assert BacktestEngine._resolve_minute_fill(mdf, None, "buy") == expected + + +def test_resolve_minute_fill_empty_returns_none(): + """空 DataFrame → None (降级到日K口径)。""" + assert BacktestEngine._resolve_minute_fill(pl.DataFrame(), None, "buy") is None + + +class _FakeRepo: + """最小 repo 桩: get_minute_range 直接返回预构造的混合列 DataFrame。""" + + def __init__(self, df: pl.DataFrame) -> None: + self._df = df + + def get_minute_range(self, symbols, start, end, asset_type="stock") -> pl.DataFrame: # noqa: ANN001 + return self._df + + +def test_load_minute_for_fills_returns_dataframe_dict(): + """_load_minute_for_fills 返回 {(symbol, date_str): DataFrame}, 而非 ndarray。 + + 锁定: cache 值类型必须是 pl.DataFrame (旧实现返回的对象后续被 .to_numpy() + 退化成 object 数组触发 bug)。 + """ + df = _sample_minute_df() + repo = _FakeRepo(df) + result = BacktestEngine._load_minute_for_fills( + repo, ["000001.SZ"], {"2024-01-02"}, "stock", + ) + assert ("000001.SZ", "2024-01-02") in result + val = result[("000001.SZ", "2024-01-02")] + # 关键断言: 返回的是 DataFrame, 可直接喂给 _resolve_minute_fill + assert isinstance(val, pl.DataFrame) + assert not val.is_empty() + # 端到端: load → resolve 不抛异常 + price = BacktestEngine._resolve_minute_fill(val, None, "buy") + assert price is not None and price > 0 diff --git a/frontend/src/components/data/MinuteSyncConfig.tsx b/frontend/src/components/data/MinuteSyncConfig.tsx index 31d50f6..bf03056 100644 --- a/frontend/src/components/data/MinuteSyncConfig.tsx +++ b/frontend/src/components/data/MinuteSyncConfig.tsx @@ -11,17 +11,20 @@ export function MinuteSyncConfig({ caps, onJobStart }: { caps: { label: string; queryFn: api.preferences, }) const update = useMutation({ - mutationFn: ({ enabled, days }: { enabled: boolean; days: number }) => - api.updateMinuteSync(enabled, days), + mutationFn: ({ enabled, days, segmentDays }: { enabled: boolean; days: number; segmentDays?: number }) => + api.updateMinuteSync(enabled, days, segmentDays), onSuccess: () => qc.invalidateQueries({ queryKey: QK.preferences }), }) const hasMinuteCap = !!caps?.capabilities?.['kline.minute.batch'] const enabled = prefs.data?.minute_sync_enabled ?? false const days = prefs.data?.minute_sync_days ?? 5 + const segmentDays = prefs.data?.minute_sync_segment_days ?? 20 const [localDays, setLocalDays] = useState(days) + const [localSegment, setLocalSegment] = useState(segmentDays) useEffect(() => { setLocalDays(days) }, [days]) + useEffect(() => { setLocalSegment(segmentDays) }, [segmentDays]) const handleToggle = () => { if (!hasMinuteCap) return @@ -34,6 +37,13 @@ export function MinuteSyncConfig({ caps, onJobStart }: { caps: { label: string; update.mutate({ enabled, days: clamped }) } + // 段大小: 交易日/段, 步进 ±5, 范围 [5, 30]。越小越省内存但越慢。 + const setSegment = (v: number) => { + const clamped = Math.max(5, Math.min(30, v)) + setLocalSegment(clamped) + update.mutate({ enabled, days: localDays, segmentDays: clamped }) + } + // 清空分钟K数据 (二次确认) const [confirmClear, setConfirmClear] = useState(false) const clearMutation = useMutation({ @@ -48,8 +58,8 @@ export function MinuteSyncConfig({ caps, onJobStart }: { caps: { label: string; const [fetchingMode, setFetchingMode] = useState<'' | '40d' | '1y'>('') const handleFetch = (mode: '40d' | '1y') => { if (!hasMinuteCap) return - // 两个按钮都用向前扩展模式: 从本地最早数据往前补, 叠加避免缺口 - const fetchDays = mode === '40d' ? 40 : 365 + // 单次获取 = 按「分段大小」拉一段 (向前扩展); 1年 = 拉365天按分段切多段 + const fetchDays = mode === '40d' ? localSegment : 365 setFetchingMode(mode) api.syncMinute(fetchDays, true).then((res) => { qc.invalidateQueries({ queryKey: QK.pipelineJobs }) @@ -61,7 +71,8 @@ export function MinuteSyncConfig({ caps, onJobStart }: { caps: { label: string; return (
- {/* 第 1 行: 自动同步开关 + 天数 */} + {/* 区块 A: 自动同步 (盘后定时拉取的偏好设置) */} +
- {enabled ? '盘后自动同步' : '已关闭'} + 自动同步{enabled ? '已开启' : '已关闭'}
@@ -104,8 +115,45 @@ export function MinuteSyncConfig({ caps, onJobStart }: { caps: { label: string;
- {/* 第 2 行: 两个手动获取按钮 (40天快速 / 1年分段) */} -
+ {/* 分段大小: 控制单次 SDK 请求覆盖的交易日数。每段拉完即落盘,避免全量攒内存 OOM。 + 「往前获取」与「获取 1 年」共用此设置 (两者都经过 sync_and_persist_minute)。 */} +
+
+ 分段大小 + 内存优化 +
+
+
+ +
+ {localSegment} +
+ +
+ 交易日/段 +
+
+
+ 每段拉完即写盘,避免内存堆积。越小越省内存但越慢,默认 20 平衡。 +
+
+ + {/* 区块 B: 手动获取 (一次性操作, 独立于上方自动同步开关) */} +
+
+ + 手动获取 + 不受自动同步开关影响 +
+
+
+
+ A股标的 · 前复权价格 · 从本地最早数据向前叠加 ·{' '} + 均按上方「分段大小」分段拉取、每段即落盘 +
- {/* 第 3 行: 清空 */} + {/* 区块 C: 清空 (危险操作, 独立分隔) */} - {/* 说明 */} -
- A股标的 · 前复权价格 · 均从本地最早数据向前叠加 ·{' '} - 单次拉满约 40 个交易日,{' '} - 1 年按月分段 (速度较慢) -
- {/* 清空确认弹窗 */} {confirmClear && (
diff --git a/frontend/src/lib/api.ts b/frontend/src/lib/api.ts index 66e1867..6b550b0 100644 --- a/frontend/src/lib/api.ts +++ b/frontend/src/lib/api.ts @@ -802,6 +802,7 @@ export interface Preferences { indices_nav_pinned: boolean minute_sync_enabled: boolean minute_sync_days: number + minute_sync_segment_days: number daily_data_provider?: string adj_factor_provider?: string minute_data_provider?: string @@ -942,10 +943,14 @@ export const api = { '/api/settings/preferences/data-providers', { method: 'PUT', body: JSON.stringify(cfg) }, ), - updateMinuteSync: (enabled: boolean, days: number) => + updateMinuteSync: (enabled: boolean, days: number, segmentDays?: number) => request('/api/settings/preferences/minute-sync', { method: 'PUT', - body: JSON.stringify({ minute_sync_enabled: enabled, minute_sync_days: days }), + body: JSON.stringify({ + minute_sync_enabled: enabled, + minute_sync_days: days, + ...(segmentDays != null ? { minute_sync_segment_days: segmentDays } : {}), + }), }), updatePipelinePullTypes: (cfg: Partial>) => request<{