perf(portfolio): G5-P2收尾—all_weather三调用点向量化:①filter_roic逐只循环→一次批量+fields=['roic']短路(9×)②_trend_mean get_price长表+pivot→get_closes_panel宽表(340×,count=N→start-N*2自然日+.tail(N)同momentum口径,fq均raw)③prepare_stock_list get_price(close+high_limit)→get_limit_status_batch精确涨跌停口径;测试fixture补panel/limit_batch mock(原MagicMock碰巧truthy蒙混paused剔除语义,默认改全正常);+6回归测试,portfolio全套302绿 [vps]
CI/CD / test (push) Successful in 13s
CI/CD / nas-deploy (push) Successful in 39s
CI/CD / nas-verify (push) Successful in 14s

This commit is contained in:
2026-08-14 23:17:13 +08:00
parent e9aec041db
commit 2cd158286b
2 changed files with 144 additions and 52 deletions
+53 -43
View File
@@ -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"]
+91 -9
View File
@@ -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"]