- 新增 finance/tests/ 6 个测试套件(agents/backtest/dao_upsert/factors/features/fundamental_lookahead) - 数据层: data_manager / dao 优化,新增 upsert 逻辑 - 因子层: 基本面因子抽象定位 _mapping、ROE/PE/PB 重构 - 回测层: vectorbt/engine 大改动(251 行),report 增强 - ML 层: features/backtest_integration 特征工程与回测优化 - CLI: agent_cli 重构 - config/settings 扩充配置项
68 lines
2.2 KiB
Python
68 lines
2.2 KiB
Python
"""
|
||
因子引擎回归测试: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") |