""" 因子引擎回归测试:RSI 除零修复 + 注册表名一致 + 34 因子核实。 """ import sys import os sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) import numpy as np import pandas as pd import pytest from factors.technical.rsi import RSIFactor from factors.registry import get_factor, list_factors, list_categories from factors.registry import FACTOR_CATEGORIES @pytest.fixture def rising_close(): """60 天纯上涨:loss 恒 0,标准 RSI 应为 100(绝非 NaN)。""" return pd.DataFrame({"close": np.arange(100, 160, 1.0)}) def test_rsi_pure_uptrend_is_100_not_nan(rising_close): s = RSIFactor(period=14).calculate(rising_close) assert not s.dropna().empty, "纯上涨 RSI 不应全 NaN(旧版 replace(0,nan) 的 bug)" assert s.dropna().iloc[-1] == 100.0 def test_rsi_pure_downtrend_is_0(): f = pd.DataFrame({"close": np.arange(160, 100, -1.0)}) s = RSIFactor(period=14).calculate(f) assert s.dropna().iloc[-1] < 1.0 def test_rsi_mixed_in_range(): f = pd.DataFrame({"close": np.sin(np.arange(100) * 0.5) * 10 + 100}) s = RSIFactor(period=14).calculate(f) assert (s.dropna().between(0, 100)).all() assert not s.isna().all() def test_rsi_flat_no_crash(): s = RSIFactor(period=14).calculate(pd.DataFrame({"close": [100.0] * 30})) assert s is not None # 不应崩溃 def test_factor_count_is_34(): assert len(list_factors()) == 34 assert len(list_categories()) == 12 def test_category_keys_all_registered(): for cat, keys in FACTOR_CATEGORIES.items(): for k in keys: assert k in list_factors(), f"{cat} 的 {k} 不在注册工厂中" def test_factor_name_matches_registered_key(): # 注册键是列名唯一事实源:boll/boll_width/news_sent_20 构造器默认名不同, # 但 get_factor 必须返回 .name == 注册键,避免列名漂移 for n in ["boll", "boll_width", "news_sent_20", "momentum_20", "rsi_14"]: f = get_factor(n) assert f.name == n, f"get_factor('{n}').name = '{f.name}' 应为 '{n}'" def test_get_factor_unknown_raises(): with pytest.raises(KeyError): get_factor("not_a_factor_xyz")