mirror of
https://ghfast.top/https://github.com/aeroxw/tick-stock-panel.git
synced 2026-09-12 15:34:16 +08:00
fix(minute): 自定义分钟数据源用户单股补拉/监控页整改 (承接 #121 review)
This commit is contained in:
@@ -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
@@ -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(),
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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"]
|
||||
Reference in New Issue
Block a user