diff --git a/backend/app/data_providers/custom/provider.py b/backend/app/data_providers/custom/provider.py index 2e62148..d7a277a 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 +import time from collections.abc import Callable from datetime import datetime, timedelta from pathlib import Path @@ -132,6 +133,23 @@ class GenericHTTPProvider: ) return errors + def _request_rows_retry( + self, cfg, symbols: list[str], *, start_time=None, end_time=None, retries: int = 1 + ) -> list[dict]: + """单批请求 + 短退避重试。仍失败抛出, 由调用方决定隔离粒度 (#226)。""" + last: Exception | None = None + for attempt in range(retries + 1): + try: + return self._request_rows( + cfg, symbols=symbols, start_time=start_time, end_time=end_time + ) + except Exception as e: # noqa: BLE001 + last = e + if attempt < retries: + time.sleep(1.0 * (attempt + 1)) + assert last is not None + raise last + def get_daily( self, symbols: list[str], @@ -143,15 +161,35 @@ class GenericHTTPProvider: cfg = self._dataset("daily") frames: list[pl.DataFrame] = [] chunks = chunked(symbols, cfg.batch) + failed: list[str] = [] 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) + try: + rows = self._request_rows_retry( + cfg, chunk, start_time=start_time, end_time=end_time + ) + except Exception as e: # noqa: BLE001 + # 单批失败只隔离该批 (#226): 之前任一批 502 会让整个 stage + # 抛异常, 已成功批次的结果留在内存里全部丢弃 + failed.extend(chunk) + logger.warning( + "custom daily: batch %d/%d failed (%d symbols), skipped: %s", + i + 1, len(chunks), len(chunk), e, + ) + if on_chunk_done: + on_chunk_done(i + 1, len(chunks)) + continue df = self._mapped_frame(cfg, rows) df = normalize_daily(df, source=self.name) if not df.is_empty(): frames.append(df) if on_chunk_done: on_chunk_done(i + 1, len(chunks)) + if failed: + logger.warning( + "custom daily: %d/%d symbols missing due to batch failures: %s", + len(failed), len(symbols), ", ".join(failed[:20]), + ) return pl.concat(frames, how="diagonal_relaxed") if frames else pl.DataFrame() def get_adj_factors( @@ -165,15 +203,33 @@ class GenericHTTPProvider: cfg = self._dataset("adj_factor") frames: list[pl.DataFrame] = [] chunks = chunked(symbols, cfg.batch) + failed: list[str] = [] 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) + try: + rows = self._request_rows_retry( + cfg, chunk, start_time=start_time, end_time=end_time + ) + except Exception as e: # noqa: BLE001 + failed.extend(chunk) + logger.warning( + "custom adj_factor: batch %d/%d failed (%d symbols), skipped: %s", + i + 1, len(chunks), len(chunk), e, + ) + if on_chunk_done: + on_chunk_done(i + 1, len(chunks)) + continue df = self._mapped_frame(cfg, rows) df = normalize_adj_factors(df, source=self.name) if not df.is_empty(): frames.append(df) if on_chunk_done: on_chunk_done(i + 1, len(chunks)) + if failed: + logger.warning( + "custom adj_factor: %d/%d symbols missing due to batch failures: %s", + len(failed), len(symbols), ", ".join(failed[:20]), + ) return pl.concat(frames, how="diagonal_relaxed") if frames else pl.DataFrame() def get_realtime(self) -> list[dict]: @@ -293,12 +349,18 @@ class GenericHTTPProvider: return pl.DataFrame() return pl.concat(frames, how="diagonal_relaxed") - @staticmethod - def _normalize_minute(df: pl.DataFrame) -> pl.DataFrame: + @classmethod + def _normalize_minute(cls, df: pl.DataFrame) -> pl.DataFrame: """把映射后的 df 规范成 minute canonical 列。""" if df.is_empty(): return df if "datetime" in df.columns and df.schema["datetime"] != pl.Datetime("us"): + if df.schema["datetime"] == pl.Utf8: + # 字符串 datetime 直接 cast 会整体置 null (polars 不做字符串解析); + # 先解析再对齐微秒精度 (#225, 参照 + # kline_sync._enforce_minute_beijing_wallclock 的处理)。 + # Series 级立即解析: 表达式错误要到 collect 才抛, 无法按格式回退 + df = df.with_columns(cls._parse_datetime_series(df["datetime"])) df = df.with_columns(pl.col("datetime").cast(pl.Datetime("us"), strict=False)) for col in ("open", "high", "low", "close", "volume", "amount"): if col in df.columns: @@ -306,6 +368,27 @@ class GenericHTTPProvider: keep = [c for c in ("symbol", "datetime", "open", "high", "low", "close", "volume", "amount") if c in df.columns] return df.select(keep) if keep else pl.DataFrame() + _DATETIME_STR_FORMATS = ( + None, # 自动推断 + "%Y-%m-%d %H:%M:%S", + "%Y-%m-%dT%H:%M:%S", + "%Y/%m/%d %H:%M:%S", + "%Y-%m-%d %H:%M", + ) + + @classmethod + def _parse_datetime_series(cls, s: pl.Series) -> pl.Series: + """逐格式尝试解析字符串 datetime; 均失败返回全 null (宽松语义)。""" + for fmt in cls._DATETIME_STR_FORMATS: + try: + return ( + s.str.to_datetime(strict=False, format=fmt) + if fmt else s.str.to_datetime(strict=False) + ) + except Exception: # noqa: BLE001 — 该格式不适用, 换下一个 + continue + return pl.Series("datetime", [None] * s.len(), dtype=pl.Datetime("us")) + def test_dataset(self, dataset: str, symbols: list[str] | None = None) -> dict: cfg = self._dataset(dataset) test_symbols = symbols or ["000001.SZ"] diff --git a/backend/app/services/backtest.py b/backend/app/services/backtest.py index 4f4e942..769b17a 100644 --- a/backend/app/services/backtest.py +++ b/backend/app/services/backtest.py @@ -7,7 +7,7 @@ from __future__ import annotations import logging import uuid from dataclasses import dataclass, field -from datetime import date +from datetime import date, timedelta from typing import Literal import numpy as np @@ -20,6 +20,10 @@ from app.tickflow.repository import KlineRepository logger = logging.getLogger(__name__) +# 旧信号回测的指标 warmup 日历窗口 (#201): 与 backtest.factor.FACTOR_WARMUP_DAYS +# 同源 (120 交易日 → 保守取日历日), 覆盖 MA60/MACD/BOLL 等最长回看 +_WARMUP_CALENDAR_DAYS = 120 * 1.6 + # vectorbt 是 optional extras(见 pyproject.toml).未装时只有 backtest 不可用,其他功能正常. _vbt = None _vbt_unavailable_reason: str | None = None @@ -162,11 +166,17 @@ class BacktestService: try: from app.tickflow.repository import enriched_dirname enriched_glob = str(self.repo.store.data_dir / enriched_dirname(asset_type) / "**" / "*.parquet") + # 指标 warmup (#201): MA/MACD/RSI/BOLL 需要区间前的历史窗口, + # 直接按 [start,end] 过滤后 compute_all 会让区间头部的指标失真。 + # 与挖掘侧同款公式 (mining_runtime: warmup = max(120, bars*1.6)), + # 此处指标最长回看约 120 交易日, 取保守日历日窗口; 数据不足时 + # 自然退化 (有多少算多少)。计算完成后裁回 [start,end]。 + warmup_start = start - timedelta(days=_WARMUP_CALENDAR_DAYS) df = ( scan_enriched_parquet(enriched_glob) .filter( (pl.col("symbol").is_in(symbols)) - & (pl.col("date") >= start) + & (pl.col("date") >= warmup_start) & (pl.col("date") <= end) ) .sort(["date", "symbol"]) @@ -182,6 +192,7 @@ class BacktestService: # 即时计算指标 + 信号 from app.indicators.pipeline import compute_all df = compute_all(df) + df = df.filter(pl.col("date") >= start) # 选择需要的列 needed_cols = [ diff --git a/backend/tests/test_backtest_warmup.py b/backend/tests/test_backtest_warmup.py new file mode 100644 index 0000000..3f00134 --- /dev/null +++ b/backend/tests/test_backtest_warmup.py @@ -0,0 +1,66 @@ +"""#201 回归: 旧信号回测的 _load_panel 必须带指标 warmup 窗口。 + +直接按 [start,end] 过滤后 compute_all, 区间头部的 MA/MACD/RSI 会因缺 +历史窗口而失真 (回测起始段信号不可信)。 +""" +from __future__ import annotations + +from datetime import date, timedelta +from unittest.mock import MagicMock + +import polars as pl + +from app.services.backtest import BacktestService + + +def _synthetic_enriched(n_days: int) -> pl.DataFrame: + base = date(2026, 1, 1) + days = [base + timedelta(days=i) for i in range(n_days)] + n = len(days) + closes = [10.0 + (i % 7) * 0.3 + i * 0.01 for i in range(n)] + return pl.DataFrame( + { + "symbol": ["600000.SH"] * n, + "date": days, + "open": [c - 0.05 for c in closes], + "high": [c + 0.1 for c in closes], + "low": [c - 0.1 for c in closes], + "close": closes, + "volume": [10000.0] * n, + "amount": [c * 10000.0 for c in closes], + "raw_close": closes, + "raw_high": [c + 0.1 for c in closes], + "raw_low": [c - 0.1 for c in closes], + } + ) + + +def test_load_panel_warms_up_indicators(monkeypatch) -> None: + df = _synthetic_enriched(250) + monkeypatch.setattr( + "app.services.backtest.scan_enriched_parquet", lambda glob: df.lazy() + ) + svc = BacktestService(repo=MagicMock()) + + start = df["date"][-30] + end = df["date"][-1] + panel = svc._load_panel(["600000.SH"], start, end) + + # warmup 行不进入结果面板 (pandas datetime64 与 date 直接比较会类型不符) + assert str(panel["date"].min())[:10] == start.isoformat() + assert str(panel["date"].max())[:10] == end.isoformat() + # 区间首日的指标已有历史窗口可用, 不再是 NaN + first = panel.iloc[0] + assert first["rsi_14"] == first["rsi_14"] # NaN != NaN + + +def test_load_panel_insufficient_history_degrades_gracefully(monkeypatch) -> None: + # 数据起点晚于 warmup 起点时自然退化: 有多少算多少, 不抛异常 + df = _synthetic_enriched(20) + monkeypatch.setattr( + "app.services.backtest.scan_enriched_parquet", lambda glob: df.lazy() + ) + svc = BacktestService(repo=MagicMock()) + + panel = svc._load_panel(["600000.SH"], df["date"][0], df["date"][-1]) + assert len(panel) == 20 diff --git a/backend/tests/test_custom_batch_isolation.py b/backend/tests/test_custom_batch_isolation.py new file mode 100644 index 0000000..8bf3461 --- /dev/null +++ b/backend/tests/test_custom_batch_isolation.py @@ -0,0 +1,140 @@ +"""#225/#226 回归: 自定义源分钟K字符串日期解析 + 日K分批失败隔离。""" +from __future__ import annotations + +from datetime import date, datetime + +import polars as pl + +from app.data_providers.custom.config import CustomSourceConfig, DatasetConfig +from app.data_providers.custom.loader import GenericHTTPProvider + + +def _provider(datasets: dict[str, DatasetConfig]) -> GenericHTTPProvider: + return GenericHTTPProvider(CustomSourceConfig( + name="test_source", + display_name="Test Source", + datasets=datasets, + )) + + +def _daily_config(batch: int = 2) -> DatasetConfig: + return DatasetConfig( + url="https://example.test/daily", + field_map={ + "symbol": "symbol", "date": "date", "open": "open", "high": "high", + "low": "low", "close": "close", "volume": "volume", "amount": "amount", + }, + batch=batch, + ) + + +# ── #225: 分钟K字符串 datetime 不得被 cast 成 null ──────────────── + +def test_normalize_minute_parses_string_datetime() -> None: + df = pl.DataFrame( + { + "symbol": ["600000.SH"] * 2, + # 上游 YAML 映射后仍是字符串; 旧代码直接 cast → 全 null (#225) + "datetime": ["2026-09-01 09:35:00", "2026-09-01 09:40:00"], + "close": [10.0, 10.5], + } + ) + out = GenericHTTPProvider._normalize_minute(df) + assert out.schema["datetime"] == pl.Datetime("us") + assert out["datetime"].null_count() == 0 + assert out["datetime"][0] == datetime(2026, 9, 1, 9, 35) + + # 非法字符串维持 strict=False 宽松行为 (null, 不抛异常) + bad = pl.DataFrame( + {"symbol": ["600000.SH"], "datetime": ["not-a-date"], "close": [1.0]} + ) + out_bad = GenericHTTPProvider._normalize_minute(bad) + assert out_bad["datetime"].null_count() == 1 + + +def test_normalize_minute_datetime_already_typed_unchanged() -> None: + df = pl.DataFrame( + { + "symbol": ["600000.SH"], + "datetime": [datetime(2026, 9, 1, 9, 35)], + "close": [10.0], + } + ) + out = GenericHTTPProvider._normalize_minute(df) + assert out["datetime"][0] == datetime(2026, 9, 1, 9, 35) + + +# ── #226: get_daily 单批失败只隔离该批 ──────────────────────── + +def _canonical_frame(symbols: list[str], day: str) -> pl.DataFrame: + return pl.DataFrame( + { + "symbol": symbols, + "date": [date.fromisoformat(day)] * len(symbols), + "open": [10.0] * len(symbols), + "high": [11.0] * len(symbols), + "low": [9.0] * len(symbols), + "close": [10.5] * len(symbols), + "volume": [100.0] * len(symbols), + "amount": [1050.0] * len(symbols), + } + ) + + +def test_get_daily_isolates_failed_batch() -> None: + provider = _provider({"daily": _daily_config(batch=2)}) + calls: list[list[str]] = [] + + def request_rows(cfg, symbols=None, **kwargs): + calls.append(list(symbols)) + if symbols == ["s3", "s4"]: + raise RuntimeError("502 Bad Gateway") + return [{"_rows": list(symbols)}] + + provider._request_rows = request_rows + provider._mapped_frame = lambda cfg, rows: _canonical_frame( + rows[0]["_rows"], "2026-09-01" + ) + + df = provider.get_daily( + ["s1", "s2", "s3", "s4", "s5", "s6"], + datetime(2026, 8, 1), datetime(2026, 9, 1), + ) + + # 3 批都请求过 (失败批重试 1 次后跳过、流程继续), 返回第 1、3 批共 4 行 + assert calls == [["s1", "s2"], ["s3", "s4"], ["s3", "s4"], ["s5", "s6"]] + assert df.height == 4 + assert set(df["symbol"]) == {"s1", "s2", "s5", "s6"} + + +def test_get_daily_progress_callback_fires_for_failed_batch() -> None: + provider = _provider({"daily": _daily_config(batch=2)}) + + def request_rows(cfg, symbols=None, **kwargs): + if symbols == ["s3", "s4"]: + raise RuntimeError("timeout") + return [{"_rows": list(symbols)}] + + progress: list[tuple[int, int]] = [] + provider._request_rows = request_rows + provider._mapped_frame = lambda cfg, rows: _canonical_frame( + rows[0]["_rows"], "2026-09-01" + ) + + provider.get_daily( + ["s1", "s2", "s3", "s4"], datetime(2026, 8, 1), datetime(2026, 9, 1), + on_chunk_done=lambda cur, tot: progress.append((cur, tot)), + ) + # 失败批也推进进度, 前端进度条不会卡死 + assert progress == [(1, 2), (2, 2)] + + +def test_get_daily_all_batches_fail_returns_empty() -> None: + provider = _provider({"daily": _daily_config(batch=2)}) + + def request_rows(cfg, symbols=None, **kwargs): + raise RuntimeError("down") + + provider._request_rows = request_rows + df = provider.get_daily(["s1", "s2"], datetime(2026, 8, 1), datetime(2026, 9, 1)) + assert df.is_empty() diff --git a/frontend/src/pages/Watchlist.tsx b/frontend/src/pages/Watchlist.tsx index 32d21f4..7ec36f4 100644 --- a/frontend/src/pages/Watchlist.tsx +++ b/frontend/src/pages/Watchlist.tsx @@ -57,6 +57,10 @@ import { const BOARDS = ['沪主板', '深主板', '创业板', '科创板', '北交所'] as const type BoardType = typeof BOARDS[number] +// 板块筛选选项 = 股票板块 + ETF(ETF 无 symbol 板块语义,按 asset_type 匹配) +const ETF_BOARD = 'ETF' +const BOARD_OPTIONS = [...BOARDS, ETF_BOARD] + function getBoardType(symbol: string): BoardType | null { if (/^(300|301)/.test(symbol)) return '创业板' if (/^688/.test(symbol)) return '科创板' @@ -1112,9 +1116,10 @@ export function Watchlist() { const [filters, setFilters] = useState>({}) // 板块筛选(持久化) + // 兼容: 旧存储不含 ETF 键 → 加载时补上,保持 ETF 行默认可见 const [boardFilter, setBoardFilter] = useState>(() => { const saved = storage.watchlistBoardFilter.get([]) - return saved.length > 0 ? new Set(saved) : new Set(BOARDS) // 默认全选 + return saved.length > 0 ? new Set([...saved, ETF_BOARD]) : new Set(BOARD_OPTIONS) // 默认全选 }) const persistBoardFilter = useCallback((next: Set) => { setBoardFilter(next) @@ -1156,7 +1161,7 @@ export function Watchlist() { const resetAllFilters = useCallback(() => { setFilters({}) - persistBoardFilter(new Set(BOARDS)) + persistBoardFilter(new Set(BOARD_OPTIONS)) setExcludeST(false) storage.watchlistExcludeST.set(false) }, [persistBoardFilter]) @@ -1184,9 +1189,10 @@ export function Watchlist() { const filteredRows = useMemo(() => { // 板块筛选(全选时跳过) let result = rowsInSelectedGroup - if (boardFilter.size > 0 && boardFilter.size < BOARDS.length) { + if (boardFilter.size > 0 && boardFilter.size < BOARD_OPTIONS.length) { result = result.filter(r => { - // 非股票 (指数/ETF) 无板块语义, 不受板块筛选影响 (顺带修复 ETF 行被误过滤) + if (r.asset_type === 'etf') return boardFilter.has(ETF_BOARD) + // 其他非股票 (指数等) 无板块语义, 不受板块筛选影响 if (r.asset_type && r.asset_type !== 'stock') return true const board = getBoardType(r.symbol) return board != null && boardFilter.has(board) @@ -1219,7 +1225,7 @@ export function Watchlist() { }, [rowsInSelectedGroup, filters, columns, boardFilter, excludeST]) const activeFilterCount = Object.values(filters).filter(v => v.min || v.max || v.text).length - const hasBoardFilter = boardFilter.size > 0 && boardFilter.size < BOARDS.length + const hasBoardFilter = boardFilter.size > 0 && boardFilter.size < BOARD_OPTIONS.length const hasActiveFilters = activeFilterCount > 0 || hasBoardFilter || excludeST // 排序(复用共享三态排序 hook)。分时列按「最新分钟收盘 vs 昨收」排序(分时图最后一点同口径), @@ -1527,7 +1533,7 @@ export function Watchlist() {
板块
- {BOARDS.map(board => { + {BOARD_OPTIONS.map(board => { const active = boardFilter.has(board) return (
- - setFees(e.target.value)} + + +
+
+ + setFees(e.target.value)} + className={INPUT_CLS} /> +
+
+ + setSlippage(e.target.value)} className={INPUT_CLS} />
@@ -322,7 +338,7 @@ export function FactorBacktest({ initialFactorName = 'momentum_20d' }: { initial {/* 结果面板 */}
- {result?.error && !result.ic_mean && ( + {result?.error && (
{result.error}
@@ -352,7 +368,7 @@ export function FactorBacktest({ initialFactorName = 'momentum_20d' }: { initial )} - {result && result.ic_mean != null && ( + {result && !result.error && ( - Rank IC · 日度调仓 + Rank IC · {rebalance === 'daily' ? '日度' : rebalance === 'weekly' ? '周度' : '月度'}调仓 {result.elapsed_ms > 0 && ( @@ -384,24 +400,28 @@ export function FactorBacktest({ initialFactorName = 'momentum_20d' }: { initial )} -
- 0.03 ? 'bull' : result.ic_mean < -0.03 ? 'bear' : 'neutral' - : undefined} - /> - - 0.5 ? (result.ir > 0 ? 'bull' : 'bear') : 'neutral' - : undefined} - /> - -
+ {result.ic_mean != null ? ( +
+ 0.03 ? 'bull' : result.ic_mean < -0.03 ? 'bear' : 'neutral'} + /> + + 0.5 ? (result.ir > 0 ? 'bull' : 'bear') : 'neutral' + : undefined} + /> + +
+ ) : ( +
+ 标的数量过少,无法计算 IC/IR(需 ≥2 只)。可清空标的使用全市场,或补充更多标的。 +
+ )} {/* IC 时序图 */}