fix(minute): 自定义分钟数据源用户单股补拉/监控页整改 (承接 #121 review)

This commit is contained in:
intfoo
2026-07-19 13:35:35 +08:00
parent e2ffde99d5
commit 4c85f99633
8 changed files with 451 additions and 45 deletions
+1 -1
View File
@@ -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"),
+34 -10
View File
@@ -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(),
+8 -1
View File
@@ -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,
@@ -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] = []
@@ -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()
+6 -4
View File
@@ -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)
+82 -26
View File
@@ -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:
+312
View File
@@ -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"]