From d9a77a534acaae08920dadce069b6c9d1f3b81c9 Mon Sep 17 00:00:00 2001 From: intfoo Date: Sun, 19 Jul 2026 15:53:00 +0800 Subject: [PATCH] =?UTF-8?q?fix(minute):=20=E6=89=BF=E6=8E=A5=20#122=20revi?= =?UTF-8?q?ew=20=E4=B8=89=E5=A4=84=E9=98=BB=E6=96=AD=E9=97=AE=E9=A2=98?= =?UTF-8?q?=E6=95=B4=E6=94=B9=20+=20resolver=20=E5=BC=82=E5=B8=B8=E8=BE=B9?= =?UTF-8?q?=E7=95=8C=E5=8A=A0=E5=9B=BA?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- backend/app/data_providers/custom/config.py | 4 + backend/app/data_providers/custom/provider.py | 21 +- backend/app/services/kline_sync.py | 69 +++- backend/tests/test_minute_routing.py | 331 +++++++++++++++++- 4 files changed, 407 insertions(+), 18 deletions(-) diff --git a/backend/app/data_providers/custom/config.py b/backend/app/data_providers/custom/config.py index 078d650..ae52033 100644 --- a/backend/app/data_providers/custom/config.py +++ b/backend/app/data_providers/custom/config.py @@ -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, ) diff --git a/backend/app/data_providers/custom/provider.py b/backend/app/data_providers/custom/provider.py index 1f80e91..d726c5b 100644 --- a/backend/app/data_providers/custom/provider.py +++ b/backend/app/data_providers/custom/provider.py @@ -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(): diff --git a/backend/app/services/kline_sync.py b/backend/app/services/kline_sync.py index 2489718..f7c9f6a 100644 --- a/backend/app/services/kline_sync.py +++ b/backend/app/services/kline_sync.py @@ -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): diff --git a/backend/tests/test_minute_routing.py b/backend/tests/test_minute_routing.py index f4e6c10..7b6ea3e 100644 --- a/backend/tests/test_minute_routing.py +++ b/backend/tests/test_minute_routing.py @@ -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