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]
This commit is contained in:
@@ -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"]
|
||||
|
||||
@@ -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"]
|
||||
|
||||
Reference in New Issue
Block a user