Files
Simon 73d191b43a feat: 量化引擎加固 — 新增测试 + 数据/因子/回测层优化
- 新增 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 扩充配置项
2026-08-31 14:01:06 +08:00

68 lines
2.2 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
因子引擎回归测试: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")