From 32dbcb8958dbfd401f8dfa9c87b15b1b12a8839e Mon Sep 17 00:00:00 2001 From: claude_dev Date: Wed, 29 Jul 2026 07:49:35 +0800 Subject: [PATCH] =?UTF-8?q?feat(portfolio):=20P2/P3=20=E7=AD=96=E7=95=A5?= =?UTF-8?q?=E5=B1=82=E5=90=91=E9=87=8F=E5=8C=96=20+=20fundamentals?= =?UTF-8?q?=E6=89=B9=E9=87=8F=E6=8F=90=E9=80=9F=E8=A7=A3=E9=94=81=E9=95=BF?= =?UTF-8?q?=E5=9B=9E=E6=B5=8B?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit P2 行情向量化(get_price→get_closes_panel,口径实证 max_abs_diff=0.0 零偏差): - momentum_timing: _cal_rps/_select_stocks/_cal_buy_sign 三处向量化 - small_cap: _cal_momentum_score 用 close.min/max 代理 low/high(方案A) P3 fundamentals 批量(small_cap _pick_stocks 加 fields=[market_cap,eps], 对接数据session f416a17 get_fundamentals_df 按需短路): - 02 000985 全市场 5128只 ~19min卡死 → 195s 跑通解锁 验收: 72单测全过; VPS 03短回测+138%(口径与get_price一致diff=0)/ 02全市场-11%(2024Q1小盘股灾期合理)/01沪深300可跑 --- .../strategies/momentum_timing.py | 150 +++++---- sanguo_portfolio/strategies/small_cap.py | 85 +++--- tests/portfolio/test_momentum_timing.py | 285 +++++++++--------- tests/portfolio/test_small_cap.py | 142 ++++----- 4 files changed, 325 insertions(+), 337 deletions(-) diff --git a/sanguo_portfolio/strategies/momentum_timing.py b/sanguo_portfolio/strategies/momentum_timing.py index e88cfb6..da28358 100644 --- a/sanguo_portfolio/strategies/momentum_timing.py +++ b/sanguo_portfolio/strategies/momentum_timing.py @@ -210,6 +210,9 @@ class MomentumTimingStrategy: 改为取 ``preDate ~ curDate`` 区间算**百分比涨跌幅**(更符合 RPS 语义, 原代码用绝对差值排序会偏向高价股,见 notes.md「移植记录」)。 + 性能改造:用 ``get_closes_panel`` 一次批量取宽表(替代 get_price+pivot, + 5128 只从逐只循环秒级降到 UNION ALL 批量)。数值口径不变(fq='raw')。 + Returns: DataFrame[code, rps_value],按 rps_value 降序;``rps_value = 99 - 100*i/n``。 """ @@ -217,31 +220,18 @@ class MomentumTimingStrategy: if n == 0: return pd.DataFrame({"code": [], "rps_value": []}) try: - df = self.provider.get_price( - stocks, - start_date=pre_date, - end_date=cur_date, - frequency="daily", - fields=["close"], - panel=False, - fill_paused=False, + panel = self.provider.get_closes_panel( + stocks, pre_date, cur_date, fq="raw", ) except Exception as exc: - logger.warning("_cal_rps get_price 失败: %s", exc) + logger.warning("_cal_rps get_closes_panel 失败: %s", exc) return pd.DataFrame({"code": [], "rps_value": []}) - if df is None or df.empty: - return pd.DataFrame({"code": [], "rps_value": []}) - try: - pivot = df.pivot(index="time", columns="code", values="close") - except Exception as exc: - logger.warning("_cal_rps pivot 失败: %s", exc) - return pd.DataFrame({"code": [], "rps_value": []}) - if pivot.empty or len(pivot) < 2: + if panel is None or panel.empty or len(panel) < 2: return pd.DataFrame({"code": [], "rps_value": []}) - # 每只股票涨跌幅(末值/首值 - 1) - first = pivot.iloc[0] - last = pivot.iloc[-1] + # 每只股票涨跌幅(末值/首值 - 1) — 向量化 + first = panel.iloc[0] + last = panel.iloc[-1] with np.errstate(divide="ignore", invalid="ignore"): returns = (last / first) - 1.0 # 过滤 NaN/Inf(数据不全或首值为 0) @@ -282,46 +272,43 @@ class MomentumTimingStrategy: """均线动量过滤:``close > MA_short`` 且 ``MA_short > MA_long``。 原策略 ``data[security].mavg(5,'close')`` (聚宽 Security.mavg), - 翻译为 provider.get_price(count=ma_long) 后段求均值。 + 翻译为 ``get_closes_panel`` 批量取 close 宽表,向量化算双均线。 + + 性能改造:用 ``get_closes_panel`` 替代 ``get_price(count=N)``。 + count=N → ``start = cur - N*2 自然日``、``end = cur``、宽表 ``.tail(N)`` 切片 + (避免自然日 vs 交易日的换算歧义,``.tail(N)`` 最稳)。数值口径不变(fq='raw')。 """ cfg = self.config if not stocks: return [] + start_date = _shift_date(cur_date, -cfg.ma_long * 2) try: - df = self.provider.get_price( - stocks, - end_date=cur_date, - frequency="daily", - fields=["close"], - count=cfg.ma_long, - panel=False, - fill_paused=False, + panel = self.provider.get_closes_panel( + stocks, start_date, cur_date, fq="raw", ) except Exception as exc: - logger.warning("_select_stocks get_price 失败: %s", exc) + logger.warning("_select_stocks get_closes_panel 失败: %s", exc) return [] - if df is None or df.empty: + if panel is None or panel.empty: return [] - try: - pivot = df.pivot(index="time", columns="code", values="close") - except Exception: - return [] - if pivot.empty: + panel = panel.tail(cfg.ma_long) + if panel.empty: return [] - out: List[str] = [] - for col in pivot.columns: - series = pivot[col].dropna() - if len(series) < cfg.ma_long: - continue - close = float(series.iloc[-1]) - ma_short = float(series.tail(cfg.ma_short).mean()) - ma_long = float(series.tail(cfg.ma_long).mean()) - if np.isnan(close) or np.isnan(ma_short) or np.isnan(ma_long): - continue - if close > ma_short and ma_short > ma_long: - out.append(col) - return out + # 向量化:对每列算 close/ma_short/ma_long,过滤条件 close>ma_short>ma_long + valid_count = panel.notna().sum() # 每列非 NaN 计数,等价于原 dropna 后长度 + close = panel.iloc[-1] + ma_short = panel.tail(cfg.ma_short).mean() + ma_long = panel.mean() # panel 已 tail(ma_long),整体均值即 MA_long + mask = ( + (valid_count >= cfg.ma_long) + & close.notna() + & ma_short.notna() + & ma_long.notna() + & (close > ma_short) + & (ma_short > ma_long) + ) + return list(mask.index[mask]) # =================== calBuySign (牛熊分界) =================== def _cal_buy_sign( @@ -333,8 +320,12 @@ class MomentumTimingStrategy: """统计 past_day 均线上方的指数占比 > index_thre → 牛市(True)。 原策略 'index' 模式(第 110-115 行):对每个指数算 ``mavg(past_day,'close')`` - 与 ``mavg(1,'close')`` 比较。翻译为取 past_day 日 close(含当日), - 算均值与最后一根 close 比较。 + 与 ``mavg(1,'close')`` 比较。翻译为 ``get_closes_panel`` 批量取 close 宽表, + 向量化算 past_day 均线。 + + 性能改造:用 ``get_closes_panel`` 替代 ``get_price(count=N)``。 + count=N → ``start = cur - N*2 自然日``、``end = cur``、宽表 ``.tail(N)`` 切片。 + 数值口径不变(fq='raw')。 ⚠️ 原代码 ``float(count)/len(indexList)`` 在 py2 是浮点除法(因 float()强转), 与 py3 一致。这里保留浮点除法语义。 @@ -342,39 +333,31 @@ class MomentumTimingStrategy: cfg = self.config if not index_list: return False + start_date = _shift_date(cur_date, -past_day * 2) try: - df = self.provider.get_price( - index_list, - end_date=cur_date, - frequency="daily", - fields=["close"], - count=past_day, - panel=False, - fill_paused=False, + panel = self.provider.get_closes_panel( + index_list, start_date, cur_date, fq="raw", ) except Exception as exc: - logger.warning("_cal_buy_sign get_price 失败: %s", exc) + logger.warning("_cal_buy_sign get_closes_panel 失败: %s", exc) return False - if df is None or df.empty: + if panel is None or panel.empty: return False - try: - pivot = df.pivot(index="time", columns="code", values="close") - except Exception: - return False - if pivot.empty: + panel = panel.tail(past_day) + if panel.empty: return False - count = 0 - for col in pivot.columns: - series = pivot[col].dropna() - if len(series) < 2: - continue - ma_past = float(series.tail(past_day).mean()) - cur_close = float(series.iloc[-1]) - if np.isnan(ma_past) or np.isnan(cur_close): - continue - if cur_close > ma_past: - count += 1 + # 向量化:对每列算 cur_close 与 past_day 均值,统计 close > ma_past 的占比 + valid_count = panel.notna().sum() + cur_close = panel.iloc[-1] + ma_past = panel.mean() + mask = ( + (valid_count >= 2) + & cur_close.notna() + & ma_past.notna() + & (cur_close > ma_past) + ) + count = int(mask.sum()) return (count / len(index_list)) > cfg.index_thre # =================== 调仓辅助 =================== @@ -418,4 +401,17 @@ def _to_date_str(value: Any) -> str: return str(value)[:10] +def _shift_date(date_str: str, days: int) -> str: + """字符串日期加减天数,返回 YYYY-MM-DD。 + + 用于 ``count=N`` → ``start = end - N*2 自然日`` 的换算(配合宽表 ``.tail(N)`` 切片, + 避免自然日 vs 交易日的歧义)。 + """ + try: + dt = datetime.datetime.strptime(date_str[:10], "%Y-%m-%d") + except (ValueError, TypeError): + return date_str + return (dt + datetime.timedelta(days=days)).strftime("%Y-%m-%d") + + __all__ = ["MomentumTimingStrategy", "MomentumTimingConfig"] diff --git a/sanguo_portfolio/strategies/small_cap.py b/sanguo_portfolio/strategies/small_cap.py index f17d714..6429083 100644 --- a/sanguo_portfolio/strategies/small_cap.py +++ b/sanguo_portfolio/strategies/small_cap.py @@ -28,6 +28,7 @@ """ from __future__ import annotations +import datetime import logging from dataclasses import dataclass from typing import Any, List, Optional @@ -191,8 +192,14 @@ class SmallCapStrategy: return [] # 2) get_fundamentals_df 一次性取 market_cap + eps + # fields= 按需短路源表(P3 批量提速,数据session commit f416a17: + # 5128 只 ~19min→~2min)。_pick_stocks 只用 market_cap(排序)+eps(>0过滤)。 + # ⚠️ 若以后给 _pick_stocks 加新过滤(ROE/营收等),必须把列名加进 fields=, + # 否则该列返 NaN→过滤静默失效;fields=None 仍全列(向后兼容但慢)。 try: - df = self.provider.get_fundamentals_df(candidates, date=previous_date) + df = self.provider.get_fundamentals_df( + candidates, date=previous_date, fields=["market_cap", "eps"], + ) except Exception as exc: logger.warning("get_fundamentals_df 失败: %s", exc) return [] @@ -260,6 +267,12 @@ class SmallCapStrategy: - ``score = (cur-low_130) + (cur-high_130) + (cur-avg_15)`` - 升序(分数越低越靠前:price 接近 130 日低 / 低于均线 → 偏底部) + 性能改造(决策方案 A):用 ``get_closes_panel`` 批量取 close 宽表向量化, + **用 close.rolling(130).min/max 代理 low.min()/high.max()**。 + 原因:``get_closes_panel`` 只返 close(不含 high/low);用 close 极值代理是有意 + 决策(spec 明确允许)——对动量评分的"底部反弹偏好"语义无实质影响(都衡量 + 当前价在 130 日极值区间的位置),换 5128 只逐只循环 → 一次批量(33s → 秒级)。 + py2→py3:``df.sort(columns=)`` → ``df.sort_values(by=)``。 Returns: @@ -269,49 +282,36 @@ class SmallCapStrategy: if not stocks: return pd.DataFrame(columns=["score"]) - # 一次性取 ma_window=130 日 close/high/low(对所有候选) + # 一次性取 ma_window=130 日 close(批量宽表,用 min/max 代理 low/high) + start_date = _shift_date(end_date, -cfg.ma_window * 2) try: - df = self.provider.get_price( - stocks, - end_date=end_date, - frequency="daily", - fields=["close", "high", "low"], - count=cfg.ma_window, - panel=False, - fill_paused=False, + panel = self.provider.get_closes_panel( + stocks, start_date, end_date, fq="raw", ) except Exception as exc: - logger.warning("_cal_momentum_score get_price 失败: %s", exc) + logger.warning("_cal_momentum_score get_closes_panel 失败: %s", exc) return pd.DataFrame(columns=["score"]) - if df is None or df.empty: + if panel is None or panel.empty: + return pd.DataFrame(columns=["score"]) + panel = panel.tail(cfg.ma_window) + if panel.empty: return pd.DataFrame(columns=["score"]) - scores: dict[str, float] = {} - for code in stocks: - sub = df[df["code"] == code] if "code" in df.columns else df - if sub is None or sub.empty: - continue - close_series = sub["close"].dropna() if "close" in sub.columns else None - high_series = sub["high"].dropna() if "high" in sub.columns else None - low_series = sub["low"].dropna() if "low" in sub.columns else None - if close_series is None or close_series.empty: - continue - cur_price = float(close_series.iloc[-1]) - if not np.isfinite(cur_price): - continue - # 130 日最低 / 最高(skip_paused=True 后 dropna) - low_130 = float(low_series.min()) if low_series is not None and not low_series.empty else cur_price - high_130 = float(high_series.max()) if high_series is not None and not high_series.empty else cur_price - # 15 日均线:close 序列最后 15 根均值 - ma15 = float(close_series.tail(cfg.ma_short).mean()) if len(close_series) >= 1 else cur_price - if not (np.isfinite(low_130) and np.isfinite(high_130) and np.isfinite(ma15)): - continue - score = (cur_price - low_130) + (cur_price - high_130) + (cur_price - ma15) - scores[code] = score + # 向量化算 score = (cur-low) + (cur-high) + (cur-ma15) + # 低/高用 close 序列代理(原 high.max()/low.min()) + cur_price = panel.iloc[-1] + low_proxy = panel.min() # 130 日 close 最低(代理 low.min()) + high_proxy = panel.max() # 130 日 close 最高(代理 high.max()) + ma15 = panel.tail(cfg.ma_short).mean() - if not scores: + # 过滤:cur_price 必须有效 + 至少 1 个有效值(原代码 close_series.empty 跳过) + valid_count = panel.notna().sum() + score = (cur_price - low_proxy) + (cur_price - high_proxy) + (cur_price - ma15) + mask = (valid_count >= 1) & cur_price.notna() & np.isfinite(cur_price) + score = score[mask].dropna() + if score.empty: return pd.DataFrame(columns=["score"]) - out = pd.DataFrame.from_dict(scores, orient="index", columns=["score"]) + out = score.to_frame("score") # 升序:分数越低越靠前(原策略 df.sort(columns='score', ascending=True)) out = out.sort_values("score", ascending=True) return out @@ -407,4 +407,17 @@ def _is_valid_positive_number(v: Any) -> bool: return fv > 0 +def _shift_date(date_str: str, days: int) -> str: + """字符串日期加减天数,返回 YYYY-MM-DD。 + + 用于 ``count=N`` → ``start = end - N*2 自然日`` 的换算(配合宽表 ``.tail(N)`` 切片, + 避免自然日 vs 交易日的歧义)。 + """ + try: + dt = datetime.datetime.strptime(date_str[:10], "%Y-%m-%d") + except (ValueError, TypeError): + return date_str + return (dt + datetime.timedelta(days=days)).strftime("%Y-%m-%d") + + __all__ = ["SmallCapStrategy", "SmallCapConfig"] diff --git a/tests/portfolio/test_momentum_timing.py b/tests/portfolio/test_momentum_timing.py index 10daa6e..893c6f8 100644 --- a/tests/portfolio/test_momentum_timing.py +++ b/tests/portfolio/test_momentum_timing.py @@ -2,6 +2,9 @@ 策略层只测**逻辑分支正确**(RPS / 均线 / 牛熊信号 / 调仓),不测真实数据。 真实数据回测在 VPS 跑,这里只保证策略翻译等价 + 两个原始 bug 已修复。 + +性能改造后:策略用 ``get_closes_panel`` 批量取 close 宽表(替代 get_price+pivot), +测试 mock 同步切到 ``get_closes_panel`` 返回宽表 DataFrame。 """ from __future__ import annotations @@ -25,13 +28,13 @@ from tests.portfolio.conftest import FakeContext, FakePosition def make_strategy( *, index_stocks_map: Optional[Dict[str, List[str]]] = None, - price_df_map: Optional[Dict[Any, pd.DataFrame]] = None, + panel_map: Optional[Dict[Any, pd.DataFrame]] = None, config: Optional[MomentumTimingConfig] = None, ) -> MomentumTimingStrategy: """构造一个 mock provider + mock broker 驱动的策略。 - index_stocks_map: get_index_stocks 返回,dict[index] -> List[code] - - price_df_map: get_price 按 (security, fields, count) 或 (security, start, end) 缓存的返回 + - panel_map: get_closes_panel 按 (symbols_tuple) 或 (symbols_tuple, start, end) 缓存的宽表返回 """ provider = MagicMock(name="provider") @@ -59,29 +62,28 @@ def make_strategy( "high_limit": 11.0, "low_limit": 9.0, } - # get_price 按 key 缓存(支持 count 模式 + start/end 模式) - # 规范化:把 key 第一项(list)转 tuple 以保证可 hash - def _normalize_key(k: Any) -> Any: + # get_closes_panel 按 key 缓存:支持精确 (tuple(symbols), start, end) 与 + # 通配 (tuple(symbols),) 两种 key(精确优先,fallback 忽略日期) + def _normalize_syms(k: Any) -> Any: if isinstance(k, tuple) and k and isinstance(k[0], (list, tuple)): return (tuple(k[0]),) + tuple(k[1:]) return k - price_df_map = {_normalize_key(k): v for k, v in (price_df_map or {}).items()} + panel_map = {_normalize_syms(k): v for k, v in (panel_map or {}).items()} - def _get_price(security, **kwargs): - # 构造 cache key:两种取数模式 - # 1) count 模式:(sec_key, fields, count) - # 2) start/end 模式:(sec_key, fields, start_date, end_date) - # 注意:list 不可 hash → 转 tuple - sec_key = tuple(security) if isinstance(security, list) else security - fields = tuple(kwargs.get("fields") or []) - if kwargs.get("count") is not None: - key = (sec_key, fields, kwargs.get("count")) - else: - key = (sec_key, fields, kwargs.get("start_date"), kwargs.get("end_date")) - return price_df_map.get(key, pd.DataFrame()) + def _get_closes_panel(symbols, start=None, end=None, interval="d", fq="raw"): + syms_key = tuple(symbols) if isinstance(symbols, list) else symbols + # 精确 key 优先 + exact = panel_map.get((syms_key, start, end)) + if exact is not None: + return exact + # fallback: 忽略 start/end + fallback = panel_map.get((syms_key,)) + if fallback is not None: + return fallback + return pd.DataFrame(index=pd.DatetimeIndex([])) - provider.get_price.side_effect = _get_price + provider.get_closes_panel.side_effect = _get_closes_panel broker = BrokerFacade() broker.order_target_value = MagicMock(return_value=MagicMock(filled=100)) @@ -94,29 +96,30 @@ def make_strategy( return MomentumTimingStrategy(provider=provider, broker=broker, config=config) -def _make_close_panel( +def _make_close_wide( codes: List[str], closes: List[List[float]], end_date: str = "2024-09-30", days: int = 30, ) -> pd.DataFrame: - """构造 panel=False 风格的 close DataFrame。 + """构造 get_closes_panel 风格的 close 宽表 (index=DatetimeIndex, columns=codes)。 Args: - codes: 股票代码列表 + codes: 股票代码列表(列名) closes: 每只股票的 close 序列(长度 <= days, 不足重复首值) end_date: 最后一根 K 线日期 days: 总 K 线根数(默认 30) """ end_dt = datetime.strptime(end_date, "%Y-%m-%d") - dates = [(end_dt - timedelta(days=days - 1 - i)).strftime("%Y-%m-%d") for i in range(days)] - rows = [] + dates = pd.DatetimeIndex([ + end_dt - timedelta(days=days - 1 - i) for i in range(days) + ]) + data: Dict[str, List[float]] = {} for code, close_list in zip(codes, closes): # 不足 days 的补首值 full = list(close_list) + [close_list[-1]] * (days - len(close_list)) - for d, c in zip(dates, full): - rows.append({"time": pd.Timestamp(d), "code": code, "close": float(c)}) - return pd.DataFrame(rows) + data[code] = [float(c) for c in full] + return pd.DataFrame(data, index=dates) # =================== initialize =================== @@ -153,17 +156,16 @@ class TestCalRps: # 3 只股票,涨幅依次为 +100% / +50% / 0% # preDate 首值 = 10, curDate 末值 = 20 / 15 / 10 codes = ["A.XSHG", "B.XSHG", "C.XSHG"] - df = pd.DataFrame([ - {"time": pd.Timestamp("2024-09-01"), "code": "A.XSHG", "close": 10.0}, - {"time": pd.Timestamp("2024-09-30"), "code": "A.XSHG", "close": 20.0}, # +100% - {"time": pd.Timestamp("2024-09-01"), "code": "B.XSHG", "close": 10.0}, - {"time": pd.Timestamp("2024-09-30"), "code": "B.XSHG", "close": 15.0}, # +50% - {"time": pd.Timestamp("2024-09-01"), "code": "C.XSHG", "close": 10.0}, - {"time": pd.Timestamp("2024-09-30"), "code": "C.XSHG", "close": 10.0}, # 0% - ]) - s = make_strategy(price_df_map={ - # 按 start_date/end_date 取数,确认 _cal_rps 走的是区间查询 - ((str(codes),) if False else (tuple(codes), ("close",), "2024-09-01", "2024-09-30")): df, + panel = pd.DataFrame( + { + "A.XSHG": [10.0, 20.0], # +100% + "B.XSHG": [10.0, 15.0], # +50% + "C.XSHG": [10.0, 10.0], # 0% + }, + index=pd.DatetimeIndex(["2024-09-01", "2024-09-30"]), + ) + s = make_strategy(panel_map={ + (tuple(codes), "2024-09-01", "2024-09-30"): panel, }) out = s._cal_rps(codes, cur_date="2024-09-30", pre_date="2024-09-01") @@ -177,14 +179,15 @@ class TestCalRps: def test_rps_descending_by_return(self): """涨幅大的排前(降序)。""" codes = ["X.XSHG", "Y.XSHG"] - df = pd.DataFrame([ - {"time": pd.Timestamp("2024-09-01"), "code": "X.XSHG", "close": 10.0}, - {"time": pd.Timestamp("2024-09-30"), "code": "X.XSHG", "close": 12.0}, # +20% - {"time": pd.Timestamp("2024-09-01"), "code": "Y.XSHG", "close": 10.0}, - {"time": pd.Timestamp("2024-09-30"), "code": "Y.XSHG", "close": 15.0}, # +50% - ]) - s = make_strategy(price_df_map={ - (tuple(codes), ("close",), "2024-09-01", "2024-09-30"): df, + panel = pd.DataFrame( + { + "X.XSHG": [10.0, 12.0], # +20% + "Y.XSHG": [10.0, 15.0], # +50% + }, + index=pd.DatetimeIndex(["2024-09-01", "2024-09-30"]), + ) + s = make_strategy(panel_map={ + (tuple(codes), "2024-09-01", "2024-09-30"): panel, }) out = s._cal_rps(codes, cur_date="2024-09-30", pre_date="2024-09-01") # Y 涨幅大,排前 @@ -194,16 +197,16 @@ class TestCalRps: def test_rps_filters_nan_and_zero_first(self): """首值为 0(除零)或 NaN → 过滤掉。""" codes = ["GOOD.XSHG", "ZERO.XSHG", "NAN.XSHG"] - df = pd.DataFrame([ - {"time": pd.Timestamp("2024-09-01"), "code": "GOOD.XSHG", "close": 10.0}, - {"time": pd.Timestamp("2024-09-30"), "code": "GOOD.XSHG", "close": 20.0}, - {"time": pd.Timestamp("2024-09-01"), "code": "ZERO.XSHG", "close": 0.0}, - {"time": pd.Timestamp("2024-09-30"), "code": "ZERO.XSHG", "close": 10.0}, - {"time": pd.Timestamp("2024-09-01"), "code": "NAN.XSHG", "close": np.nan}, - {"time": pd.Timestamp("2024-09-30"), "code": "NAN.XSHG", "close": 10.0}, - ]) - s = make_strategy(price_df_map={ - (tuple(codes), ("close",), "2024-09-01", "2024-09-30"): df, + panel = pd.DataFrame( + { + "GOOD.XSHG": [10.0, 20.0], + "ZERO.XSHG": [0.0, 10.0], + "NAN.XSHG": [np.nan, 10.0], + }, + index=pd.DatetimeIndex(["2024-09-01", "2024-09-30"]), + ) + s = make_strategy(panel_map={ + (tuple(codes), "2024-09-01", "2024-09-30"): panel, }) out = s._cal_rps(codes, cur_date="2024-09-30", pre_date="2024-09-01") assert list(out["code"]) == ["GOOD.XSHG"] @@ -219,10 +222,9 @@ class TestSelectStocks: """close > MA5 且 MA5 > MA15 → 保留。""" # 构造 15 日 close 序列:上升 → close(末) > MA5 > MA15 rising = [10.0 + i * 0.5 for i in range(15)] # 10→17 - df = _make_close_panel(["UP.XSHG"], [rising], end_date="2024-09-30", days=15) - s = make_strategy(price_df_map={ - # 注意:_select_stocks 传 list,get_price 内部转 tuple → key 第一项必须是 tuple - (("UP.XSHG",), ("close",), 15): df, + panel = _make_close_wide(["UP.XSHG"], [rising], end_date="2024-09-30", days=15) + s = make_strategy(panel_map={ + (("UP.XSHG",),): panel, }) out = s._select_stocks(["UP.XSHG"], cur_date="2024-09-30") assert out == ["UP.XSHG"] @@ -230,9 +232,9 @@ class TestSelectStocks: def test_filter_close_below_ma_short(self): """close < MA5 → 剔除(下行趋势)。""" falling = [20.0 - i * 0.5 for i in range(15)] # 20→13 - df = _make_close_panel(["DOWN.XSHG"], [falling], end_date="2024-09-30", days=15) - s = make_strategy(price_df_map={ - (("DOWN.XSHG"), ("close",), 15): df, + panel = _make_close_wide(["DOWN.XSHG"], [falling], end_date="2024-09-30", days=15) + s = make_strategy(panel_map={ + (("DOWN.XSHG",),): panel, }) out = s._select_stocks(["DOWN.XSHG"], cur_date="2024-09-30") assert out == [] @@ -241,9 +243,9 @@ class TestSelectStocks: """close > MA5 但 MA5 < MA15(下跌但末值小反弹)→ 剔除。""" # 前 10 日大涨(20→30),后 5 日跌(30→26):MA5 < MA15 series = [20 + i for i in range(10)] + [30 - i for i in range(1, 6)] # 20..29, 29..25 - df = _make_close_panel(["FLAT.XSHG"], [series], end_date="2024-09-30", days=15) - s = make_strategy(price_df_map={ - (("FLAT.XSHG"), ("close",), 15): df, + panel = _make_close_wide(["FLAT.XSHG"], [series], end_date="2024-09-30", days=15) + s = make_strategy(panel_map={ + (("FLAT.XSHG",),): panel, }) out = s._select_stocks(["FLAT.XSHG"], cur_date="2024-09-30") # close=25, MA5 = mean(29,28,27,26,25)=27, MA15 = mean(all)=24.67 @@ -252,17 +254,14 @@ class TestSelectStocks: def test_insufficient_data_skipped(self): """不足 ma_long=15 根 → 跳过。""" - short_df = _make_close_panel(["NEW.XSHG"], [[10, 11, 12]], end_date="2024-09-30", days=15) - s = make_strategy(price_df_map={ - (("NEW.XSHG"), ("close",), 15): short_df, + # 直接给短宽表(3 行):panel.notna().sum()=3 < 15 → 全部剔除 + short_panel = pd.DataFrame( + {"NEW.XSHG": [10.0, 11.0, 12.0]}, + index=pd.DatetimeIndex(["2024-09-28", "2024-09-29", "2024-09-30"]), + ) + s = make_strategy(panel_map={ + (("NEW.XSHG",),): short_panel, }) - # 序列被 _make_close_panel 补齐到 15,这里改为真短数据 - s.provider.get_price.side_effect = None - s.provider.get_price.return_value = pd.DataFrame([ - {"time": pd.Timestamp("2024-09-28"), "code": "NEW.XSHG", "close": 10.0}, - {"time": pd.Timestamp("2024-09-29"), "code": "NEW.XSHG", "close": 11.0}, - {"time": pd.Timestamp("2024-09-30"), "code": "NEW.XSHG", "close": 12.0}, - ]) out = s._select_stocks(["NEW.XSHG"], cur_date="2024-09-30") assert out == [] @@ -278,9 +277,9 @@ class TestCalBuySign: # 上升序列:末值远高于均值 idx_list = ["000300.XSHG", "000905.XSHG"] rising = [10.0 + i for i in range(30)] # 10→39 - df = _make_close_panel(idx_list, [rising, rising], end_date="2024-09-30", days=30) - s = make_strategy(price_df_map={ - (tuple(idx_list), ("close",), 30): df, + panel = _make_close_wide(idx_list, [rising, rising], end_date="2024-09-30", days=30) + s = make_strategy(panel_map={ + (tuple(idx_list),): panel, }) assert s._cal_buy_sign(idx_list, past_day=30, cur_date="2024-09-30") is True @@ -288,9 +287,9 @@ class TestCalBuySign: """所有指数都跌破 30 日均线 → 占比 0% < 20% → 熊市(False)。""" idx_list = ["000300.XSHG", "000905.XSHG"] falling = [40.0 - i for i in range(30)] # 40→11 - df = _make_close_panel(idx_list, [falling, falling], end_date="2024-09-30", days=30) - s = make_strategy(price_df_map={ - (tuple(idx_list), ("close",), 30): df, + panel = _make_close_wide(idx_list, [falling, falling], end_date="2024-09-30", days=30) + s = make_strategy(panel_map={ + (tuple(idx_list),): panel, }) assert s._cal_buy_sign(idx_list, past_day=30, cur_date="2024-09-30") is False @@ -301,18 +300,18 @@ class TestCalBuySign: falling = [40.0 - i for i in range(30)] # 2 个 rising + 7 个 falling series_list = [rising, rising] + [falling] * 7 - df = _make_close_panel(idx_list, series_list, end_date="2024-09-30", days=30) - s = make_strategy(price_df_map={ - (tuple(idx_list), ("close",), 30): df, + panel = _make_close_wide(idx_list, series_list, end_date="2024-09-30", days=30) + s = make_strategy(panel_map={ + (tuple(idx_list),): panel, }) # 2/9 ≈ 0.222 > 0.2 → 牛市 assert s._cal_buy_sign(idx_list, past_day=30, cur_date="2024-09-30") is True # 改为 1 个 rising:1/9 ≈ 0.111 < 0.2 → 熊市 series_list_1 = [rising] + [falling] * 8 - df_1 = _make_close_panel(idx_list, series_list_1, end_date="2024-09-30", days=30) - s.provider.get_price.side_effect = None - s.provider.get_price.return_value = df_1 + panel_1 = _make_close_wide(idx_list, series_list_1, end_date="2024-09-30", days=30) + s.provider.get_closes_panel.side_effect = None + s.provider.get_closes_panel.return_value = panel_1 assert s._cal_buy_sign(idx_list, past_day=30, cur_date="2024-09-30") is False @@ -322,9 +321,9 @@ class TestHandleData: """熊市信号 → 全部持仓清掉。""" cfg = MomentumTimingConfig(index_list=["IDX.XSHG"]) s = make_strategy(config=cfg) - # 触发熊市:get_price 返回下行 close - s.provider.get_price.side_effect = None - s.provider.get_price.return_value = _make_close_panel( + # 触发熊市:get_closes_panel 返回下行 close + s.provider.get_closes_panel.side_effect = None + s.provider.get_closes_panel.return_value = _make_close_wide( ["IDX.XSHG"], [[40.0 - i for i in range(30)]], end_date="2024-10-08", days=30, ) @@ -356,35 +355,22 @@ class TestHandleData: rising_30 = [10.0 + i for i in range(30)] # 牛市信号用 rising_15 = [10.0 + i for i in range(15)] # 均线筛选用 - # 提供所有可能查询路径的 price 数据 + # 提供所有可能查询路径的 close 宽表 idx_codes = ["IDX.XSHG"] stock_codes = ["CAND.XSHG"] - def _gp(security, **kwargs): - fields = tuple(kwargs.get("fields") or []) - # 1) _cal_buy_sign: idx_list, count=30 - if security == idx_codes and kwargs.get("count") == 30: - return _make_close_panel(idx_codes, [rising_30], days=30) - # 2) _cal_rps for index 股池:股票, start/end 模式 - if ( - isinstance(security, list) - and security == stock_codes - and kwargs.get("start_date") - ): - return pd.DataFrame([ - {"time": pd.Timestamp("2024-09-01"), "code": "CAND.XSHG", "close": 10.0}, - {"time": pd.Timestamp("2024-10-08"), "code": "CAND.XSHG", "close": 39.0}, - ]) - # 3) _select_stocks: count=15 - if ( - isinstance(security, list) - and security == stock_codes - and kwargs.get("count") == cfg.ma_long - ): - return _make_close_panel(stock_codes, [rising_15], days=15) - return pd.DataFrame() + def _gcp(symbols, start=None, end=None, interval="d", fq="raw"): + if symbols == idx_codes: + # _cal_buy_sign: idx_list, tail(30) + return _make_close_wide(idx_codes, [rising_30], days=30) + if isinstance(symbols, list) and symbols == stock_codes: + # _cal_rps (2 行 start/end) 与 _select_stocks (15 行) 都用同一份上升序列 + # 取 15 行,_cal_rps 用首末(_shift_date 算出 start), + # _select_stocks 用全部 15 行(ma_long=15) + return _make_close_wide(stock_codes, [rising_15], days=15) + return pd.DataFrame(index=pd.DatetimeIndex([])) - s.provider.get_price.side_effect = _gp + s.provider.get_closes_panel.side_effect = _gcp ctx = FakeContext( current_dt=datetime(2024, 10, 8, 9, 30), @@ -405,27 +391,27 @@ class TestHandleData: """⚠️ 修复原始 bug 验证:handle_data 必须用 context.current_dt 计算 cur_date, 不能用 datetime.date.today()(后者取真实今天)。 """ - # 用一个明显不同的 current_dt,确认 get_price 的 end_date 跟随它 + # 用一个明显不同的 current_dt,确认 get_closes_panel 的 end 跟随它 cfg = MomentumTimingConfig(index_list=["IDX.XSHG"]) s = make_strategy(config=cfg) - captured_end_dates: List[Any] = [] + captured_ends: List[Any] = [] - def _gp(security, **kwargs): - # 记录 end_date 用于断言 - if kwargs.get("end_date"): - captured_end_dates.append(str(kwargs["end_date"])) + def _gcp(symbols, start=None, end=None, interval="d", fq="raw"): + # 记录 end 用于断言 + if end is not None: + captured_ends.append(str(end)) # 下行 → 熊市(快速 return,不查其他) - return _make_close_panel( + return _make_close_wide( ["IDX.XSHG"], [[40.0 - i for i in range(30)]], - end_date=str(kwargs.get("end_date", "2024-10-08"))[:10], + end_date=str(end or "2024-10-08")[:10], days=30, ) - s.provider.get_price.side_effect = _gp + s.provider.get_closes_panel.side_effect = _gcp ctx = FakeContext(current_dt=datetime(2024, 10, 8, 9, 30)) s.handle_data(ctx) - # 至少一次 get_price 的 end_date 是 "2024-10-08"(来自 current_dt),非今天 - assert any("2024-10-08" in d for d in captured_end_dates) + # 至少一次 get_closes_panel 的 end 是 "2024-10-08"(来自 current_dt),非今天 + assert any("2024-10-08" in d for d in captured_ends) # =================== _find_stock_pool (取强舍弱) =================== @@ -441,30 +427,31 @@ class TestFindStockPool: config=cfg, ) # 涨幅:A=+100%, B=+50%, C=0%, D=+30%, E=-10% - rps_df_1 = pd.DataFrame([ - {"time": pd.Timestamp("2024-09-01"), "code": "A.XSHG", "close": 10.0}, - {"time": pd.Timestamp("2024-09-30"), "code": "A.XSHG", "close": 20.0}, - {"time": pd.Timestamp("2024-09-01"), "code": "B.XSHG", "close": 10.0}, - {"time": pd.Timestamp("2024-09-30"), "code": "B.XSHG", "close": 15.0}, - {"time": pd.Timestamp("2024-09-01"), "code": "C.XSHG", "close": 10.0}, - {"time": pd.Timestamp("2024-09-30"), "code": "C.XSHG", "close": 10.0}, - ]) - rps_df_2 = pd.DataFrame([ - {"time": pd.Timestamp("2024-09-01"), "code": "D.XSHG", "close": 10.0}, - {"time": pd.Timestamp("2024-09-30"), "code": "D.XSHG", "close": 13.0}, - {"time": pd.Timestamp("2024-09-01"), "code": "E.XSHG", "close": 10.0}, - {"time": pd.Timestamp("2024-09-30"), "code": "E.XSHG", "close": 9.0}, - ]) + rps_panel_1 = pd.DataFrame( + { + "A.XSHG": [10.0, 20.0], + "B.XSHG": [10.0, 15.0], + "C.XSHG": [10.0, 10.0], + }, + index=pd.DatetimeIndex(["2024-09-01", "2024-09-30"]), + ) + rps_panel_2 = pd.DataFrame( + { + "D.XSHG": [10.0, 13.0], + "E.XSHG": [10.0, 9.0], + }, + index=pd.DatetimeIndex(["2024-09-01", "2024-09-30"]), + ) - def _gp(security, **kwargs): - if isinstance(security, list): - if "A.XSHG" in security: - return rps_df_1 - if "D.XSHG" in security: - return rps_df_2 - return pd.DataFrame() + def _gcp(symbols, start=None, end=None, interval="d", fq="raw"): + if isinstance(symbols, list): + if "A.XSHG" in symbols: + return rps_panel_1 + if "D.XSHG" in symbols: + return rps_panel_2 + return pd.DataFrame(index=pd.DatetimeIndex([])) - s.provider.get_price.side_effect = _gp + s.provider.get_closes_panel.side_effect = _gcp out = s._find_stock_pool( ["IDX1.XSHG", "IDX2.XSHG"], cur_date="2024-09-30", pre_date="2024-09-01", ) diff --git a/tests/portfolio/test_small_cap.py b/tests/portfolio/test_small_cap.py index 984816f..4dc216b 100644 --- a/tests/portfolio/test_small_cap.py +++ b/tests/portfolio/test_small_cap.py @@ -9,6 +9,10 @@ - ✅ 5 日调仓周期:day_count % tc == 0 时选股+调仓,其他日 no-op - ✅ 等权 20 只 - ❌ 对冲部分(已删,不测) + +性能改造后:策略用 ``get_closes_panel`` 批量取 close 宽表(替代 get_price+逐只循环), +**用 close.min/max 代理 low/high**(有意决策,见 small_cap._cal_momentum_score docstring)。 +测试 mock 同步切到 ``get_closes_panel`` 返回 close 宽表。 """ from __future__ import annotations @@ -33,14 +37,16 @@ def make_strategy( *, universe_stocks: Optional[List[str]] = None, fundamentals_df: Optional[pd.DataFrame] = None, - price_df_map: Optional[Dict[Any, pd.DataFrame]] = None, + panel_map: Optional[Dict[Any, pd.DataFrame]] = None, + default_panel: Optional[pd.DataFrame] = None, config: Optional[SmallCapConfig] = None, ) -> SmallCapStrategy: """构造一个 mock provider + mock broker 驱动的策略。 - universe_stocks: get_index_stocks(universe, date) 返回的全市场候选列表 - fundamentals_df: get_fundamentals_df 返回(index=code, cols=[market_cap, eps, ...]) - - price_df_map: get_price 按 (security_tuple, fields_tuple, count) 缓存的返回 + - panel_map: get_closes_panel 按 (symbols_tuple,) 缓存的宽表返回 + - default_panel: 所有未命中 panel_map 的查询返回的默认宽表(便于 _pick_stocks 流程测试) """ provider = MagicMock(name="provider") @@ -73,24 +79,27 @@ def make_strategy( else: provider.get_fundamentals_df.return_value = pd.DataFrame() - # get_price 按 key 缓存 - def _normalize_key(k: Any) -> Any: + # get_closes_panel 按 key 缓存:支持 (tuple(symbols),) 通配与精确 (tuple, start, end) + def _normalize_syms(k: Any) -> Any: if isinstance(k, tuple) and k and isinstance(k[0], (list, tuple)): return (tuple(k[0]),) + tuple(k[1:]) return k - price_df_map = {_normalize_key(k): v for k, v in (price_df_map or {}).items()} + panel_map = {_normalize_syms(k): v for k, v in (panel_map or {}).items()} - def _get_price(security, **kwargs): - sec_key = tuple(security) if isinstance(security, list) else security - fields = tuple(kwargs.get("fields") or []) - if kwargs.get("count") is not None: - key = (sec_key, fields, kwargs.get("count")) - else: - key = (sec_key, fields, kwargs.get("start_date"), kwargs.get("end_date")) - return price_df_map.get(key, pd.DataFrame()) + def _get_closes_panel(symbols, start=None, end=None, interval="d", fq="raw"): + syms_key = tuple(symbols) if isinstance(symbols, list) else symbols + exact = panel_map.get((syms_key, start, end)) + if exact is not None: + return exact + fallback = panel_map.get((syms_key,)) + if fallback is not None: + return fallback + if default_panel is not None: + return default_panel + return pd.DataFrame(index=pd.DatetimeIndex([])) - provider.get_price.side_effect = _get_price + provider.get_closes_panel.side_effect = _get_closes_panel broker = BrokerFacade() broker.order_target_value = MagicMock(return_value=MagicMock(filled=100)) @@ -119,46 +128,29 @@ def _make_fundamentals_df( return df.set_index("code", drop=False) -def _make_hlc_panel( +def _make_close_wide( stocks: List[str], closes: List[List[float]], *, - highs: Optional[List[List[float]]] = None, - lows: Optional[List[List[float]]] = None, end_date: str = "2024-09-30", days: int = 130, ) -> pd.DataFrame: - """构造 panel=False 风格的 close+high+low DataFrame。 + """构造 get_closes_panel 风格的 close 宽表 (index=DatetimeIndex, columns=stocks)。 Args: - stocks: 股票代码列表 + stocks: 股票代码列表(列名) closes: 每只股票的 close 序列(长度 <= days, 不足重复首值) - highs: 同 close,None → 取 close - lows: 同 close,None → 取 close days: 总 K 线根数(默认 130) """ end_dt = datetime.strptime(end_date, "%Y-%m-%d") - dates = [ - (end_dt - timedelta(days=days - 1 - i)).strftime("%Y-%m-%d") - for i in range(days) - ] - rows = [] - for idx, code in enumerate(stocks): - close_list = closes[idx] - high_list = highs[idx] if highs else close_list - low_list = lows[idx] if lows else close_list - c_full = list(close_list) + [close_list[-1]] * (days - len(close_list)) - h_full = list(high_list) + [high_list[-1]] * (days - len(high_list)) - l_full = list(low_list) + [low_list[-1]] * (days - len(low_list)) - for d, c, h, l in zip(dates, c_full, h_full, l_full): - rows.append({ - "time": pd.Timestamp(d), - "code": code, - "close": float(c), - "high": float(h), - "low": float(l), - }) - return pd.DataFrame(rows) + dates = pd.DatetimeIndex([ + end_dt - timedelta(days=days - 1 - i) for i in range(days) + ]) + data: Dict[str, List[float]] = {} + for code, close_list in zip(stocks, closes): + full = list(close_list) + [close_list[-1]] * (days - len(close_list)) + data[code] = [float(c) for c in full] + return pd.DataFrame(data, index=dates) # =================== initialize =================== @@ -239,6 +231,7 @@ class TestCalMomentumScore: def test_score_formula_is_cur_minus_low_high_ma15(self): """score = (cur-low_130) + (cur-high_130) + (cur-ma15)。 + 改造后用 close.min/max 代理 low/high(有意决策,见 small_cap docstring)。 构造已知序列验证公式: - close 全 10(平):low=high=ma15=10,cur=10,score=0 - close 上升:cur>low/high/ma15 → score 正 @@ -248,13 +241,13 @@ class TestCalMomentumScore: rising = [10.0 + i * 0.1 for i in range(130)] # 10→22.9,cur=22.9 falling = [23.0 - i * 0.1 for i in range(130)] # 23→10.1,cur=10.1 - df = _make_hlc_panel( + panel = _make_close_wide( ["FLAT.XSHG", "UP.XSHG", "DOWN.XSHG"], [flat, rising, falling], end_date="2024-09-30", days=130, ) - s = make_strategy(price_df_map={ - (("FLAT.XSHG", "UP.XSHG", "DOWN.XSHG"), ("close", "high", "low"), 130): df, + s = make_strategy(panel_map={ + (("FLAT.XSHG", "UP.XSHG", "DOWN.XSHG"),): panel, }) out = s._cal_momentum_score( ["FLAT.XSHG", "UP.XSHG", "DOWN.XSHG"], end_date="2024-09-30", @@ -262,11 +255,11 @@ class TestCalMomentumScore: # FLAT: score = 0(全部相同) assert out.loc["FLAT.XSHG", "score"] == pytest.approx(0.0, abs=0.01) - # UP: cur=22.9, low=10, high=22.9, ma15=mean([21.5..22.9])≈22.2 - # score = (22.9-10) + (22.9-22.9) + (22.9-22.2) = 12.9 + 0 + 0.7 ≈ 13.6 + # UP: 用 close 代理后 low=10(close min)、high=22.9(close max)、 + # ma15=mean([21.5..22.9])≈22.2,score = (22.9-10) + (22.9-22.9) + (22.9-22.2) ≈ 13.6 assert out.loc["UP.XSHG", "score"] > 0 - # DOWN: cur=10.1, low=10.1, high=23, ma15≈10.8 - # score = (10.1-10.1) + (10.1-23) + (10.1-10.8) ≈ 0 + (-12.9) + (-0.7) ≈ -13.6 + # DOWN: 用 close 代理后 low=10.1、high=23、ma15≈10.8 + # score = (10.1-10.1) + (10.1-23) + (10.1-10.8) ≈ -13.6 assert out.loc["DOWN.XSHG", "score"] < 0 def test_score_sorted_ascending(self): @@ -274,13 +267,13 @@ class TestCalMomentumScore: flat = [10.0] * 130 rising = [10.0 + i * 0.1 for i in range(130)] falling = [23.0 - i * 0.1 for i in range(130)] - df = _make_hlc_panel( + panel = _make_close_wide( ["FLAT.XSHG", "UP.XSHG", "DOWN.XSHG"], [flat, rising, falling], end_date="2024-09-30", days=130, ) - s = make_strategy(price_df_map={ - (("FLAT.XSHG", "UP.XSHG", "DOWN.XSHG"), ("close", "high", "low"), 130): df, + s = make_strategy(panel_map={ + (("FLAT.XSHG", "UP.XSHG", "DOWN.XSHG"),): panel, }) out = s._cal_momentum_score( ["FLAT.XSHG", "UP.XSHG", "DOWN.XSHG"], end_date="2024-09-30", @@ -291,9 +284,9 @@ class TestCalMomentumScore: def test_insufficient_data_skipped(self): """K 线序列不足/空 → 该股跳过(不在结果里)。""" s = make_strategy() - # 让 provider.get_price 返回空 DataFrame - s.provider.get_price.side_effect = None - s.provider.get_price.return_value = pd.DataFrame() + # 让 provider.get_closes_panel 返回空 DataFrame + s.provider.get_closes_panel.side_effect = None + s.provider.get_closes_panel.return_value = pd.DataFrame(index=pd.DatetimeIndex([])) out = s._cal_momentum_score(["EMPTY.XSHG"], end_date="2024-09-30") assert out.empty @@ -320,14 +313,13 @@ class TestPickStocks: # 不传 price → _cal_momentum_score 会拿到空 df → 结果可能为空 # 我们只验证 eps 过滤生效:在 fundamentals 过滤后 top_candidates 不含 B/C # 直接调 _pick_stocks 会因 price 空导致评分为空 → 返回空 - # 这里通过 mock price 给所有候选相同 close,看最终名单 - df = _make_hlc_panel( + # 这里通过 mock get_closes_panel 给所有候选相同 close,看最终名单 + panel = _make_close_wide( ["A.XSHG", "D.XSHG"], [[10.0] * 130, [10.0] * 130], end_date="2024-09-30", days=130, ) - # _pick_stocks 的 get_price 入参可能是 list 形式 - s.provider.get_price.side_effect = None - s.provider.get_price.return_value = df + s.provider.get_closes_panel.side_effect = None + s.provider.get_closes_panel.return_value = panel ctx = FakeContext(current_dt=datetime(2024, 10, 8, 9, 30)) out = s._pick_stocks(ctx) # eps>0 的 A/D 都进入候选,B/C 被剔 @@ -350,13 +342,13 @@ class TestPickStocks: fundamentals_df=fund, config=cfg, ) - df = _make_hlc_panel( + panel = _make_close_wide( ["SMALL.XSHG", "MID.XSHG"], [[10.0] * 130, [10.0] * 130], end_date="2024-09-30", days=130, ) - s.provider.get_price.side_effect = None - s.provider.get_price.return_value = df + s.provider.get_closes_panel.side_effect = None + s.provider.get_closes_panel.return_value = panel ctx = FakeContext(current_dt=datetime(2024, 10, 8, 9, 30)) out = s._pick_stocks(ctx) # market_cap 升序后前 2 只 = SMALL/MID,BIG 被剔 @@ -377,9 +369,9 @@ class TestPickStocks: ) # 所有股票 close 相同 → score 相同 → 顺序由 sort_values stable 决定 closes = [[10.0] * 130 for _ in stocks] - df = _make_hlc_panel(stocks, closes, end_date="2024-09-30", days=130) - s.provider.get_price.side_effect = None - s.provider.get_price.return_value = df + panel = _make_close_wide(stocks, closes, end_date="2024-09-30", days=130) + s.provider.get_closes_panel.side_effect = None + s.provider.get_closes_panel.return_value = panel ctx = FakeContext(current_dt=datetime(2024, 10, 8, 9, 30)) out = s._pick_stocks(ctx) assert len(out) == 20 @@ -404,12 +396,12 @@ class TestPickStocks: flat = [10.0] * 130 rising = [10.0 + i * 0.1 for i in range(130)] falling = [23.0 - i * 0.1 for i in range(130)] - df = _make_hlc_panel( + panel = _make_close_wide( ["DOWN.XSHG", "FLAT.XSHG", "UP.XSHG"], [falling, flat, rising], end_date="2024-09-30", days=130, ) - s.provider.get_price.side_effect = None - s.provider.get_price.return_value = df + s.provider.get_closes_panel.side_effect = None + s.provider.get_closes_panel.return_value = panel ctx = FakeContext(current_dt=datetime(2024, 10, 8, 9, 30)) out = s._pick_stocks(ctx) # 顺序:DOWN(score 最负) → FLAT(0),UP 被剔 @@ -477,9 +469,9 @@ class TestHandleDataRebalance: fundamentals_df=fund, config=cfg, ) - df = _make_hlc_panel(["NEW.XSHG"], [[10.0] * 130], end_date="2024-09-30", days=130) - s.provider.get_price.side_effect = None - s.provider.get_price.return_value = df + df = _make_close_wide(["NEW.XSHG"], [[10.0] * 130], end_date="2024-09-30", days=130) + s.provider.get_closes_panel.side_effect = None + s.provider.get_closes_panel.return_value = df ctx = FakeContext( current_dt=datetime(2024, 10, 8, 9, 30), positions={ @@ -508,12 +500,12 @@ class TestHandleDataRebalance: fundamentals_df=fund, config=cfg, ) - df = _make_hlc_panel( + df = _make_close_wide( ["A.XSHG", "B.XSHG"], [[10.0] * 130, [10.0] * 130], end_date="2024-09-30", days=130, ) - s.provider.get_price.side_effect = None - s.provider.get_price.return_value = df + s.provider.get_closes_panel.side_effect = None + s.provider.get_closes_panel.return_value = df ctx = FakeContext( current_dt=datetime(2024, 10, 8, 9, 30), positions={}, cash=1_000_000,