diff --git a/sanguo_portfolio/filters.py b/sanguo_portfolio/filters.py index fcb425e..a9cda33 100644 --- a/sanguo_portfolio/filters.py +++ b/sanguo_portfolio/filters.py @@ -1,15 +1,22 @@ """A 股选股过滤器(ST/停牌/科创北交/次新/涨跌停)。 聚宽全天候策略用到的 filter_* 函数的本地实现。 -聚宽原版用 ``get_current_data()`` 拿快照,我们用注入的 provider 拿同样的字段。 设计:接受一个 ``provider`` 参数(duck typing),不依赖具体类。 provider 需要提供: - ``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 -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 @@ -67,30 +74,46 @@ def filter_st_stock(stocks: Iterable[str], provider: Any) -> List[str]: 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] = [] for stock in stocks: - tick = _safe_call(_get_tick, provider, stock) - if tick and tick.get("paused"): + st = status_map.get(stock) + if st is None: + # 无 bar = 停牌(数据 session 指定)→ 剔除 + continue + if st.get("is_paused") is True: continue result.append(stock) 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]: """过滤科创北交所:代码 4/8(北交)/68(科创)/3(创业板)开头。 @@ -152,33 +175,44 @@ def filter_limitup_stock( stocks: Iterable[str], provider: Any, positions: Optional[Sequence[str]] = None, - last_prices: Optional[dict] = None, + status_map: Optional[dict] = None, + date: Optional[str] = None, ) -> 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 []) - prices = last_prices or {} result: List[str] = [] for stock in stocks: if stock in pos_set: result.append(stock) continue - price = prices.get(stock) - high_limit = None - if price is None: - 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: + st = status_map.get(stock) + if st is None: + # 无 bar → 保留(不确定涨停,宁错过允许买) result.append(stock) continue - if price >= high_limit: + if st.get("is_limit_up") is True: continue result.append(stock) return result @@ -188,32 +222,44 @@ def filter_limitdown_stock( stocks: Iterable[str], provider: Any, positions: Optional[Sequence[str]] = None, - last_prices: Optional[dict] = None, + status_map: Optional[dict] = None, + date: Optional[str] = None, ) -> 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 []) - prices = last_prices or {} result: List[str] = [] for stock in stocks: if stock in pos_set: result.append(stock) continue - price = prices.get(stock) - low_limit = None - if price is None: - 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: + st = status_map.get(stock) + if st is None: + # 无 bar → 保留(不确定跌停,宁错过允许卖) result.append(stock) continue - if price <= low_limit: + if st.get("is_limit_down") is True: continue result.append(stock) return result diff --git a/sanguo_portfolio/strategies/momentum_timing.py b/sanguo_portfolio/strategies/momentum_timing.py index da28358..0e7046f 100644 --- a/sanguo_portfolio/strategies/momentum_timing.py +++ b/sanguo_portfolio/strategies/momentum_timing.py @@ -163,13 +163,19 @@ class MomentumTimingStrategy: stocks = list(rps_df["code"])[: cfg.top_k] # 5) 过滤涨停/跌停/停牌(复用 sanguo_portfolio.filters) + # 批量预取当日涨跌停/停牌状态(数据 session 判断好),三个 filter 共享一次查询 + status_map = self._get_limit_status(stocks, cur_date) 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, 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) # 6) 调仓:先清掉不在 stocks 的 @@ -370,6 +376,23 @@ class MomentumTimingStrategy: 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]: """成分股 + 过滤 ST/科创北交/次新。""" try: diff --git a/sanguo_portfolio/strategies/small_cap.py b/sanguo_portfolio/strategies/small_cap.py index 6429083..e4ede43 100644 --- a/sanguo_portfolio/strategies/small_cap.py +++ b/sanguo_portfolio/strategies/small_cap.py @@ -231,13 +231,19 @@ class SmallCapStrategy: ) # 6) 过滤 ST/停牌/涨跌停(原策略 current_data 过滤) + # 批量预取当日涨跌停/停牌状态(数据 session 判断好),三个 filter 共享一次查询 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, 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, 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) @@ -372,6 +378,23 @@ class SmallCapStrategy: 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]: """全市场候选池 = universe 成份股 + 过滤创业板/科创北交。 diff --git a/sanguo_portfolio/strategies/value_selection.py b/sanguo_portfolio/strategies/value_selection.py index 02841cd..704977c 100644 --- a/sanguo_portfolio/strategies/value_selection.py +++ b/sanguo_portfolio/strategies/value_selection.py @@ -170,14 +170,20 @@ class ValueSelectionStrategy: logger.info("[%s] 6条过滤后候选:%d/%d", previous_date, len(buy_list), len(candidates)) # 3) 过滤涨停/跌停/停牌(复用 sanguo_portfolio.filters) + # 批量预取当日涨跌停/停牌状态(数据 session 判断好),三个 filter 共享一次查询 positions = _get_positions(context) + status_map = self._get_limit_status(buy_list, previous_date) 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, 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) # 4) 调仓:卖出不在 buy_list 的(原策略 sell 函数) @@ -220,12 +226,11 @@ class ValueSelectionStrategy: return [] # 1) 取所有候选股的多期指标(provider 实现 NOTICE_DATE 过滤) - metrics: dict[str, dict[str, Any]] = {} - for stock in stocks: - m = self._load_value_metrics(stock, date_str) - if m is None: - continue - metrics[stock] = m + # 批量一次取(provider 层 ThreadPool 并发),等价于逐只 get_value_metrics 但快 5~8x + raw_map = self._load_value_metrics_batch(stocks, date_str) + metrics: dict[str, dict[str, Any]] = { + s: m for s, m in raw_map.items() if m is not None + } if not metrics: logger.warning("[%s] 所有股票多期指标都为空,返回空列表", date_str) @@ -380,6 +385,23 @@ class ValueSelectionStrategy: 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]: """成份股 + 过滤 ST/科创北交/次新。""" try: @@ -396,23 +418,26 @@ class ValueSelectionStrategy: ) return stocks - def _load_value_metrics( - self, stock: str, date_str: str, - ) -> Optional[dict[str, Any]]: - """从 provider 取该股的多期价值精选指标。 + def _load_value_metrics_batch( + self, stocks: List[str], date_str: str, + ) -> dict[str, Optional[dict[str, Any]]]: + """从 provider 批量取多期价值精选指标(ThreadPool 并发提速)。 - 调用 provider 的 ``get_value_metrics(stock, date_str)`` 接口(由 provider 层 - 实现 NOTICE_DATE 过滤和聚宽→东财字段映射)。provider 未实现该接口 / 返回 - None / 异常 → 该股被跳过(不入选)。 + 调用 provider 的 ``get_value_metrics_batch(stocks, date_str)`` 接口 + (provider 层 ThreadPool 并发逐只 get_value_metrics,NOTICE_DATE 过滤和 + 聚宽→东财字段映射不变)。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: - return None + return {} try: - return fn(stock, date_str) + return fn(stocks, date_str) except Exception as exc: - logger.debug("get_value_metrics(%s) 失败: %s", stock, exc) - return None + logger.debug("get_value_metrics_batch 失败: %s", exc) + return {} # ======================== 数值辅助 ======================== diff --git a/tests/portfolio/test_filters.py b/tests/portfolio/test_filters.py index 131d2aa..46d9087 100644 --- a/tests/portfolio/test_filters.py +++ b/tests/portfolio/test_filters.py @@ -46,25 +46,48 @@ class TestFilterSt: # =================== filter_paused_stock =================== class TestFilterPaused: def test_paused_stock_filtered(self, mock_provider): - mock_provider.get_live_current.side_effect = lambda code: { - "paused": True, "last_price": 10.0, - "high_limit": 11.0, "low_limit": 9.0, - } - out = filters.filter_paused_stock(["600519.XSHG"], mock_provider) + # is_paused=True → 剔除 + sm = {"600519.XSHG": {"is_paused": True, + "is_limit_up": False, "is_limit_down": False}} + out = filters.filter_paused_stock( + ["600519.XSHG"], mock_provider, status_map=sm, + ) assert out == [] def test_trading_stock_kept(self, mock_provider): - mock_provider.get_live_current.side_effect = lambda code: { - "paused": False, "last_price": 10.0, - "high_limit": 11.0, "low_limit": 9.0, - } - out = filters.filter_paused_stock(["600519.XSHG"], mock_provider) + sm = {"600519.XSHG": {"is_paused": False, + "is_limit_up": False, "is_limit_down": False}} + out = filters.filter_paused_stock( + ["600519.XSHG"], mock_provider, status_map=sm, + ) assert out == ["600519.XSHG"] - def test_tick_failure_keeps_stock(self, mock_provider): - mock_provider.get_live_current.side_effect = Exception("no tick") - out = filters.filter_paused_stock(["600519.XSHG"], mock_provider) - assert out == ["600519.XSHG"] + def test_no_bar_filtered_as_paused(self, mock_provider): + # status_map[code] = None(无 bar = 停牌,数据 session 指定)→ 剔除 + sm = {"600519.XSHG": None} + 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 =================== @@ -148,69 +171,82 @@ class TestFilterNewStock: # =================== filter_limitup_stock =================== class TestFilterLimitUp: def test_hits_limit_filtered_out(self, mock_provider): - mock_provider.get_live_current.side_effect = lambda code: { - "paused": False, "last_price": 11.0, - "high_limit": 11.0, "low_limit": 9.0, - } - out = filters.filter_limitup_stock(["600519.XSHG"], mock_provider) + # is_limit_up=True → 剔除 + sm = {"600519.XSHG": {"is_limit_up": True, + "is_paused": False, "is_limit_down": False}} + out = filters.filter_limitup_stock( + ["600519.XSHG"], mock_provider, status_map=sm, + ) assert out == [] def test_below_limit_kept(self, mock_provider): - mock_provider.get_live_current.side_effect = lambda code: { - "paused": False, "last_price": 10.0, - "high_limit": 11.0, "low_limit": 9.0, - } - out = filters.filter_limitup_stock(["600519.XSHG"], mock_provider) + sm = {"600519.XSHG": {"is_limit_up": False, + "is_paused": False, "is_limit_down": False}} + out = filters.filter_limitup_stock( + ["600519.XSHG"], mock_provider, status_map=sm, + ) + 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"] def test_position_held_keeps_even_at_limit(self, mock_provider): - # 持仓中的涨停股不过滤 - mock_provider.get_live_current.side_effect = lambda code: { - "paused": False, "last_price": 11.0, - "high_limit": 11.0, "low_limit": 9.0, - } + # 持仓中的涨停股不过滤(涨停还能继续持有) + sm = {"600519.XSHG": {"is_limit_up": True, + "is_paused": False, "is_limit_down": False}} 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"] - def test_last_prices_override_skips_tick(self, mock_provider): - mock_provider.get_live_current.side_effect = lambda code: { - "paused": False, "last_price": 99.0, - "high_limit": 11.0, "low_limit": 9.0, - } - 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"] + def test_no_status_map_no_date_keeps_all(self, mock_provider): + # 向后兼容:无 status_map 无 date → 保留所有 + out = filters.filter_limitup_stock(["A.XSHG", "B.XSHG"], mock_provider) + assert out == ["A.XSHG", "B.XSHG"] # =================== filter_limitdown_stock =================== class TestFilterLimitDown: def test_hits_limit_down_filtered_out(self, mock_provider): - mock_provider.get_live_current.side_effect = lambda code: { - "paused": False, "last_price": 9.0, - "high_limit": 11.0, "low_limit": 9.0, - } - out = filters.filter_limitdown_stock(["600519.XSHG"], mock_provider) + sm = {"600519.XSHG": {"is_limit_down": True, + "is_paused": False, "is_limit_up": False}} + out = filters.filter_limitdown_stock( + ["600519.XSHG"], mock_provider, status_map=sm, + ) assert out == [] def test_above_limit_kept(self, mock_provider): - mock_provider.get_live_current.side_effect = lambda code: { - "paused": False, "last_price": 10.0, - "high_limit": 11.0, "low_limit": 9.0, - } - out = filters.filter_limitdown_stock(["600519.XSHG"], mock_provider) + sm = {"600519.XSHG": {"is_limit_down": False, + "is_paused": False, "is_limit_up": False}} + out = filters.filter_limitdown_stock( + ["600519.XSHG"], mock_provider, status_map=sm, + ) + 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"] def test_position_held_keeps_even_at_limit_down(self, mock_provider): # 跌停要能卖才平 → 持仓不过滤 - mock_provider.get_live_current.side_effect = lambda code: { - "paused": False, "last_price": 9.0, - "high_limit": 11.0, "low_limit": 9.0, - } + sm = {"600519.XSHG": {"is_limit_down": True, + "is_paused": False, "is_limit_up": False}} 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"] + + 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"] diff --git a/tests/portfolio/test_momentum_timing.py b/tests/portfolio/test_momentum_timing.py index 893c6f8..8205363 100644 --- a/tests/portfolio/test_momentum_timing.py +++ b/tests/portfolio/test_momentum_timing.py @@ -85,6 +85,14 @@ def make_strategy( 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.order_target_value = MagicMock(return_value=MagicMock(filled=100)) broker.order_value = MagicMock(return_value=MagicMock(filled=100)) diff --git a/tests/portfolio/test_small_cap.py b/tests/portfolio/test_small_cap.py index 4dc216b..7273994 100644 --- a/tests/portfolio/test_small_cap.py +++ b/tests/portfolio/test_small_cap.py @@ -101,6 +101,14 @@ def make_strategy( 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.order_target_value = MagicMock(return_value=MagicMock(filled=100)) broker.order_value = MagicMock(return_value=MagicMock(filled=100)) diff --git a/tests/portfolio/test_value_selection.py b/tests/portfolio/test_value_selection.py index 8c9b6b4..4f523b2 100644 --- a/tests/portfolio/test_value_selection.py +++ b/tests/portfolio/test_value_selection.py @@ -93,20 +93,23 @@ def make_strategy( return metrics_map.get(stock) 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_security_info.return_value = { "display_name": "NORMAL", "name": "600519", "start_date": datetime(2000, 1, 1), } - provider.get_live_current.return_value = { - "paused": False, "last_price": 10.0, - "high_limit": 11.0, "low_limit": 9.0, - } - provider.get_current_tick.return_value = { - "paused": False, "last_price": 10.0, - "high_limit": 11.0, "low_limit": 9.0, - } + # 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.order_target_value = MagicMock(return_value=MagicMock(filled=100)) @@ -364,15 +367,18 @@ class TestNoticeDateFiltering: """ 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] = [] - def _capture(stock, date): + def _capture_batch(stocks, date=None): 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.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_security_info.return_value = { "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, } - # 让第二只 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._get_stock_list(["HIGH.XSHG", "LOW.XSHG"], "2024-09-30") # provider 收到的 date 应是 "2024-09-30"(由策略层传过去) - # 验证 side_effect 被调用时收到 date 参数 - assert provider.get_value_metrics.called - for call in provider.get_value_metrics.call_args_list: - # call.args = (stock, date) 或 call.args = (stock,) + kwargs + assert provider.get_value_metrics_batch.called + for call in provider.get_value_metrics_batch.call_args_list: + # call.args = (stocks, date) 或 call.args = (stocks,) + kwargs if len(call.args) >= 2: assert call.args[1] == "2024-09-30" else: @@ -414,18 +414,24 @@ class TestEmptyDataSkip: assert "BAD.XSHG" not in out # HIGH 单只剩下的情况 → 均值=自身,不过(预期行为,不阻塞主流程) - def test_provider_raises_stock_excluded(self): - """provider 异常 → 跳过,不污染整批。""" - provider = MagicMock() - # HIGH 正常, LOW 抛异常 - def _gnm(stock, date=None): - if "LOW" in stock: - raise RuntimeError("三表损坏") - return _make_low_metrics() + def test_provider_returns_none_stock_excluded(self): + """单只异常由 provider 吞为 None → 跳过,不污染整批。 - 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()) - # 不抛异常(异常被吞) + # 不抛异常 out = s._get_stock_list(["HIGH.XSHG", "LOW.XSHG"], "2024-09-30") assert "LOW.XSHG" not in out # HIGH 因均值=自身不过(预期), 但**没有崩**