diff --git a/sanguo_portfolio/strategies/all_weather.py b/sanguo_portfolio/strategies/all_weather.py index 8d66f56..4acdfa1 100644 --- a/sanguo_portfolio/strategies/all_weather.py +++ b/sanguo_portfolio/strategies/all_weather.py @@ -141,26 +141,18 @@ class AllWeatherStrategy: previous_date = _previous_date_str(context) if not previous_date: return - try: - df = self.provider.get_price( - self.hold_list, - end_date=previous_date, - frequency="daily", - fields=["close", "high_limit"], - count=1, - panel=False, - fill_paused=False, - ) - except Exception as exc: - logger.warning("prepare_stock_list get_price 失败: %s", exc) - return - if df is None or len(df) == 0: + fn = getattr(self.provider, "get_limit_status_batch", None) + if fn is None: return try: - hit = df[df["close"] == df["high_limit"]] - self.yesterday_hl_list = list(hit.get("code", [])) + status_map = fn(self.hold_list, previous_date) or {} except Exception as exc: - logger.debug("prepare_stock_list 解析涨停失败: %s", exc) + logger.warning("prepare_stock_list get_limit_status_batch 失败: %s", exc) + return + self.yesterday_hl_list = [ + s for s in self.hold_list + if (status_map.get(s) or {}).get("is_limit_up") + ] # =================== stop_loss =================== def stop_loss(self, context: Any) -> None: @@ -375,18 +367,26 @@ class AllWeatherStrategy: """聚宽 filter_roic:只保留 ROIC > 0.08。 聚宽原版用 ``get_factor_values(stock, 'roic_ttm')``,这里走我们自算的 roic - (provider 已合并到 df['roic'])。逐只查比 df 慢,但保留原签名便于实证对账。 + (provider 已合并到 df['roic'])。批量 + ``fields=['roic']`` 按需短路 + (只读 income/balance 源表,fields 短路 9 倍速,见 provider f416a17)。 """ + if not stock_list: + return [] threshold = self.config.roic_threshold - out: List[str] = [] - for stock in stock_list: - df = self.provider.get_fundamentals_df([stock], date=previous_date) - if df.empty: - continue - roic = float(df["roic"].iloc[0]) if "roic" in df.columns else float("nan") - if roic == roic and roic > threshold: - out.append(stock) - return out + try: + df = self.provider.get_fundamentals_df( + stock_list, date=previous_date, fields=["roic"], + ) + except Exception as exc: + logger.warning("filter_roic get_fundamentals_df 失败: %s", exc) + return [] + if df is None or df.empty or "roic" not in df.columns: + return [] + roic = df["roic"] + return [ + s for s in stock_list + if s in roic.index and roic[s] == roic[s] and roic[s] > threshold + ] # =================== 调仓辅助 =================== def _close_position(self, code: str) -> bool: @@ -430,30 +430,28 @@ class AllWeatherStrategy: return list(df.index)[:n] def _trend_mean(self, stocks: List[str], end_date: str, n: int) -> float: - """N 日涨幅(% )的均值。聚宽原版取 close 涨幅,这里同样。""" + """N 日涨幅(% )的均值。聚宽原版取 close 涨幅,这里同样。 + + 批量取数:``get_closes_panel`` 宽表(UNION ALL 走复合索引 340×,只取 close) + 替 ``get_price`` 长表+pivot;``count=n`` → ``start=end-n*2 自然日`` + ``.tail(n)`` + 切片(避免自然日 vs 交易日换算歧义,同 momentum_timing 口径)。fq 均为 raw 不变。 + """ if not stocks: return 0.0 + start_date = _shift_date(end_date, -n * 2) try: - df = self.provider.get_price( - stocks, - end_date=end_date, - frequency="1d", - fields=["close"], - count=n, - panel=False, + panel = self.provider.get_closes_panel( + stocks, start_date, end_date, fq="raw", ) except Exception as exc: - logger.warning("get_price trend 失败: %s", exc) + logger.warning("get_closes_panel trend 失败: %s", exc) return 0.0 - if df is None or df.empty: + if panel is None or panel.empty: return 0.0 - try: - pivot = df.pivot(index="time", columns="code", values="close") - except Exception: + panel = panel.tail(n) + if len(panel) < 2: return 0.0 - if len(pivot) < 2: - return 0.0 - change = (pivot.iloc[-1] / pivot.iloc[0] - 1) * 100 + change = (panel.iloc[-1] / panel.iloc[0] - 1) * 100 arr = np.nan_to_num(change.to_numpy()) return float(np.mean(arr)) @@ -578,4 +576,16 @@ def _dedup(seq: Sequence[str]) -> List[str]: return out +def _shift_date(date_str: str, days: int) -> str: + """字符串日期加减自然日,返回 YYYY-MM-DD(供 count=N → start=end-N*2 换算)。 + + 同 momentum_timing._shift_date 口径;本地定义避免互相 import 循环。 + """ + 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__ = ["AllWeatherStrategy", "AllWeatherConfig", "BrokerFacade"] diff --git a/tests/portfolio/test_all_weather.py b/tests/portfolio/test_all_weather.py index 09d5f16..0f77489 100644 --- a/tests/portfolio/test_all_weather.py +++ b/tests/portfolio/test_all_weather.py @@ -68,6 +68,30 @@ def make_strategy( provider.get_price.side_effect = _get_price + # get_closes_panel:从 price_df_map 同源构造 close 宽表(忽略 count 差异) + def _get_closes_panel(symbols, start=None, end=None, interval="d", fq="raw"): + frames = [ + df[["time", "code", "close"]] + for (sec, fields, _cnt), df in price_df_map.items() + if sec == str(symbols) and tuple(fields) == ("close",) + and df is not None and len(df) + ] + if not frames: + return pd.DataFrame() + long_df = pd.concat(frames) + return long_df.pivot(index="time", columns="code", values="close") + + provider.get_closes_panel.side_effect = _get_closes_panel + + # 涨跌停/停牌批量:默认全正常(filter_paused 无bar=剔除语义,MagicMock/{} 会误杀) + def _get_limit_status_batch(codes, date=None): + return { + c: {"is_limit_up": False, "is_limit_down": False, "is_paused": False} + for c in codes + } + + provider.get_limit_status_batch.side_effect = _get_limit_status_batch + broker = BrokerFacade() broker.order_target_value = MagicMock(return_value=MagicMock(filled=100)) broker.order_value = MagicMock(return_value=MagicMock(filled=100)) @@ -136,23 +160,69 @@ class TestPrepareStockList: def test_records_yesterday_limit_up(self): s = make_strategy() - # 清掉 make_strategy 设置的 side_effect,直接用 return_value - s.provider.get_price.side_effect = None pos = MagicMock(); pos.security = "600519.XSHG" ctx = MagicMock() ctx.portfolio.positions = {"600519.XSHG": pos} ctx.previous_date = "2024-09-30" - # close == high_limit 视为涨停 - s.provider.get_price.return_value = pd.DataFrame({ - "code": ["600519.XSHG"], - "close": [10.0], - "high_limit": [10.0], - }) + # get_limit_status_batch 返回昨日涨停(G5-P2:原 get_price close==high_limit 改批量口径) + s.provider.get_limit_status_batch.side_effect = \ + lambda codes, date=None: { + c: {"is_limit_up": True, "is_limit_down": False, "is_paused": False} + for c in codes + } s.prepare_stock_list(ctx) assert "600519.XSHG" in s.yesterday_hl_list + def test_uses_batch_endpoint_once(self): + """G5-P2:原 get_price(close+high_limit) → get_limit_status_batch 一次批量。""" + s = make_strategy() + pos = MagicMock(); pos.security = "600519.XSHG" + ctx = MagicMock() + ctx.portfolio.positions = {"600519.XSHG": pos, "000001.XSHE": MagicMock(security="000001.XSHE")} + ctx.previous_date = "2024-09-30" + s.prepare_stock_list(ctx) + assert s.provider.get_limit_status_batch.call_count == 1 + args, _kwargs = s.provider.get_limit_status_batch.call_args + assert sorted(args[0]) == ["000001.XSHE", "600519.XSHG"] + s.provider.get_price.assert_not_called() -# =================== stop_loss =================== + +# =================== trend_mean =================== +class TestTrendMean: + def _panel_seed(self): + return pd.DataFrame({ + "time": pd.to_datetime(["2024-09-20", "2024-09-30"]), + "code": ["600519.XSHG"] * 2, + "close": [10.0, 15.0], # 涨 50% + }) + + def test_uses_closes_panel_not_get_price(self): + """G5-P2:原 get_price 长表+pivot → get_closes_panel 宽表。""" + s = make_strategy(price_df_map={ + ("['600519.XSHG']", ("close",), 10): self._panel_seed(), + }) + out = s._trend_mean(["600519.XSHG"], "2024-09-30", 10) + assert out == 50.0 + s.provider.get_closes_panel.assert_called_once() + s.provider.get_price.assert_not_called() + + def test_empty_or_short_panel_returns_zero(self): + """panel 缺失/不足 2 行 → 0.0 不崩(降级同原 get_price 空df语义)。""" + s = make_strategy() + assert s._trend_mean([], "2024-09-30", 10) == 0.0 + assert s._trend_mean(["600519.XSHG"], "2024-09-30", 10) == 0.0 # 无种子数据 + + def test_nan_only_change_not_propagated(self): + """某列 NaN(close 缺失)→ nan_to_num 归 0,不产出 NaN。""" + wide = pd.DataFrame( + {"600519.XSHG": [np.nan, 15.0], "000001.XSHE": [10.0, 11.0]}, + index=pd.to_datetime(["2024-09-20", "2024-09-30"]), + ) + s = make_strategy() + s.provider.get_closes_panel.side_effect = None + s.provider.get_closes_panel.return_value = wide + # 600519 NaN→0%, 000001 +10% → 均值 5(浮点容差) + assert s._trend_mean(["600519.XSHG", "000001.XSHE"], "2024-09-30", 10) == pytest.approx(5.0) class TestStopLoss: def test_stop_loss_triggers_when_price_drops_8pct(self): """avg_cost=100, price=91 (< 100*0.92=92) → 止损。""" @@ -467,3 +537,15 @@ class TestFilterRoic: s = make_strategy() out = s.filter_roic([], previous_date="2024-09-30") assert out == [] + + def test_batch_single_call_with_fields_shortcut(self): + """G5-P2:原逐只循环 → 一次批量 + fields=['roic'] 短路。""" + df = make_fund_df([{"code": "A.XSHG", "roic": 0.15}]) + s = make_strategy() + s.provider.get_fundamentals_df.return_value = df + s.filter_roic(["A.XSHG", "B.XSHG"], previous_date="2024-09-30") + # 一次批量调用(非逐只 N 次),且带 fields 短路 + assert s.provider.get_fundamentals_df.call_count == 1 + args, kwargs = s.provider.get_fundamentals_df.call_args + assert list(args[0]) == ["A.XSHG", "B.XSHG"] + assert kwargs.get("fields") == ["roic"]