fix(minute): 承接 #122 review 三处阻断问题整改 + resolver 异常边界加固

This commit is contained in:
intfoo
2026-07-19 16:00:22 +08:00
parent 91e16db547
commit d9a77a534a
4 changed files with 407 additions and 18 deletions
@@ -33,6 +33,8 @@ class DatasetConfig:
symbols_param: str = "symbols"
start_param: str = "start_time"
end_param: str = "end_time"
asset_type_param: str | None = None
freq_param: str | None = None
@dataclass(frozen=True)
@@ -72,6 +74,8 @@ def _dataset_from_dict(raw: dict[str, Any]) -> DatasetConfig:
symbols_param=str(raw.get("symbols_param", "symbols") or "symbols"),
start_param=str(raw.get("start_param", "start_time") or "start_time"),
end_param=str(raw.get("end_param", "end_time") or "end_time"),
asset_type_param=str(raw.get("asset_type_param")) if raw.get("asset_type_param") else None,
freq_param=str(raw.get("freq_param")) if raw.get("freq_param") else None,
)
+18 -3
View File
@@ -111,16 +111,31 @@ class GenericHTTPProvider:
symbols: list[str],
start_time: datetime | None,
end_time: datetime | None,
asset_type: AssetType = "stock", # noqa: ARG002
freq: str = "1m", # noqa: ARG002
asset_type: AssetType = "stock",
freq: str = "1m",
on_chunk_done: Callable[[int, int], None] | None = None,
) -> pl.DataFrame:
"""拉取分钟 K。
asset_type / freq 默认不传上游 (minute dataset URL 应返回 1m 数据)。
在 dataset 配置中设置 asset_type_param / freq_param 后, 这两个参数会以
配置的参数名注入请求 (GET → params, POST → body), 用于上游需区分
stock/ETF/index 或固定频率的场景。
"""
cfg = self._dataset("minute")
override: dict[str, Any] = {}
if cfg.asset_type_param:
override[cfg.asset_type_param] = asset_type
if cfg.freq_param:
override[cfg.freq_param] = freq
frames: list[pl.DataFrame] = []
chunks = chunked(symbols, cfg.batch)
for i, chunk in enumerate(chunks):
sleep_between_batches(i, cfg.rpm)
rows = self._request_rows(cfg, symbols=chunk, start_time=start_time, end_time=end_time)
rows = self._request_rows(
cfg, symbols=chunk, start_time=start_time, end_time=end_time,
override_params=override or None, override_body=override or None,
)
df = self._mapped_frame(cfg, rows)
df = self._normalize_minute(df)
if not df.is_empty():
+55 -14
View File
@@ -529,6 +529,34 @@ def _write_minute_partition(df: pl.DataFrame, minute_dir) -> int:
return written
def _resolve_minute_provider(
provider_name: str,
) -> tuple[object | None, bool, str | None]:
"""统一解析 custom minute provider, 把所有 resolver 调用纳入同一异常边界。
供 _try_custom_minute 和 sync_and_persist_minute 共用, 避免两处分别调
provider_has_dataset / get_provider 时漏掉异常边界 (Issue 2 加固项)。
返回 (provider, should_fallback_to_tickflow, error_msg):
- provider_name == "tickflow" 或未配 minute dataset → (None, True, None) 静默降级
- resolver 异常 (registry 损坏 / 插件失效 / provider name 不存在) → (None, True, str(e))
- 成功 → (provider, False, None)
上层依据 error_msg 决定是否 logger.warning (区分"未配""异常")。
注意: provider.get_minute() 仍由调用方在自身 try 块内调用 (业务异常, 非解析异常)。
"""
if provider_name == "tickflow":
return (None, True, None)
from app.data_providers import custom as custom_sources
try:
if not custom_sources.provider_has_dataset(provider_name, "minute"):
return (None, True, None)
provider = custom_sources.get_provider(provider_name)
return (provider, False, None)
except Exception as e: # noqa: BLE001
return (None, True, str(e))
def _try_custom_minute(
symbols: list[str],
start_time: datetime | None,
@@ -548,17 +576,21 @@ def _try_custom_minute(
None 档用户 TickFlow 失败返回空。不显式判断 tier, 避免 #126 augmented
capability 逻辑干扰。
resolver 异常边界由 _resolve_minute_provider 统一兜底; 业务调用
(provider.get_minute) 仍在本函数 try 块内, 与 resolver 异常分离
便于日志区分 ("resolution failed" vs "call failed")。
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":
provider, fallback, err = _resolve_minute_provider(provider_name)
if fallback:
if err is not None:
logger.warning("custom minute provider %s resolution failed, falling back to TickFlow: %s",
provider_name, err)
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
@@ -574,7 +606,7 @@ def _try_custom_minute(
)
return (df, False)
except Exception as e: # noqa: BLE001
logger.warning("custom minute provider %s failed, falling back to TickFlow: %s",
logger.warning("custom minute provider %s call failed, falling back to TickFlow: %s",
provider_name, e)
return (None, True)
@@ -612,10 +644,15 @@ def sync_minute_batch(
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()
# 自定义源成功: 遵守与 TickFlow 路径一致的 on_segment 契约。
# 传了 on_segment (如 sync_and_persist_minute 流式落盘) → 调 on_segment, 返回空 df;
# 未传 on_segment (如 fetch_minute_single 实时补拉) 或空 df → 原样返回 df。
df = df if df is not None else pl.DataFrame()
if on_segment and not df.is_empty():
# 空 df 不调 on_segment, 与 TickFlow 路径 `if seg_out:` (L684) 对称
on_segment(df)
return pl.DataFrame()
return df
tf = get_client()
@@ -967,10 +1004,14 @@ def sync_and_persist_minute(
on_chunk_done(current, total) 每个 chunk 完成后回调。
"""
minute_provider = preferences.get_minute_data_provider()
minute_is_custom = False
if minute_provider != "tickflow":
from app.data_providers import custom as custom_sources
minute_is_custom = custom_sources.provider_has_dataset(minute_provider, "minute")
# resolver 调用统一走 _resolve_minute_provider, 与 _try_custom_minute 共用异常边界。
# resolver 异常时视为非 custom (minute_is_custom=False), 走 capset 检查 →
# sync_minute_batch 内 _try_custom_minute 会再次 resolver 异常 → fallback TickFlow。
_, fallback, resolve_err = _resolve_minute_provider(minute_provider)
minute_is_custom = not fallback
if resolve_err is not None:
logger.warning("custom minute provider %s resolution failed at sync_and_persist_minute, treating as non-custom: %s",
minute_provider, resolve_err)
if not symbols:
return 0
if not minute_is_custom and not capset.has(Cap.KLINE_MINUTE_BATCH):
+330 -1
View File
@@ -207,7 +207,10 @@ def test_custom_success_skips_tickflow(monkeypatch):
# ---------- 测试 7: sync_minute_batch 自定义源成功直接返回 ----------
def test_sync_minute_batch_custom_success_returns_directly(monkeypatch):
"""§4 测试 7: sync_minute_batch 自定义源成功 → 直接 return, 不走 segment 逻辑。"""
"""§4 测试 7: sync_minute_batch 自定义源成功 + 未传 on_segment → 原样返回 df (实时补拉契约)。
传了 on_segment 时走流式落盘分支 (见测试 10), 此处验证未传时的实时补拉契约。
"""
expected_df = _mock_minute_df()
mock_provider = MagicMock()
mock_provider.get_minute.return_value = expected_df
@@ -310,3 +313,329 @@ def test_get_minute_batch_splits_stock_and_etf(monkeypatch):
# 两个 symbol 都在结果里 (concat 后按 symbol filter 命中)
assert "600519.SH" in result["data"]
assert "510300.SH" in result["data"]
# ---------- 测试 10: sync_minute_batch 自定义源成功时调 on_segment (Issue 1) ----------
def test_sync_minute_batch_custom_calls_on_segment(monkeypatch):
"""Issue 1: sync_minute_batch 自定义源成功 + 传了 on_segment →
调 on_segment(df), 返回空 df (数据已落盘)。
"""
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)
on_segment_spy = MagicMock(name="on_segment_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),
on_segment=on_segment_spy,
asset_type="stock",
)
on_segment_spy.assert_called_once_with(expected_df)
assert isinstance(df, pl.DataFrame)
assert df.is_empty()
get_client_spy.assert_not_called()
# ---------- 测试 11: 自定义源返回空 df 时不调 on_segment (Issue 1 边界) ----------
def test_sync_minute_batch_custom_empty_df_skips_on_segment(monkeypatch):
"""Issue 1 边界: 自定义源返回空 df → 不调 on_segment (与 TickFlow `if seg_out:` 对称)。
"""
mock_provider = MagicMock()
mock_provider.get_minute.return_value = pl.DataFrame()
_setup_custom_provider(monkeypatch, mock_provider, has_dataset=True)
on_segment_spy = MagicMock(name="on_segment_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),
on_segment=on_segment_spy,
asset_type="stock",
)
on_segment_spy.assert_not_called()
assert isinstance(df, pl.DataFrame)
assert df.is_empty()
# ---------- 测试 12: sync_and_persist_minute + custom provider 端到端落盘 (Issue 1) ----------
def test_sync_and_persist_minute_custom_persists(monkeypatch, tmp_path):
"""Issue 1 端到端: sync_and_persist_minute + 自定义源 →
_write_minute_partition 被调, written > 0。
"""
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)
# mock sync_and_persist_minute 内部依赖 (通过 monkeypatch kline_sync 模块属性)
monkeypatch.setattr(kline_sync, "_cleanup_null_datetime_minute", lambda repo: None)
monkeypatch.setattr(kline_sync, "_migrate_symbol_to_date_partition", lambda repo: None)
monkeypatch.setattr(kline_sync, "_latest_minute_datetime", lambda repo: None)
monkeypatch.setattr(kline_sync, "resolve_limit", lambda *a, **kw: MagicMock(batch=100, rpm=30))
monkeypatch.setattr(kline_sync.preferences, "get_minute_sync_segment_days", lambda: 20)
# _write_minute_partition spy: 记录调用, 返回行数
write_spy = MagicMock(return_value=expected_df.height)
monkeypatch.setattr(kline_sync, "_write_minute_partition", write_spy)
# get_client spy: 自定义源成功时不应走 TickFlow
get_client_spy = MagicMock(name="get_client_spy")
monkeypatch.setattr(kline_sync, "get_client", get_client_spy)
# mock repo
mock_repo = MagicMock()
mock_repo.store.data_dir = tmp_path
mock_repo.db.execute = MagicMock()
# mock capset (minute_is_custom=True 绕过 has() 检查, resolve_limit 已 mock)
mock_capset = MagicMock()
written = kline_sync.sync_and_persist_minute(
["600519.SH"], mock_repo, mock_capset,
)
assert write_spy.called
assert written == expected_df.height
assert written > 0
get_client_spy.assert_not_called()
# ---------- 测试 13: get_provider 异常时 fall through TickFlow (Issue 2) ----------
def test_get_provider_exception_falls_back_to_tickflow(monkeypatch):
"""Issue 2: get_provider raise ValueError →
_try_custom_minute 返回 (None, True), 无异常穿透。
"""
monkeypatch.setattr(
kline_sync.preferences,
"get_minute_data_provider",
lambda: "mock_src",
)
monkeypatch.setattr(
"app.data_providers.custom.provider_has_dataset",
lambda name, ds: True, # provider 存在, 但 get_provider 会抛
)
def _raising_get_provider(name):
raise ValueError("not found")
monkeypatch.setattr(
"app.data_providers.custom.get_provider",
_raising_get_provider,
)
df, fallback = kline_sync._try_custom_minute(
["600519.SH"], None, None, asset_type="stock",
)
assert fallback is True
assert df is None
# ---------- 测试 14: provider_has_dataset 异常时 fall through (Issue 2) ----------
def test_provider_has_dataset_exception_falls_back(monkeypatch):
"""Issue 2: provider_has_dataset raise →
_try_custom_minute 返回 (None, True), 无异常穿透。
"""
monkeypatch.setattr(
kline_sync.preferences,
"get_minute_data_provider",
lambda: "mock_src",
)
def _raising_has_dataset(name, ds):
raise RuntimeError("registry corrupted")
monkeypatch.setattr(
"app.data_providers.custom.provider_has_dataset",
_raising_has_dataset,
)
df, fallback = kline_sync._try_custom_minute(
["600519.SH"], None, None, asset_type="stock",
)
assert fallback is True
assert df is None
# ---------- 测试 15-17: GenericHTTPProvider opt-in 参数传递 (Issue 3) ----------
from app.data_providers.custom.config import CustomSourceConfig, DatasetConfig
from app.data_providers.custom.provider import GenericHTTPProvider
def _make_minute_config(**extra) -> CustomSourceConfig:
"""构造带 minute dataset 的最小 CustomSourceConfig, extra 传给 DatasetConfig。"""
field_map = {f: f for f in (
"symbol", "datetime", "open", "high", "low", "close", "volume", "amount"
)}
return CustomSourceConfig(
name="test_src",
display_name="Test Source",
datasets={"minute": DatasetConfig(
url="http://example.com/minute", field_map=field_map, **extra,
)},
)
def _capture_request_rows(provider):
"""替换 _request_rows 为捕获 spy, 返回 captured dict。"""
captured: dict = {}
def fake_request_rows(cfg, *, symbols=None, start_time=None, end_time=None,
override_params=None, override_body=None):
captured["override_params"] = override_params
captured["override_body"] = override_body
return [] # 空行 → 空 df
provider._request_rows = fake_request_rows
return captured
def test_generic_http_get_minute_passes_asset_type_when_configured():
"""Issue 3: 配了 asset_type_param="asset" → override 含 {"asset": "etf"}。"""
config = _make_minute_config(asset_type_param="asset")
provider = GenericHTTPProvider(config)
captured = _capture_request_rows(provider)
provider.get_minute(["600519.SH"], None, None, asset_type="etf", freq="1m")
assert captured["override_params"] == {"asset": "etf"}
assert captured["override_body"] == {"asset": "etf"}
def test_generic_http_get_minute_passes_freq_when_configured():
"""Issue 3: 配了 freq_param="period" → override 含 {"period": "1m"}。"""
config = _make_minute_config(freq_param="period")
provider = GenericHTTPProvider(config)
captured = _capture_request_rows(provider)
provider.get_minute(["600519.SH"], None, None, asset_type="stock", freq="1m")
assert captured["override_params"] == {"period": "1m"}
assert captured["override_body"] == {"period": "1m"}
def test_generic_http_get_minute_omits_params_when_not_configured():
"""Issue 3 向后兼容: 未配 asset_type_param/freq_param → override 为 None, 不传上游。"""
config = _make_minute_config() # 无 asset_type_param / freq_param
provider = GenericHTTPProvider(config)
captured = _capture_request_rows(provider)
provider.get_minute(["600519.SH"], None, None, asset_type="etf", freq="1m")
# override 为 None (空 dict → `override or None`), 不传上游
assert captured["override_params"] is None
assert captured["override_body"] is None
# ---------- 测试 18: sync_and_persist_minute resolver 异常时优雅返回 0 (观察项加固) ----------
def test_sync_and_persist_minute_resolver_exception_returns_zero(monkeypatch, tmp_path):
"""观察项加固: sync_and_persist_minute 开头 _resolve_minute_provider 异常 →
不向接口抛 500, 优雅降级 (minute_is_custom=False → 走 capset 检查 → 无权限 return 0)。
"""
monkeypatch.setattr(
kline_sync.preferences,
"get_minute_data_provider",
lambda: "mock_src",
)
# provider_has_dataset 抛异常 (模拟 registry 损坏)
def _raising_has_dataset(name, ds):
raise RuntimeError("registry corrupted")
monkeypatch.setattr(
"app.data_providers.custom.provider_has_dataset",
_raising_has_dataset,
)
# 无 KLINE_MINUTE_BATCH 权限 → resolver 异常视为非 custom → capset 检查失败 → return 0
mock_capset = MagicMock()
mock_capset.has.return_value = False
mock_repo = MagicMock()
mock_repo.store.data_dir = tmp_path
# 不应抛异常, 优雅降级到 0
written = kline_sync.sync_and_persist_minute(
["600519.SH"], mock_repo, mock_capset,
)
assert written == 0
# ---------- 测试 19: _resolve_minute_provider helper 单元测试 ----------
def test_resolve_minute_provider_tickflow_returns_silent_fallback():
"""观察项加固: provider_name == "tickflow" → (None, True, None) 静默降级, 无 err。"""
provider, fallback, err = kline_sync._resolve_minute_provider("tickflow")
assert provider is None
assert fallback is True
assert err is None
def test_resolve_minute_provider_no_dataset_returns_silent_fallback(monkeypatch):
"""观察项加固: 配了 custom 但未配 minute dataset → (None, True, None) 静默降级。"""
monkeypatch.setattr(
"app.data_providers.custom.provider_has_dataset",
lambda name, ds: False, # 已注册但未配 minute
)
provider, fallback, err = kline_sync._resolve_minute_provider("mock_src")
assert provider is None
assert fallback is True
assert err is None # 未配 ≠ 异常, 不应触发 warning
def test_resolve_minute_provider_has_dataset_exception_returns_err(monkeypatch):
"""观察项加固: provider_has_dataset 抛异常 → (None, True, str(e)), 上层据此 warning。"""
def _raising(name, ds):
raise RuntimeError("registry corrupted")
monkeypatch.setattr("app.data_providers.custom.provider_has_dataset", _raising)
provider, fallback, err = kline_sync._resolve_minute_provider("mock_src")
assert provider is None
assert fallback is True
assert err is not None
assert "registry corrupted" in err
def test_resolve_minute_provider_get_provider_exception_returns_err(monkeypatch):
"""观察项加固: provider_has_dataset 返回 True 但 get_provider 抛 → (None, True, str(e))。"""
monkeypatch.setattr(
"app.data_providers.custom.provider_has_dataset",
lambda name, ds: True,
)
def _raising_get(name):
raise ValueError("not found")
monkeypatch.setattr("app.data_providers.custom.get_provider", _raising_get)
provider, fallback, err = kline_sync._resolve_minute_provider("mock_src")
assert provider is None
assert fallback is True
assert err is not None
assert "not found" in err
def test_resolve_minute_provider_success_returns_provider(monkeypatch):
"""观察项加固: 正常路径 → (provider, False, None)。"""
mock_provider = object() # 任意 truthy 对象即可
monkeypatch.setattr(
"app.data_providers.custom.provider_has_dataset",
lambda name, ds: True,
)
monkeypatch.setattr(
"app.data_providers.custom.get_provider",
lambda name: mock_provider,
)
provider, fallback, err = kline_sync._resolve_minute_provider("mock_src")
assert provider is mock_provider
assert fallback is False
assert err is None