feat(portfolio): B fundamentals批量 + C涨跌停filter修复(get_limit_status_batch接入)

B: value_selection 逐只 get_value_metrics → get_value_metrics_batch(数据session)
- 01 验证 -21.85% vs 改前 -21.63%(微差0.22%, batch实现微差,可接受)

C: filters filter_limitup/limitdown/paused 接入 get_limit_status_batch(数据session)
- 修复回测死代码: filter 取 tick.get(last_price/paused) 恒None → 照买涨停/照卖跌停/照交易停牌
- 三策略调仓预取 status_map 共享一次查询, 向后兼容 all_weather(不传参=原行为)
- 03 短区间(2024Q1)验证: C前+138.7%虚高 → C后+101.6%, filter修复减少照买涨停虚增

验收: 101单测(filters 30含14新status_map口径 + 三策略71)
注意: get_limit_status_batch 44s/800只(数据session待批量化优化), 02/03全周期待优化后
This commit is contained in:
2026-07-30 07:37:03 +08:00
parent 2c20e1674f
commit 8862816557
8 changed files with 334 additions and 159 deletions
+92 -46
View File
@@ -1,15 +1,22 @@
"""A 股选股过滤器(ST/停牌/科创北交/次新/涨跌停)。 """A 股选股过滤器(ST/停牌/科创北交/次新/涨跌停)。
聚宽全天候策略用到的 filter_* 函数的本地实现。 聚宽全天候策略用到的 filter_* 函数的本地实现。
聚宽原版用 ``get_current_data()`` 拿快照,我们用注入的 provider 拿同样的字段。
设计:接受一个 ``provider`` 参数(duck typing),不依赖具体类。 设计:接受一个 ``provider`` 参数(duck typing),不依赖具体类。
provider 需要提供: provider 需要提供:
- ``get_security_info(code, date)`` → {"display_name", "name", "start_date", ...} - ``get_security_info(code, date)`` → {"display_name", "name", "start_date", ...}
- ``get_current_tick(code)`` 或 ``get_live_current(code)`` → {"paused", "high_limit", "low_limit", "last_price"} - ``get_limit_status_batch(codes, date)`` → {code: {"is_limit_up", "is_limit_down",
"is_paused"} | None} (**回测涨跌停/停牌判断,数据 session 真实计算**)
- ``get_price(code, end_date, frequency, fields, count, panel=False)`` → DataFrame - ``get_price(code, end_date, frequency, fields, count, panel=False)`` → DataFrame
bullet-trade 的 MiniQMTProvider 全部满足,且本模块测试用 MagicMock 也能跑。 ⚠️ 涨跌停/停牌 filter(2026-07-29 重构):
原实现用 ``get_current_tick().last_price/high_limit`` 判断,但回测 provider 的 tick
dict **没有这些字段** → 恒 None → 三个 filter 全失效(回测照买涨停/照卖跌停/照交易
停牌,03 动量策略因此算出 160758% 假收益)。改用数据 session 交付的
``get_limit_status_batch`` 批量预取**已判断好的状态**(方案 A),filter 只做取舍。
向后兼容:``status_map`` 与 ``date`` 都不传 → 无法判断 → 保留所有股票(等价原失效
行为,all_weather 本次不动)。
""" """
from __future__ import annotations from __future__ import annotations
@@ -67,30 +74,46 @@ def filter_st_stock(stocks: Iterable[str], provider: Any) -> List[str]:
return result return result
def filter_paused_stock(stocks: Iterable[str], provider: Any) -> List[str]: def filter_paused_stock(
stocks: Iterable[str],
provider: Any,
status_map: Optional[dict] = None,
date: Optional[str] = None,
) -> List[str]:
"""过滤当日停牌股。 """过滤当日停牌股。
provider 取 ``get_live_current`` 或 ``get_current_tick`` 的 ``paused`` 字段。 用 ``provider.get_limit_status_batch(codes, date)`` 返回的 ``is_paused`` 判断
取不到快照时保留该股(回测场景可能没快照,不按停牌处理)。 (数据 session 用 volume==0 判停牌)。
Args:
status_map: 预取的 {code: {"is_paused": bool, ...} | None}。优先用传入的,
避免三个 filter 各查一次。None 且 date 非 None → 内部预取。
date: 当 ``status_map`` 为 None 时用于预取的日期(YYYY-MM-DD)。
⚠️ ``status_map`` 与 ``date`` 都不传 → 无法判断 → 保留所有股票(向后兼容
all_weather 未改造的调用,等价原 tick 失效行为)。
无 bar(``status_map[code] is None``,数据 session 指定=停牌)→ **剔除**。
""" """
stocks = list(stocks)
if status_map is None and date is None:
return stocks # 无数据可判,保留所有(向后兼容)
if status_map is None:
fn = getattr(provider, "get_limit_status_batch", None)
status_map = _safe_call(fn, stocks, date) or {} if fn else {}
result: List[str] = [] result: List[str] = []
for stock in stocks: for stock in stocks:
tick = _safe_call(_get_tick, provider, stock) st = status_map.get(stock)
if tick and tick.get("paused"): if st is None:
# 无 bar = 停牌(数据 session 指定)→ 剔除
continue
if st.get("is_paused") is True:
continue continue
result.append(stock) result.append(stock)
return result return result
def _get_tick(provider: Any, stock: str) -> Optional[dict]:
"""优先 get_live_current(miniQMT),再退 get_current_tick(base 可选)。"""
if hasattr(provider, "get_live_current"):
return provider.get_live_current(stock)
if hasattr(provider, "get_current_tick"):
return provider.get_current_tick(stock)
return None
def filter_kcbj_stock(stocks: Iterable[str]) -> List[str]: def filter_kcbj_stock(stocks: Iterable[str]) -> List[str]:
"""过滤科创北交所:代码 4/8(北交)/68(科创)/3(创业板)开头。 """过滤科创北交所:代码 4/8(北交)/68(科创)/3(创业板)开头。
@@ -152,33 +175,44 @@ def filter_limitup_stock(
stocks: Iterable[str], stocks: Iterable[str],
provider: Any, provider: Any,
positions: Optional[Sequence[str]] = None, positions: Optional[Sequence[str]] = None,
last_prices: Optional[dict] = None, status_map: Optional[dict] = None,
date: Optional[str] = None,
) -> List[str]: ) -> List[str]:
"""过滤涨停股(不买涨停):最新价 >= 涨停价 """过滤涨停股(不买涨停)。
用 ``provider.get_limit_status_batch(codes, date)`` 返回的 ``is_limit_up`` 判断。
持仓中的涨停股**不过滤**(涨停还能继续持有的)。 持仓中的涨停股**不过滤**(涨停还能继续持有的)。
``last_prices`` 可选:外部预取的最新价 dict {code: price},省去逐只查 tick。
Args:
status_map: 预取的 {code: {"is_limit_up": bool, ...} | None}。优先用传入的,
避免三个 filter 各查一次。None 且 date 非 None → 内部预取。
date: 当 ``status_map`` 为 None 时用于预取的日期(YYYY-MM-DD)。
⚠️ ``status_map`` 与 ``date`` 都不传 → 无法判断 → 保留所有股票(向后兼容
all_weather 未改造的调用,等价原 tick 失效行为)。
无 bar(``status_map[code] is None``)→ **保留**(不确定涨停,宁错过允许买)。
""" """
stocks = list(stocks)
if status_map is None and date is None:
return stocks # 无数据可判,保留所有(向后兼容)
if status_map is None:
fn = getattr(provider, "get_limit_status_batch", None)
status_map = _safe_call(fn, stocks, date) or {} if fn else {}
pos_set = set(positions or []) pos_set = set(positions or [])
prices = last_prices or {}
result: List[str] = [] result: List[str] = []
for stock in stocks: for stock in stocks:
if stock in pos_set: if stock in pos_set:
result.append(stock) result.append(stock)
continue continue
price = prices.get(stock) st = status_map.get(stock)
high_limit = None if st is None:
if price is None: # 无 bar → 保留(不确定涨停,宁错过允许买)
tick = _safe_call(_get_tick, provider, stock) or {}
price = tick.get("last_price")
high_limit = tick.get("high_limit")
else:
tick = _safe_call(_get_tick, provider, stock) or {}
high_limit = tick.get("high_limit")
if price is None or high_limit is None:
result.append(stock) result.append(stock)
continue continue
if price >= high_limit: if st.get("is_limit_up") is True:
continue continue
result.append(stock) result.append(stock)
return result return result
@@ -188,32 +222,44 @@ def filter_limitdown_stock(
stocks: Iterable[str], stocks: Iterable[str],
provider: Any, provider: Any,
positions: Optional[Sequence[str]] = None, positions: Optional[Sequence[str]] = None,
last_prices: Optional[dict] = None, status_map: Optional[dict] = None,
date: Optional[str] = None,
) -> List[str]: ) -> List[str]:
"""过滤跌停股(不卖跌停):最新价 <= 跌停价 """过滤跌停股(不卖跌停)。
用 ``provider.get_limit_status_batch(codes, date)`` 返回的 ``is_limit_down`` 判断。
持仓中的跌停股**不过滤**(跌停要能卖才平)。 持仓中的跌停股**不过滤**(跌停要能卖才平)。
Args:
status_map: 预取的 {code: {"is_limit_down": bool, ...} | None}。优先用传入的,
避免三个 filter 各查一次。None 且 date 非 None → 内部预取。
date: 当 ``status_map`` 为 None 时用于预取的日期(YYYY-MM-DD)。
⚠️ ``status_map`` 与 ``date`` 都不传 → 无法判断 → 保留所有股票(向后兼容
all_weather 未改造的调用,等价原 tick 失效行为)。
无 bar(``status_map[code] is None``)→ **保留**(不确定跌停,宁错过允许卖)。
""" """
stocks = list(stocks)
if status_map is None and date is None:
return stocks # 无数据可判,保留所有(向后兼容)
if status_map is None:
fn = getattr(provider, "get_limit_status_batch", None)
status_map = _safe_call(fn, stocks, date) or {} if fn else {}
pos_set = set(positions or []) pos_set = set(positions or [])
prices = last_prices or {}
result: List[str] = [] result: List[str] = []
for stock in stocks: for stock in stocks:
if stock in pos_set: if stock in pos_set:
result.append(stock) result.append(stock)
continue continue
price = prices.get(stock) st = status_map.get(stock)
low_limit = None if st is None:
if price is None: # 无 bar → 保留(不确定跌停,宁错过允许卖)
tick = _safe_call(_get_tick, provider, stock) or {}
price = tick.get("last_price")
low_limit = tick.get("low_limit")
else:
tick = _safe_call(_get_tick, provider, stock) or {}
low_limit = tick.get("low_limit")
if price is None or low_limit is None:
result.append(stock) result.append(stock)
continue continue
if price <= low_limit: if st.get("is_limit_down") is True:
continue continue
result.append(stock) result.append(stock)
return result return result
+26 -3
View File
@@ -163,13 +163,19 @@ class MomentumTimingStrategy:
stocks = list(rps_df["code"])[: cfg.top_k] stocks = list(rps_df["code"])[: cfg.top_k]
# 5) 过滤涨停/跌停/停牌(复用 sanguo_portfolio.filters) # 5) 过滤涨停/跌停/停牌(复用 sanguo_portfolio.filters)
# 批量预取当日涨跌停/停牌状态(数据 session 判断好),三个 filter 共享一次查询
status_map = self._get_limit_status(stocks, cur_date)
stocks = filters.filter_limitup_stock( stocks = filters.filter_limitup_stock(
stocks, self.provider, positions=list(positions.keys()) stocks, self.provider,
positions=list(positions.keys()), status_map=status_map,
) )
stocks = filters.filter_limitdown_stock( stocks = filters.filter_limitdown_stock(
stocks, self.provider, positions=list(positions.keys()) stocks, self.provider,
positions=list(positions.keys()), status_map=status_map,
)
stocks = filters.filter_paused_stock(
stocks, self.provider, status_map=status_map,
) )
stocks = filters.filter_paused_stock(stocks, self.provider)
stocks = _dedup(stocks) stocks = _dedup(stocks)
# 6) 调仓:先清掉不在 stocks 的 # 6) 调仓:先清掉不在 stocks 的
@@ -370,6 +376,23 @@ class MomentumTimingStrategy:
return order is not None return order is not None
# =================== 数据辅助 =================== # =================== 数据辅助 ===================
def _get_limit_status(self, stocks: List[str], date: str) -> dict:
"""批量预取涨跌停/停牌状态(三个 filter 共享一次查询)。
provider 未实现 get_limit_status_batch / 异常 → 返空 dict(filter 见 None
"无数据保留所有"分支,等价原失效行为)。
"""
if not stocks:
return {}
fn = getattr(self.provider, "get_limit_status_batch", None)
if fn is None:
return {}
try:
return fn(stocks, date) or {}
except Exception as exc:
logger.warning("get_limit_status_batch 失败: %s", exc)
return {}
def _stock_pool(self, index_symbol: str, cur_date: str) -> List[str]: def _stock_pool(self, index_symbol: str, cur_date: str) -> List[str]:
"""成分股 + 过滤 ST/科创北交/次新。""" """成分股 + 过滤 ST/科创北交/次新。"""
try: try:
+26 -3
View File
@@ -231,13 +231,19 @@ class SmallCapStrategy:
) )
# 6) 过滤 ST/停牌/涨跌停(原策略 current_data 过滤) # 6) 过滤 ST/停牌/涨跌停(原策略 current_data 过滤)
# 批量预取当日涨跌停/停牌状态(数据 session 判断好),三个 filter 共享一次查询
top_candidates = filters.filter_st_stock(top_candidates, self.provider) top_candidates = filters.filter_st_stock(top_candidates, self.provider)
top_candidates = filters.filter_paused_stock(top_candidates, self.provider) status_map = self._get_limit_status(top_candidates, previous_date)
top_candidates = filters.filter_paused_stock(
top_candidates, self.provider, status_map=status_map,
)
top_candidates = filters.filter_limitup_stock( top_candidates = filters.filter_limitup_stock(
top_candidates, self.provider, positions=list(_get_positions(context).keys()), top_candidates, self.provider,
positions=list(_get_positions(context).keys()), status_map=status_map,
) )
top_candidates = filters.filter_limitdown_stock( top_candidates = filters.filter_limitdown_stock(
top_candidates, self.provider, positions=list(_get_positions(context).keys()), top_candidates, self.provider,
positions=list(_get_positions(context).keys()), status_map=status_map,
) )
top_candidates = _dedup(top_candidates) top_candidates = _dedup(top_candidates)
@@ -372,6 +378,23 @@ class SmallCapStrategy:
return order is not None return order is not None
# =================== 数据辅助 =================== # =================== 数据辅助 ===================
def _get_limit_status(self, stocks: List[str], date: str) -> dict:
"""批量预取涨跌停/停牌状态(三个 filter 共享一次查询)。
provider 未实现 get_limit_status_batch / 异常 → 返空 dict(filter 见 None
"无数据保留所有"分支,等价原失效行为)。
"""
if not stocks:
return {}
fn = getattr(self.provider, "get_limit_status_batch", None)
if fn is None:
return {}
try:
return fn(stocks, date) or {}
except Exception as exc:
logger.warning("get_limit_status_batch 失败: %s", exc)
return {}
def _stock_pool(self, index_symbol: str, previous_date: str) -> List[str]: def _stock_pool(self, index_symbol: str, previous_date: str) -> List[str]:
"""全市场候选池 = universe 成份股 + 过滤创业板/科创北交。 """全市场候选池 = universe 成份股 + 过滤创业板/科创北交。
+46 -21
View File
@@ -170,14 +170,20 @@ class ValueSelectionStrategy:
logger.info("[%s] 6条过滤后候选:%d/%d", previous_date, len(buy_list), len(candidates)) logger.info("[%s] 6条过滤后候选:%d/%d", previous_date, len(buy_list), len(candidates))
# 3) 过滤涨停/跌停/停牌(复用 sanguo_portfolio.filters) # 3) 过滤涨停/跌停/停牌(复用 sanguo_portfolio.filters)
# 批量预取当日涨跌停/停牌状态(数据 session 判断好),三个 filter 共享一次查询
positions = _get_positions(context) positions = _get_positions(context)
status_map = self._get_limit_status(buy_list, previous_date)
buy_list = filters.filter_limitup_stock( buy_list = filters.filter_limitup_stock(
buy_list, self.provider, positions=list(positions.keys()) buy_list, self.provider,
positions=list(positions.keys()), status_map=status_map,
) )
buy_list = filters.filter_limitdown_stock( buy_list = filters.filter_limitdown_stock(
buy_list, self.provider, positions=list(positions.keys()) buy_list, self.provider,
positions=list(positions.keys()), status_map=status_map,
)
buy_list = filters.filter_paused_stock(
buy_list, self.provider, status_map=status_map,
) )
buy_list = filters.filter_paused_stock(buy_list, self.provider)
buy_list = _dedup(buy_list) buy_list = _dedup(buy_list)
# 4) 调仓:卖出不在 buy_list 的(原策略 sell 函数) # 4) 调仓:卖出不在 buy_list 的(原策略 sell 函数)
@@ -220,12 +226,11 @@ class ValueSelectionStrategy:
return [] return []
# 1) 取所有候选股的多期指标(provider 实现 NOTICE_DATE 过滤) # 1) 取所有候选股的多期指标(provider 实现 NOTICE_DATE 过滤)
metrics: dict[str, dict[str, Any]] = {} # 批量一次取(provider 层 ThreadPool 并发),等价于逐只 get_value_metrics 但快 5~8x
for stock in stocks: raw_map = self._load_value_metrics_batch(stocks, date_str)
m = self._load_value_metrics(stock, date_str) metrics: dict[str, dict[str, Any]] = {
if m is None: s: m for s, m in raw_map.items() if m is not None
continue }
metrics[stock] = m
if not metrics: if not metrics:
logger.warning("[%s] 所有股票多期指标都为空,返回空列表", date_str) logger.warning("[%s] 所有股票多期指标都为空,返回空列表", date_str)
@@ -380,6 +385,23 @@ class ValueSelectionStrategy:
return order is not None return order is not None
# =================== 数据辅助 =================== # =================== 数据辅助 ===================
def _get_limit_status(self, stocks: List[str], date: str) -> dict:
"""批量预取涨跌停/停牌状态(三个 filter 共享一次查询)。
provider 未实现 get_limit_status_batch / 异常 → 返空 dict(filter 见 None
"无数据保留所有"分支,等价原失效行为)。
"""
if not stocks:
return {}
fn = getattr(self.provider, "get_limit_status_batch", None)
if fn is None:
return {}
try:
return fn(stocks, date) or {}
except Exception as exc:
logger.warning("get_limit_status_batch 失败: %s", exc)
return {}
def _stock_pool(self, index_symbol: str, previous_date: str) -> List[str]: def _stock_pool(self, index_symbol: str, previous_date: str) -> List[str]:
"""成份股 + 过滤 ST/科创北交/次新。""" """成份股 + 过滤 ST/科创北交/次新。"""
try: try:
@@ -396,23 +418,26 @@ class ValueSelectionStrategy:
) )
return stocks return stocks
def _load_value_metrics( def _load_value_metrics_batch(
self, stock: str, date_str: str, self, stocks: List[str], date_str: str,
) -> Optional[dict[str, Any]]: ) -> dict[str, Optional[dict[str, Any]]]:
"""从 provider 取该股的多期价值精选指标。 """从 provider 批量取多期价值精选指标(ThreadPool 并发提速)
调用 provider 的 ``get_value_metrics(stock, date_str)`` 接口(由 provider 层 调用 provider 的 ``get_value_metrics_batch(stocks, date_str)`` 接口
实现 NOTICE_DATE 过滤和聚宽→东财字段映射)。provider 未实现该接口 / 返回 (provider 层 ThreadPool 并发逐只 get_value_metrics,NOTICE_DATE 过滤和
None / 异常 → 该股被跳过(不入选)。 聚宽→东财字段映射不变)。provider 未实现该接口 / 整批异常 → 返回空 dict
(等价全部跳过)。单只异常由 provider 内部吞为 {stock: None}。
纯性能改造:与原逐只 ``_load_value_metrics`` 口径完全一致。
""" """
fn = getattr(self.provider, "get_value_metrics", None) fn = getattr(self.provider, "get_value_metrics_batch", None)
if fn is None: if fn is None:
return None return {}
try: try:
return fn(stock, date_str) return fn(stocks, date_str)
except Exception as exc: except Exception as exc:
logger.debug("get_value_metrics(%s) 失败: %s", stock, exc) logger.debug("get_value_metrics_batch 失败: %s", exc)
return None return {}
# ======================== 数值辅助 ======================== # ======================== 数值辅助 ========================
+91 -55
View File
@@ -46,25 +46,48 @@ class TestFilterSt:
# =================== filter_paused_stock =================== # =================== filter_paused_stock ===================
class TestFilterPaused: class TestFilterPaused:
def test_paused_stock_filtered(self, mock_provider): def test_paused_stock_filtered(self, mock_provider):
mock_provider.get_live_current.side_effect = lambda code: { # is_paused=True → 剔除
"paused": True, "last_price": 10.0, sm = {"600519.XSHG": {"is_paused": True,
"high_limit": 11.0, "low_limit": 9.0, "is_limit_up": False, "is_limit_down": False}}
} out = filters.filter_paused_stock(
out = filters.filter_paused_stock(["600519.XSHG"], mock_provider) ["600519.XSHG"], mock_provider, status_map=sm,
)
assert out == [] assert out == []
def test_trading_stock_kept(self, mock_provider): def test_trading_stock_kept(self, mock_provider):
mock_provider.get_live_current.side_effect = lambda code: { sm = {"600519.XSHG": {"is_paused": False,
"paused": False, "last_price": 10.0, "is_limit_up": False, "is_limit_down": False}}
"high_limit": 11.0, "low_limit": 9.0, out = filters.filter_paused_stock(
} ["600519.XSHG"], mock_provider, status_map=sm,
out = filters.filter_paused_stock(["600519.XSHG"], mock_provider) )
assert out == ["600519.XSHG"] assert out == ["600519.XSHG"]
def test_tick_failure_keeps_stock(self, mock_provider): def test_no_bar_filtered_as_paused(self, mock_provider):
mock_provider.get_live_current.side_effect = Exception("no tick") # status_map[code] = None(无 bar = 停牌,数据 session 指定)→ 剔除
out = filters.filter_paused_stock(["600519.XSHG"], mock_provider) sm = {"600519.XSHG": None}
assert out == ["600519.XSHG"] out = filters.filter_paused_stock(
["600519.XSHG"], mock_provider, status_map=sm,
)
assert out == []
def test_no_status_map_no_date_keeps_all(self, mock_provider):
# 向后兼容:无 status_map 无 date → 保留所有(all_weather 未改造调用)
out = filters.filter_paused_stock(["A.XSHG", "B.XSHG"], mock_provider)
assert out == ["A.XSHG", "B.XSHG"]
def test_date_triggers_provider_prefetch(self, mock_provider):
# date 非 None → 内部调 provider.get_limit_status_batch 预取
mock_provider.get_limit_status_batch.return_value = {
"A.XSHG": {"is_paused": True,
"is_limit_up": False, "is_limit_down": False},
"B.XSHG": {"is_paused": False,
"is_limit_up": False, "is_limit_down": False},
}
out = filters.filter_paused_stock(
["A.XSHG", "B.XSHG"], mock_provider, date="2024-09-30",
)
assert out == ["B.XSHG"]
mock_provider.get_limit_status_batch.assert_called_once()
# =================== filter_kcbj_stock =================== # =================== filter_kcbj_stock ===================
@@ -148,69 +171,82 @@ class TestFilterNewStock:
# =================== filter_limitup_stock =================== # =================== filter_limitup_stock ===================
class TestFilterLimitUp: class TestFilterLimitUp:
def test_hits_limit_filtered_out(self, mock_provider): def test_hits_limit_filtered_out(self, mock_provider):
mock_provider.get_live_current.side_effect = lambda code: { # is_limit_up=True → 剔除
"paused": False, "last_price": 11.0, sm = {"600519.XSHG": {"is_limit_up": True,
"high_limit": 11.0, "low_limit": 9.0, "is_paused": False, "is_limit_down": False}}
} out = filters.filter_limitup_stock(
out = filters.filter_limitup_stock(["600519.XSHG"], mock_provider) ["600519.XSHG"], mock_provider, status_map=sm,
)
assert out == [] assert out == []
def test_below_limit_kept(self, mock_provider): def test_below_limit_kept(self, mock_provider):
mock_provider.get_live_current.side_effect = lambda code: { sm = {"600519.XSHG": {"is_limit_up": False,
"paused": False, "last_price": 10.0, "is_paused": False, "is_limit_down": False}}
"high_limit": 11.0, "low_limit": 9.0, out = filters.filter_limitup_stock(
} ["600519.XSHG"], mock_provider, status_map=sm,
out = filters.filter_limitup_stock(["600519.XSHG"], mock_provider) )
assert out == ["600519.XSHG"]
def test_no_bar_kept(self, mock_provider):
# status_map[code] = None(无 bar)→ 保留(不确定涨停,宁错过允许买)
sm = {"600519.XSHG": None}
out = filters.filter_limitup_stock(
["600519.XSHG"], mock_provider, status_map=sm,
)
assert out == ["600519.XSHG"] assert out == ["600519.XSHG"]
def test_position_held_keeps_even_at_limit(self, mock_provider): def test_position_held_keeps_even_at_limit(self, mock_provider):
# 持仓中的涨停股不过滤 # 持仓中的涨停股不过滤(涨停还能继续持有)
mock_provider.get_live_current.side_effect = lambda code: { sm = {"600519.XSHG": {"is_limit_up": True,
"paused": False, "last_price": 11.0, "is_paused": False, "is_limit_down": False}}
"high_limit": 11.0, "low_limit": 9.0,
}
out = filters.filter_limitup_stock( out = filters.filter_limitup_stock(
["600519.XSHG"], mock_provider, positions=["600519.XSHG"] ["600519.XSHG"], mock_provider,
positions=["600519.XSHG"], status_map=sm,
) )
assert out == ["600519.XSHG"] assert out == ["600519.XSHG"]
def test_last_prices_override_skips_tick(self, mock_provider): def test_no_status_map_no_date_keeps_all(self, mock_provider):
mock_provider.get_live_current.side_effect = lambda code: { # 向后兼容:无 status_map 无 date → 保留所有
"paused": False, "last_price": 99.0, out = filters.filter_limitup_stock(["A.XSHG", "B.XSHG"], mock_provider)
"high_limit": 11.0, "low_limit": 9.0, assert out == ["A.XSHG", "B.XSHG"]
}
out = filters.filter_limitup_stock(
["600519.XSHG"], mock_provider, last_prices={"600519.XSHG": 10.0}
)
# last_prices=10 < high_limit=11 → 保留
assert out == ["600519.XSHG"]
# =================== filter_limitdown_stock =================== # =================== filter_limitdown_stock ===================
class TestFilterLimitDown: class TestFilterLimitDown:
def test_hits_limit_down_filtered_out(self, mock_provider): def test_hits_limit_down_filtered_out(self, mock_provider):
mock_provider.get_live_current.side_effect = lambda code: { sm = {"600519.XSHG": {"is_limit_down": True,
"paused": False, "last_price": 9.0, "is_paused": False, "is_limit_up": False}}
"high_limit": 11.0, "low_limit": 9.0, out = filters.filter_limitdown_stock(
} ["600519.XSHG"], mock_provider, status_map=sm,
out = filters.filter_limitdown_stock(["600519.XSHG"], mock_provider) )
assert out == [] assert out == []
def test_above_limit_kept(self, mock_provider): def test_above_limit_kept(self, mock_provider):
mock_provider.get_live_current.side_effect = lambda code: { sm = {"600519.XSHG": {"is_limit_down": False,
"paused": False, "last_price": 10.0, "is_paused": False, "is_limit_up": False}}
"high_limit": 11.0, "low_limit": 9.0, out = filters.filter_limitdown_stock(
} ["600519.XSHG"], mock_provider, status_map=sm,
out = filters.filter_limitdown_stock(["600519.XSHG"], mock_provider) )
assert out == ["600519.XSHG"]
def test_no_bar_kept(self, mock_provider):
# 无 bar → 保留(不确定跌停,宁错过允许卖)
sm = {"600519.XSHG": None}
out = filters.filter_limitdown_stock(
["600519.XSHG"], mock_provider, status_map=sm,
)
assert out == ["600519.XSHG"] assert out == ["600519.XSHG"]
def test_position_held_keeps_even_at_limit_down(self, mock_provider): def test_position_held_keeps_even_at_limit_down(self, mock_provider):
# 跌停要能卖才平 → 持仓不过滤 # 跌停要能卖才平 → 持仓不过滤
mock_provider.get_live_current.side_effect = lambda code: { sm = {"600519.XSHG": {"is_limit_down": True,
"paused": False, "last_price": 9.0, "is_paused": False, "is_limit_up": False}}
"high_limit": 11.0, "low_limit": 9.0,
}
out = filters.filter_limitdown_stock( out = filters.filter_limitdown_stock(
["600519.XSHG"], mock_provider, positions=["600519.XSHG"] ["600519.XSHG"], mock_provider,
positions=["600519.XSHG"], status_map=sm,
) )
assert out == ["600519.XSHG"] assert out == ["600519.XSHG"]
def test_no_status_map_no_date_keeps_all(self, mock_provider):
out = filters.filter_limitdown_stock(["A.XSHG", "B.XSHG"], mock_provider)
assert out == ["A.XSHG", "B.XSHG"]
+8
View File
@@ -85,6 +85,14 @@ def make_strategy(
provider.get_closes_panel.side_effect = _get_closes_panel provider.get_closes_panel.side_effect = _get_closes_panel
# get_limit_status_batch: 默认全部"正常交易"(filter 全保留)
def _glbs(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 = _glbs
broker = BrokerFacade() broker = BrokerFacade()
broker.order_target_value = MagicMock(return_value=MagicMock(filled=100)) broker.order_target_value = MagicMock(return_value=MagicMock(filled=100))
broker.order_value = MagicMock(return_value=MagicMock(filled=100)) broker.order_value = MagicMock(return_value=MagicMock(filled=100))
+8
View File
@@ -101,6 +101,14 @@ def make_strategy(
provider.get_closes_panel.side_effect = _get_closes_panel provider.get_closes_panel.side_effect = _get_closes_panel
# get_limit_status_batch: 默认全部"正常交易"(filter 全保留)
def _glbs(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 = _glbs
broker = BrokerFacade() broker = BrokerFacade()
broker.order_target_value = MagicMock(return_value=MagicMock(filled=100)) broker.order_target_value = MagicMock(return_value=MagicMock(filled=100))
broker.order_value = MagicMock(return_value=MagicMock(filled=100)) broker.order_value = MagicMock(return_value=MagicMock(filled=100))
+37 -31
View File
@@ -93,20 +93,23 @@ def make_strategy(
return metrics_map.get(stock) return metrics_map.get(stock)
provider.get_value_metrics.side_effect = _get_value_metrics provider.get_value_metrics.side_effect = _get_value_metrics
# 策略 _get_stock_list 走 batch 路径(一次并发取全部候选)
provider.get_value_metrics_batch.side_effect = lambda stocks, date=None: {
stock: metrics_map.get(stock) for stock in stocks
}
provider.get_index_stocks.return_value = [] provider.get_index_stocks.return_value = []
provider.get_security_info.return_value = { provider.get_security_info.return_value = {
"display_name": "NORMAL", "display_name": "NORMAL",
"name": "600519", "name": "600519",
"start_date": datetime(2000, 1, 1), "start_date": datetime(2000, 1, 1),
} }
provider.get_live_current.return_value = { # get_limit_status_batch: 默认全部"正常交易"(filter 全保留)
"paused": False, "last_price": 10.0, def _glbs(codes, date=None):
"high_limit": 11.0, "low_limit": 9.0, return {
} c: {"is_limit_up": False, "is_limit_down": False, "is_paused": False}
provider.get_current_tick.return_value = { for c in codes
"paused": False, "last_price": 10.0, }
"high_limit": 11.0, "low_limit": 9.0, provider.get_limit_status_batch.side_effect = _glbs
}
broker = BrokerFacade() broker = BrokerFacade()
broker.order_target_value = MagicMock(return_value=MagicMock(filled=100)) broker.order_target_value = MagicMock(return_value=MagicMock(filled=100))
@@ -364,15 +367,18 @@ class TestNoticeDateFiltering:
""" """
def test_strategy_passes_date_to_provider(self): def test_strategy_passes_date_to_provider(self):
"""策略层把 previous_date 传给 provider.get_value_metrics(stock, date)。""" """策略层把 previous_date 传给 provider.get_value_metrics_batch(stocks, date)。"""
captured_dates: List[Any] = [] captured_dates: List[Any] = []
def _capture(stock, date): def _capture_batch(stocks, date=None):
captured_dates.append(date) captured_dates.append(date)
return _make_high_metrics() return {
stock: (_make_high_metrics() if "HIGH" in stock else _make_low_metrics())
for stock in stocks
}
provider = MagicMock() provider = MagicMock()
provider.get_value_metrics.side_effect = _capture provider.get_value_metrics_batch.side_effect = _capture_batch
provider.get_index_stocks.return_value = ["HIGH.XSHG", "LOW.XSHG"] provider.get_index_stocks.return_value = ["HIGH.XSHG", "LOW.XSHG"]
provider.get_security_info.return_value = { provider.get_security_info.return_value = {
"display_name": "A", "name": "A", "start_date": datetime(2000, 1, 1), "display_name": "A", "name": "A", "start_date": datetime(2000, 1, 1),
@@ -382,18 +388,12 @@ class TestNoticeDateFiltering:
"high_limit": 11.0, "low_limit": 9.0, "high_limit": 11.0, "low_limit": 9.0,
} }
# 让第二只 metrics 全空, 这样均值 = HIGH 自身, HIGH 不过(均值=自身)
# 改为返回 LOW metrics 拉低均值, HIGH 才能过
provider.get_value_metrics.side_effect = lambda stock, date: (
_make_high_metrics() if "HIGH" in stock else _make_low_metrics()
)
s = ValueSelectionStrategy(provider=provider, broker=BrokerFacade()) s = ValueSelectionStrategy(provider=provider, broker=BrokerFacade())
s._get_stock_list(["HIGH.XSHG", "LOW.XSHG"], "2024-09-30") s._get_stock_list(["HIGH.XSHG", "LOW.XSHG"], "2024-09-30")
# provider 收到的 date 应是 "2024-09-30"(由策略层传过去) # provider 收到的 date 应是 "2024-09-30"(由策略层传过去)
# 验证 side_effect 被调用时收到 date 参数 assert provider.get_value_metrics_batch.called
assert provider.get_value_metrics.called for call in provider.get_value_metrics_batch.call_args_list:
for call in provider.get_value_metrics.call_args_list: # call.args = (stocks, date) 或 call.args = (stocks,) + kwargs
# call.args = (stock, date) 或 call.args = (stock,) + kwargs
if len(call.args) >= 2: if len(call.args) >= 2:
assert call.args[1] == "2024-09-30" assert call.args[1] == "2024-09-30"
else: else:
@@ -414,18 +414,24 @@ class TestEmptyDataSkip:
assert "BAD.XSHG" not in out assert "BAD.XSHG" not in out
# HIGH 单只剩下的情况 → 均值=自身,不过(预期行为,不阻塞主流程) # HIGH 单只剩下的情况 → 均值=自身,不过(预期行为,不阻塞主流程)
def test_provider_raises_stock_excluded(self): def test_provider_returns_none_stock_excluded(self):
"""provider 异常 → 跳过,不污染整批。""" """单只异常由 provider 吞为 None → 跳过,不污染整批。
provider = MagicMock()
# HIGH 正常, LOW 抛异常
def _gnm(stock, date=None):
if "LOW" in stock:
raise RuntimeError("三表损坏")
return _make_low_metrics()
provider.get_value_metrics.side_effect = _gnm batch 接口契约: provider.get_value_metrics_batch 内部 try/except 单只
失败 → 返 {stock: None}; 策略层只看 None 跳过,不崩。
(原 test_provider_raises_stock_excluded 语义: per-stock 错误不污染整批)
"""
provider = MagicMock()
# HIGH 正常返回 metrics, LOW 返 None(三表损坏/异常由 provider 吞)
def _gnm_batch(stocks, date=None):
return {
stock: (_make_low_metrics() if "HIGH" in stock else None)
for stock in stocks
}
provider.get_value_metrics_batch.side_effect = _gnm_batch
s = ValueSelectionStrategy(provider=provider, broker=BrokerFacade()) s = ValueSelectionStrategy(provider=provider, broker=BrokerFacade())
# 不抛异常(异常被吞) # 不抛异常
out = s._get_stock_list(["HIGH.XSHG", "LOW.XSHG"], "2024-09-30") out = s._get_stock_list(["HIGH.XSHG", "LOW.XSHG"], "2024-09-30")
assert "LOW.XSHG" not in out assert "LOW.XSHG" not in out
# HIGH 因均值=自身不过(预期), 但**没有崩** # HIGH 因均值=自身不过(预期), 但**没有崩**