diff --git a/backend/app/api/indices.py b/backend/app/api/indices.py index e78c2a2..febd6bb 100644 --- a/backend/app/api/indices.py +++ b/backend/app/api/indices.py @@ -113,7 +113,7 @@ def get_index_minute( repo = request.app.state.repo info = _index_info(repo, symbol) day = trade_date or date.today() - df = kline_sync.fetch_minute_single(symbol, day) + df = kline_sync.fetch_minute_single(symbol, day, asset_type="index") return { "symbol": symbol, "name": info.get("name"), diff --git a/backend/app/api/kline.py b/backend/app/api/kline.py index b4cfe1a..2018d6f 100644 --- a/backend/app/api/kline.py +++ b/backend/app/api/kline.py @@ -540,14 +540,38 @@ def get_minute_batch(request: Request, body: dict): start_time = datetime(trade_date.year, trade_date.month, trade_date.day, 9, 25, 0) end_time = datetime(trade_date.year, trade_date.month, trade_date.day, 15, 5, 0) lim = capset.limits(Cap.KLINE_MINUTE_BATCH) - live_df = kline_sync.sync_minute_batch( - incomplete, - start_time=start_time, - end_time=end_time, - batch_size=lim.batch if lim else None, - rpm=lim.rpm if lim else None, - ) - if not live_df.is_empty(): + # etf_set 已在上方获取, 直接复用 — 按 asset_type 拆分调用 sync_minute_batch + # (自定义源 / TickFlow 路由均依赖 asset_type 正确传递) + # 契约: 本端点只接受 stock/ETF (指数分钟K走 /api/index/minute 独立路径), + # 故两分支已覆盖全部 incomplete。若未来放开指数支持, 需额外加 index 分支 + # 以避免被误路由为 stock。 + stock_incomplete = [s for s in incomplete if s not in etf_set] + etf_incomplete = [s for s in incomplete if s in etf_set] + live_parts: list[pl.DataFrame] = [] + if stock_incomplete: + df_s = kline_sync.sync_minute_batch( + stock_incomplete, + start_time=start_time, + end_time=end_time, + batch_size=lim.batch if lim else None, + rpm=lim.rpm if lim else None, + asset_type="stock", + ) + if not df_s.is_empty(): + live_parts.append(df_s) + if etf_incomplete: + df_e = kline_sync.sync_minute_batch( + etf_incomplete, + start_time=start_time, + end_time=end_time, + batch_size=lim.batch if lim else None, + rpm=lim.rpm if lim else None, + asset_type="etf", + ) + if not df_e.is_empty(): + live_parts.append(df_e) + if live_parts: + live_df = pl.concat(live_parts, how="diagonal_relaxed") for sym in incomplete: sub = live_df.filter(pl.col("symbol") == sym).sort("datetime") if not sub.is_empty(): @@ -594,7 +618,7 @@ def get_minute( if trade_date is None: # 本地无任何分钟K,尝试从 TickFlow 拉取当天 trade_date = cn_today() - df = kline_sync.fetch_minute_single(symbol, trade_date) + df = kline_sync.fetch_minute_single(symbol, trade_date, asset_type=asset_type) price_limit = _get_price_limit_info( repo, symbol, trade_date, asset_type, stock_name, ) @@ -636,7 +660,7 @@ def get_minute( } # 本地不完整或无数据 → 从 TickFlow 实时拉取 - live_df = kline_sync.fetch_minute_single(symbol, trade_date) + live_df = kline_sync.fetch_minute_single(symbol, trade_date, asset_type=asset_type) return { "symbol": symbol, "name": stock_name, "stock_info": stock_info, "date": str(trade_date), "rows": live_df.to_dicts(), diff --git a/backend/app/data_providers/base.py b/backend/app/data_providers/base.py index c17ca21..53b59f4 100644 --- a/backend/app/data_providers/base.py +++ b/backend/app/data_providers/base.py @@ -6,6 +6,7 @@ backtests stay data-source agnostic. """ from __future__ import annotations +from collections.abc import Callable from dataclasses import dataclass from datetime import datetime from typing import Literal, Protocol @@ -57,8 +58,14 @@ class MarketDataProvider(Protocol): end_time: datetime | None, asset_type: AssetType, freq: str = "1m", + on_chunk_done: Callable[[int, int], None] | None = None, ) -> pl.DataFrame: - """Return normalized minute K rows. Implementations may return empty.""" + """Return normalized minute K rows. Implementations may return empty. + + on_chunk_done 契约: provider 实现内部以 2 参 (cur, total) 调用; + 3 参 seg_label 适配由 kline_sync._try_custom_minute 包装层负责, + 不应泄漏到 provider 契约层。 + """ def get_realtime( self, diff --git a/backend/app/data_providers/custom/provider.py b/backend/app/data_providers/custom/provider.py index 53b258e..1f80e91 100644 --- a/backend/app/data_providers/custom/provider.py +++ b/backend/app/data_providers/custom/provider.py @@ -3,6 +3,7 @@ from __future__ import annotations import logging import os +from collections.abc import Callable from datetime import datetime from pathlib import Path from typing import Any @@ -11,6 +12,7 @@ import httpx import polars as pl from app.config import settings +from app.data_providers.base import AssetType from app.data_providers.custom.config import CustomSourceConfig, DatasetConfig from app.data_providers.custom.mapper import apply_transforms, datetime_payload, extract_rows, map_rows from app.data_providers.normalizer import normalize_adj_factors, normalize_daily @@ -109,8 +111,9 @@ class GenericHTTPProvider: symbols: list[str], start_time: datetime | None, end_time: datetime | None, - asset_type: str = "stock", # noqa: ARG002 - on_chunk_done=None, + asset_type: AssetType = "stock", # noqa: ARG002 + freq: str = "1m", # noqa: ARG002 + on_chunk_done: Callable[[int, int], None] | None = None, ) -> pl.DataFrame: cfg = self._dataset("minute") frames: list[pl.DataFrame] = [] diff --git a/backend/app/data_providers/tickflow_provider.py b/backend/app/data_providers/tickflow_provider.py index a086949..c69bd73 100644 --- a/backend/app/data_providers/tickflow_provider.py +++ b/backend/app/data_providers/tickflow_provider.py @@ -2,6 +2,7 @@ from __future__ import annotations import logging +from collections.abc import Callable from datetime import datetime import polars as pl @@ -97,8 +98,9 @@ class TickFlowProvider: symbols: list[str], start_time: datetime | None, end_time: datetime | None, - asset_type: AssetType, # noqa: ARG002 + asset_type: AssetType = "stock", # noqa: ARG002 freq: str = "1m", # noqa: ARG002 + on_chunk_done: Callable[[int, int], None] | None = None, # noqa: ARG002 ) -> pl.DataFrame: # Existing minute sync remains in app.services.kline_sync for now. return pl.DataFrame() diff --git a/backend/app/plugins/stocksdk/provider.py b/backend/app/plugins/stocksdk/provider.py index 6922055..c3481a4 100644 --- a/backend/app/plugins/stocksdk/provider.py +++ b/backend/app/plugins/stocksdk/provider.py @@ -10,11 +10,13 @@ Original implementation by @forrany (PR #57), migrated to plugin architecture. from __future__ import annotations import logging +from collections.abc import Callable from dataclasses import dataclass, field from datetime import datetime import polars as pl +from app.data_providers.base import AssetType from app.data_providers.normalizer import normalize_adj_factors, normalize_daily from app.plugins.stocksdk import bridge from app.tickflow.rate_limits import chunked @@ -135,13 +137,13 @@ class StockSDKProvider: symbols: list[str], start_time: datetime | None, end_time: datetime | None, - asset_type: str = "stock", # noqa: ARG002 - on_chunk_done=None, - freq: str = "5m", + asset_type: AssetType = "stock", # noqa: ARG002 + freq: str = "1m", + on_chunk_done: Callable[[int, int], None] | None = None, ) -> pl.DataFrame: if not symbols: return pl.DataFrame() - period = "".join(ch for ch in str(freq) if ch.isdigit()) or "5" + period = "".join(ch for ch in str(freq) if ch.isdigit()) or "1" logger.info("stock-sdk minute 拉取开始(%d symbols, period=%s)", len(symbols), period) frames: list[pl.DataFrame] = [] chunks = chunked(symbols, _BATCH) diff --git a/backend/app/services/kline_sync.py b/backend/app/services/kline_sync.py index e2ba0db..2489718 100644 --- a/backend/app/services/kline_sync.py +++ b/backend/app/services/kline_sync.py @@ -9,10 +9,11 @@ from __future__ import annotations import logging from collections.abc import Callable -from datetime import datetime, timedelta +from datetime import date, datetime, timedelta import polars as pl +from app.data_providers.base import AssetType from app.indicators.pipeline import filter_halt_days from app.market_time import cn_now from app.services import preferences @@ -528,6 +529,56 @@ def _write_minute_partition(df: pl.DataFrame, minute_dir) -> int: return written +def _try_custom_minute( + symbols: list[str], + start_time: datetime | None, + end_time: datetime | None, + asset_type: AssetType, + freq: str = "1m", + on_chunk_done: Callable[[int, int, str], None] | None = None, +) -> tuple[pl.DataFrame | None, bool]: + """尝试从自定义分钟源拉取。返回 (df, should_fallback_to_tickflow)。 + + 返回契约: + (None, True) → 未配自定义源 / 未配 minute dataset / 自定义源异常 → 走 TickFlow + (df, False) → 自定义源成功(含空 df) → 直接用, 不回退 + + 降级策略 (C): 自定义源异常时无条件 fall through 到 TickFlow, + 由 TickFlow 路径自身 try/except 兜底。Pro+ 用户 TickFlow 成功返回数据, + None 档用户 TickFlow 失败返回空。不显式判断 tier, 避免 #126 augmented + capability 逻辑干扰。 + + on_chunk_done 适配: 上层回调是 3 参 (cur, total, seg_label), provider + 实现内部以 2 参 (cur, total) 调用。这里包装一层, provider 调 2 参时补 + 默认 seg_label="custom" 转发给上层, 保证进度展示不降级。 + """ + provider_name = preferences.get_minute_data_provider() + if provider_name == "tickflow": + return (None, True) + from app.data_providers import custom as custom_sources + if not custom_sources.provider_has_dataset(provider_name, "minute"): + return (None, True) + provider = custom_sources.get_provider(provider_name) + + # 包装 on_chunk_done: provider 调 2 参 → 补 seg_label="custom" → 转发上层 3 参 + wrapped_cb: Callable[[int, int], None] | None = None + if on_chunk_done is not None: + def _wrapped_cb(cur: int, total: int) -> None: + on_chunk_done(cur, total, "custom") + wrapped_cb = _wrapped_cb + + try: + df = provider.get_minute( + symbols, start_time=start_time, end_time=end_time, + asset_type=asset_type, freq=freq, on_chunk_done=wrapped_cb, + ) + return (df, False) + except Exception as e: # noqa: BLE001 + logger.warning("custom minute provider %s failed, falling back to TickFlow: %s", + provider_name, e) + return (None, True) + + def sync_minute_batch( symbols: list[str], start_time: datetime | None = None, @@ -538,6 +589,7 @@ def sync_minute_batch( on_chunk_done: Callable[[int, int, str], None] | None = None, segment_trading_days: int = 20, on_segment: Callable[[pl.DataFrame], None] | None = None, + asset_type: AssetType = "stock", ) -> pl.DataFrame: """批量拉取多股分钟 K。 @@ -555,16 +607,15 @@ def sync_minute_batch( 不进入全局 out → 内存峰值从「全量」降到「单段」。适用于 sync_and_persist_minute。 不传时 (如 get_minute_batch 的实时补拉) 保持原契约: 累积进 out 末尾一次性返回。 """ - # 自定义数据源分流: minute provider - provider_name = preferences.get_minute_data_provider() - if provider_name != "tickflow": - from app.data_providers import custom as custom_sources - if custom_sources.provider_has_dataset(provider_name, "minute"): - provider = custom_sources.get_provider(provider_name) - return provider.get_minute( - symbols, start_time=start_time, end_time=end_time, on_chunk_done=on_chunk_done, - ) - # 未配置 minute → 回退 TickFlow + df, fallback = _try_custom_minute( + symbols, start_time=start_time, end_time=end_time, + asset_type=asset_type, freq="1m", on_chunk_done=on_chunk_done, + ) + if not fallback: + # _try_custom_minute 成功时返回 (df, False), df 非 None; + # fallback=True 时才返回 (None, True)。故此处 df 必非 None。 + # (旧版 `df or pl.DataFrame()` 触发 polars DataFrame __bool__ TypeError, 已修。) + return df if df is not None else pl.DataFrame() tf = get_client() @@ -572,7 +623,7 @@ def sync_minute_batch( # 按 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]] = [] + time_segments: list[tuple[datetime | None, datetime | None]] = [] if start_time and end_time: seg_start = start_time while seg_start < end_time: @@ -589,10 +640,10 @@ def sync_minute_batch( # 段内累积: 每段拉完即 flush, 避免全量攒内存 (OOM 根因) seg_out: list[pl.DataFrame] = [] - for seg_idx, (seg_start, seg_end) in enumerate(time_segments): + for seg_idx, (cur_start, cur_end) in enumerate(time_segments): # 当前的日期段描述 (供进度展示) - if seg_start and seg_end: - seg_label = f"{seg_start.strftime('%m-%d')}~{seg_end.strftime('%m-%d')}" + if cur_start and cur_end: + seg_label = f"{cur_start.strftime('%m-%d')}~{cur_end.strftime('%m-%d')}" else: seg_label = "最新" seg_total = len(time_segments) @@ -601,11 +652,11 @@ def sync_minute_batch( sleep_between_batches(step, rpm) step += 1 try: - if seg_start and seg_end: + if cur_start and cur_end: raw = tf.klines.batch( chunk, period="1m", - start_time=_datetime_to_ms(seg_start), - end_time=_datetime_to_ms(seg_end), + start_time=_datetime_to_ms(cur_start), + end_time=_datetime_to_ms(cur_end), count=10000, adjust="forward", as_dataframe=True, show_progress=False, @@ -737,7 +788,11 @@ def fetch_intraday_monitor_batch( return pl.concat(frames, how="diagonal_relaxed") if frames else pl.DataFrame() -def fetch_minute_single(symbol: str, trade_date: date) -> pl.DataFrame: +def fetch_minute_single( + symbol: str, + trade_date: date, + asset_type: AssetType = "stock", +) -> pl.DataFrame: """实时拉取单股单日分钟 K(不写入本地)。优先自定义分钟源, 回退 TickFlow。""" from datetime import datetime start_time = datetime(trade_date.year, trade_date.month, trade_date.day, 9, 25, 0) @@ -745,13 +800,13 @@ def fetch_minute_single(symbol: str, trade_date: date) -> pl.DataFrame: # 自定义数据源分流: 与 sync_minute_batch 一致, 配了自定义分钟源时走 custom provider, # 避免无 TickFlow Pro+ 权限的用户分时图首次打开(本地无数据)时补拉失败返回空。 - provider_name = preferences.get_minute_data_provider() - if provider_name != "tickflow": - from app.data_providers import custom as custom_sources - if custom_sources.provider_has_dataset(provider_name, "minute"): - provider = custom_sources.get_provider(provider_name) - return provider.get_minute([symbol], start_time=start_time, end_time=end_time) - # 未配置 minute dataset → 回退 TickFlow + df, fallback = _try_custom_minute( + [symbol], start_time=start_time, end_time=end_time, + asset_type=asset_type, freq="1m", + ) + if not fallback: + # 见 sync_minute_batch 同分支注释: df 在此必非 None。 + return df if df is not None else pl.DataFrame() tf = get_client() try: @@ -975,6 +1030,7 @@ def sync_and_persist_minute( on_chunk_done=on_chunk_done, segment_trading_days=segment_days, on_segment=_persist, + asset_type="stock", ) if written_box[0] == 0: diff --git a/backend/tests/test_minute_routing.py b/backend/tests/test_minute_routing.py new file mode 100644 index 0000000..f4e6c10 --- /dev/null +++ b/backend/tests/test_minute_routing.py @@ -0,0 +1,312 @@ +"""自定义分钟数据源路由回归测试。 + +对应设计文档 §4 测试矩阵 (docs/superpowers/specs/2026-07-18-minute-provider-unification-design.md)。 + +覆盖三个阻断问题: +1. stock-sdk 默认 freq 漂移 (5m → 1m) +2. 自定义源异常直接 500 (无 try/except) +3. 插件化路由重复 + asset_type 未透传 + +mock 范式沿用 test_stocksdk_provider.py (monkeypatch 模块属性)。 +""" +from __future__ import annotations + +from datetime import date, datetime +from unittest.mock import MagicMock + +import httpx +import polars as pl + +from app.plugins.stocksdk import provider as sp +from app.plugins.stocksdk.provider import StockSDKProvider +from app.services import kline_sync + + +# ---------- 辅助 ---------- + +def _mock_minute_df(symbol: str = "600519.SH") -> pl.DataFrame: + """构造非空分钟 K df, 用于 mock provider.get_minute 返回值。""" + return pl.DataFrame({ + "symbol": [symbol], + "datetime": [datetime(2026, 1, 15, 9, 35, 0)], + "open": [100.0], + "high": [101.0], + "low": [99.5], + "close": [100.5], + "volume": [1000.0], + "amount": [100500.0], + }) + + +def _setup_custom_provider(monkeypatch, provider: object, has_dataset: bool = True) -> None: + """统一 mock 自定义分钟源路由前置: preferences + provider_has_dataset + get_provider。 + + - preferences.get_minute_data_provider → "mock_src" + - custom.provider_has_dataset → has_dataset + - custom.get_provider → provider + """ + monkeypatch.setattr( + kline_sync.preferences, + "get_minute_data_provider", + lambda: "mock_src", + ) + monkeypatch.setattr( + "app.data_providers.custom.provider_has_dataset", + lambda name, ds: has_dataset, + ) + monkeypatch.setattr( + "app.data_providers.custom.get_provider", + lambda name: provider, + ) + + +# ---------- 测试 1: 自定义源成功返回 1 分钟 K ---------- + +def test_custom_minute_provider_returns_1m_k(monkeypatch): + """§4 测试 1: 自定义源成功返回 1m K, 且 provider 收到 freq="1m"。""" + spy = MagicMock(return_value=_mock_minute_df()) + mock_provider = MagicMock() + mock_provider.get_minute = spy + _setup_custom_provider(monkeypatch, mock_provider, has_dataset=True) + + df, fallback = kline_sync._try_custom_minute( + ["600519.SH"], + datetime(2026, 1, 15, 9, 25, 0), + datetime(2026, 1, 15, 15, 5, 0), + asset_type="stock", + ) + + assert fallback is False + assert df is not None + assert not df.is_empty() + # spy 收到 freq="1m" 和 asset_type="stock" + spy.assert_called_once() + _, kwargs = spy.call_args + assert kwargs.get("freq") == "1m" + assert kwargs.get("asset_type") == "stock" + + +# ---------- 测试 2: stock-sdk 收到 freq=1m → bridge job period="1" ---------- + +def test_stocksdk_get_minute_receives_freq_1m(monkeypatch): + """§4 测试 2: StockSDKProvider.get_minute(freq="1m") → bridge job period == "1"。 + + bridge.mjs opMinute 用 String(period), 1m → "1"。 + """ + captured: dict = {} + + def fake_run_job(job, timeout=None): + captured["job"] = job + # 返回空结果, 测试只验证 job.period + return {"ok": True, "op": job["op"], "rows": {}} + + monkeypatch.setattr(sp.bridge, "run_job", fake_run_job) + + StockSDKProvider().get_minute( + ["600519.SH"], None, None, freq="1m", + ) + + assert captured["job"]["op"] == "minute" + assert captured["job"]["period"] == "1" + + +# ---------- 测试 3: 自定义源异常 + TickFlow 也失败 → 返回空 (非 500) ---------- + +def test_custom_provider_exception_no_500(monkeypatch): + """§4 测试 3: 自定义源抛异常 + TickFlow 也失败, + fetch_minute_single / sync_minute_batch 返回空 df。 + """ + # 自定义源抛异常 + mock_provider = MagicMock() + mock_provider.get_minute.side_effect = httpx.TimeoutException("timeout") + _setup_custom_provider(monkeypatch, mock_provider, has_dataset=True) + + # mock get_client 返回 mock client, 其 klines.batch raise (TickFlow 也失败) + mock_tf = MagicMock() + mock_tf.klines.batch.side_effect = Exception("tickflow fail") + monkeypatch.setattr(kline_sync, "get_client", lambda: mock_tf) + + # fetch_minute_single: 自定义源异常 → fall through → TickFlow 异常 → 返回空 + df_single = kline_sync.fetch_minute_single( + "600519.SH", date(2026, 1, 15), asset_type="stock", + ) + assert isinstance(df_single, pl.DataFrame) + assert df_single.is_empty() + + # sync_minute_batch: 同一路径, 返回空 + df_batch = kline_sync.sync_minute_batch( + ["600519.SH"], + start_time=datetime(2026, 1, 15, 9, 25, 0), + end_time=datetime(2026, 1, 15, 15, 5, 0), + asset_type="stock", + ) + assert isinstance(df_batch, pl.DataFrame) + assert df_batch.is_empty() + + +# ---------- 测试 4: 未配 minute dataset → 回退 TickFlow ---------- + +def test_provider_without_minute_dataset_fallback(monkeypatch): + """§4 测试 4: provider_has_dataset 返回 False → (None, True) 回退 TickFlow。""" + mock_provider = MagicMock() + _setup_custom_provider(monkeypatch, mock_provider, has_dataset=False) + + df, fallback = kline_sync._try_custom_minute( + ["600519.SH"], None, None, asset_type="stock", + ) + + assert fallback is True + assert df is None + # provider.get_minute 不应被调用 (回退决策在前) + mock_provider.get_minute.assert_not_called() + + +# ---------- 测试 5: asset_type 透传到 provider ---------- + +def test_asset_type_threaded_to_provider(monkeypatch): + """§4 测试 5: stock/etf/index asset_type 透传到 provider.get_minute。""" + spy = MagicMock(return_value=_mock_minute_df()) + mock_provider = MagicMock() + mock_provider.get_minute = spy + _setup_custom_provider(monkeypatch, mock_provider, has_dataset=True) + + # 三次调用不同 asset_type + kline_sync.fetch_minute_single("600519.SH", date(2026, 1, 15), asset_type="stock") + kline_sync.fetch_minute_single("510300.SH", date(2026, 1, 15), asset_type="etf") + kline_sync.fetch_minute_single("000001.SH", date(2026, 1, 15), asset_type="index") + + # spy 被调 3 次, 每次收到对应 asset_type + assert spy.call_count == 3 + received_assets = [call.kwargs.get("asset_type") for call in spy.call_args_list] + assert received_assets == ["stock", "etf", "index"] + + +# ---------- 测试 6: 自定义源成功时不调 TickFlow ---------- + +def test_custom_success_skips_tickflow(monkeypatch): + """§4 测试 6: fetch_minute_single 自定义源成功 → 不调 get_client。""" + expected_df = _mock_minute_df() + mock_provider = MagicMock() + mock_provider.get_minute.return_value = expected_df + _setup_custom_provider(monkeypatch, mock_provider, has_dataset=True) + + # get_client 设为 spy, 若被调说明路由失败 + get_client_spy = MagicMock(name="get_client_spy") + monkeypatch.setattr(kline_sync, "get_client", get_client_spy) + + df = kline_sync.fetch_minute_single( + "600519.SH", date(2026, 1, 15), asset_type="stock", + ) + + # 返回的是 mock provider 的 df + assert df is expected_df + # TickFlow 路径未进入 + get_client_spy.assert_not_called() + + +# ---------- 测试 7: sync_minute_batch 自定义源成功直接返回 ---------- + +def test_sync_minute_batch_custom_success_returns_directly(monkeypatch): + """§4 测试 7: sync_minute_batch 自定义源成功 → 直接 return, 不走 segment 逻辑。""" + expected_df = _mock_minute_df() + mock_provider = MagicMock() + mock_provider.get_minute.return_value = expected_df + _setup_custom_provider(monkeypatch, mock_provider, has_dataset=True) + + get_client_spy = MagicMock(name="get_client_spy") + monkeypatch.setattr(kline_sync, "get_client", get_client_spy) + + df = kline_sync.sync_minute_batch( + ["600519.SH"], + start_time=datetime(2026, 1, 15, 9, 25, 0), + end_time=datetime(2026, 1, 15, 15, 5, 0), + asset_type="stock", + ) + + # 返回 mock provider 的 df, 不走 segment + assert df is expected_df + get_client_spy.assert_not_called() + + +# ---------- 测试 8: on_chunk_done 包装 (2参 → 3参补 seg_label='custom') ---------- + +def test_on_chunk_done_wrapped_to_3_args(monkeypatch): + """on_chunk_done 包装: provider 内部以 2 参 (cur, total) 调用 → + 上层 3 参 (cur, total, seg_label) spy 收到 seg_label='custom'。 + + 设计文档 §2: 保证自定义源路径进度展示不降级 (与 TickFlow 路径 3 参回调对齐)。 + """ + upper_cb = MagicMock(name="upper_3arg_cb") + + def provider_get_minute_side_effect(symbols, *, start_time, end_time, + asset_type, freq, on_chunk_done): + # 模拟 provider 实现内部以 2 参调用 on_chunk_done + # (如 GenericHTTPProvider/provider.py:127 / StockSDKProvider/provider.py:166) + if on_chunk_done is not None: + on_chunk_done(1, 3) + return _mock_minute_df() + + mock_provider = MagicMock() + mock_provider.get_minute.side_effect = provider_get_minute_side_effect + _setup_custom_provider(monkeypatch, mock_provider, has_dataset=True) + + df, fallback = kline_sync._try_custom_minute( + ["600519.SH"], + datetime(2026, 1, 15, 9, 25, 0), + datetime(2026, 1, 15, 15, 5, 0), + asset_type="stock", + on_chunk_done=upper_cb, + ) + + assert fallback is False + assert df is not None + # 上层 3 参 spy 被调用一次, 收到 (1, 3, "custom") + upper_cb.assert_called_once_with(1, 3, "custom") + + +# ---------- 测试 9: get_minute_batch 按 asset_type 拆分调用 sync_minute_batch ---------- + +def test_get_minute_batch_splits_stock_and_etf(monkeypatch): + """get_minute_batch 把 incomplete 拆成 stock/ETF 两组, 分别以 + asset_type='stock'/'etf' 调用 sync_minute_batch, 结果 concat 返回。 + + 覆盖 kline.py get_minute_batch 的双调用拼接逻辑 (本次提交改动量最大的部分)。 + 契约: 本端点只接受 stock/ETF (指数走 /api/index/minute), 故两分支覆盖全部 incomplete。 + """ + from app.api import kline as kline_api + + # mock sync_minute_batch: stock 返回 df_s, etf 返回 df_e (不同 symbol 便于 concat 后 filter 验证) + def fake_sync(symbols, *, start_time, end_time, batch_size, rpm, asset_type): + if asset_type == "stock": + return _mock_minute_df(symbol="600519.SH") + if asset_type == "etf": + return _mock_minute_df(symbol="510300.SH") + return pl.DataFrame() + sync_spy = MagicMock(side_effect=fake_sync) + monkeypatch.setattr(kline_api.kline_sync, "sync_minute_batch", sync_spy) + + # mock repo: ETF 集合含 510300.SH; 本地分钟K返回空 (强制走 incomplete 补拉) + mock_repo = MagicMock() + mock_repo.get_etf_symbol_set.return_value = {"510300.SH"} + mock_repo.get_minute_batch.return_value = pl.DataFrame() + + # mock capset: 有权限, limits 返回 None (lim.batch 访问被 `if lim else` 守护) + mock_capset = MagicMock() + mock_capset.has.return_value = True + mock_capset.limits.return_value = None + + mock_request = MagicMock() + mock_request.app.state.repo = mock_repo + mock_request.app.state.capabilities = mock_capset + + body = {"symbols": ["600519.SH", "510300.SH"], "date": "2026-01-15"} + result = kline_api.get_minute_batch(mock_request, body) + + # sync_minute_batch 被调 2 次, asset_type 分别为 stock 和 etf + assert sync_spy.call_count == 2 + call_assets = sorted(call.kwargs.get("asset_type") for call in sync_spy.call_args_list) + assert call_assets == ["etf", "stock"] + + # 两个 symbol 都在结果里 (concat 后按 symbol filter 命中) + assert "600519.SH" in result["data"] + assert "510300.SH" in result["data"]