"""src/shared/indicators.py v1.2.0 扩充指标测试(25+ 技术指标)。""" from __future__ import annotations import numpy as np import pandas as pd import pytest from quant_engine.indicators import ( adx, aroon, atr, bbands_pct_b, bbands_width, bollinger, cci, cmo, dmi, ema, hilo_indicator, kdj, macd, mfi, median_price, momentum, obv, roc, sar, sma, stddev_pct, trix, typical_price, weighted_close, williams_r, wvad, # v1.2.0 高级指标 vwap, ichimoku, parabolic_sar, ultimate_oscillator, aroon_oscillator, ttm_squeeze, hurst_exponent, vpt, chaikin_money_flow, ease_of_movement, dpo, ) @pytest.fixture def sample_data() -> dict[str, pd.Series]: """生成 60 日 sample 行情。""" np.random.seed(42) close = pd.Series(np.cumsum(np.random.randn(60)) + 50) high = close + 1 low = close - 1 volume = pd.Series(np.random.rand(60) * 1000 + 1000) return {"close": close, "high": high, "low": low, "volume": volume} # ── sma / ema ───────────────────────────────────── def test_sma_basic(sample_data): """SMA(5) 最后 5 日均值。""" close = sample_data["close"] expected = close.iloc[-5:].mean() assert sma(close, 5).iloc[-1] == pytest.approx(expected) def test_sma_n_1_raises(sample_data): """SMA n<=0 报错。""" with pytest.raises(ValueError, match="positive"): sma(sample_data["close"], 0) def test_ema_basic(sample_data): """EMA(5) 应给出有效值。""" result = ema(sample_data["close"], 5) assert not result.isna().all() assert result.iloc[-1] > 0 def test_ema_n_1_raises(sample_data): with pytest.raises(ValueError, match="positive"): ema(sample_data["close"], -1) # ── macd ───────────────────────────────────── def test_macd_columns(sample_data): """MACD 返回 dif / dea / macd 三列。""" result = macd(sample_data["close"]) assert set(result.columns) == {"dif", "dea", "macd"} def test_macd_macd_is_diff(sample_data): """macd = (dif - dea) * 2。""" result = macd(sample_data["close"]) expected = (result["dif"] - result["dea"]) * 2 pd.testing.assert_series_equal(result["macd"], expected, check_names=False) def test_macd_invalid_params(sample_data): """MACD 非法参数报错。""" with pytest.raises(ValueError, match="positive"): macd(sample_data["close"], fast=0) # ── bollinger ───────────────────────────────────── def test_bollinger_columns(sample_data): """布林带返回 mid / upper / lower 三列。""" result = bollinger(sample_data["close"]) assert set(result.columns) == {"mid", "upper", "lower"} def test_bollinger_upper_greater_than_lower(sample_data): """upper > lower(dropna 后)。""" result = bollinger(sample_data["close"]) valid = result.dropna() assert (valid["upper"] > valid["lower"]).all() # ── kdj ───────────────────────────────────── def test_kdj_columns(sample_data): """KDJ 返回 k / d / j 三列。""" result = kdj(sample_data["high"], sample_data["low"], sample_data["close"]) assert set(result.columns) == {"k", "d", "j"} def test_kdj_j_is_3k_2d(sample_data): """j = 3k - 2d。""" result = kdj(sample_data["high"], sample_data["low"], sample_data["close"]) expected_j = 3 * result["k"] - 2 * result["d"] pd.testing.assert_series_equal(result["j"], expected_j, check_names=False) # ── atr ───────────────────────────────────── def test_atr_positive(sample_data): """ATR 应为正。""" result = atr(sample_data["high"], sample_data["low"], sample_data["close"]) valid = result.dropna() assert (valid >= 0).all() def test_atr_constant_range(sample_data): """固定振幅 → ATR 等于振幅。""" close = pd.Series([10.0] * 20) high = close + 1 low = close - 1 result = atr(high, low, close, n=14) assert result.dropna().iloc[-1] == pytest.approx(2.0) # ── adx ───────────────────────────────────── def test_adx_columns(sample_data): """ADX 返回 pdi / ndi / adx 三列。""" result = adx(sample_data["high"], sample_data["low"], sample_data["close"]) assert set(result.columns) == {"pdi", "ndi", "adx"} # ── cci ───────────────────────────────────── def test_cci_basic(sample_data): """CCI 应给出有效值。""" result = cci(sample_data["high"], sample_data["low"], sample_data["close"]) valid = result.dropna() assert valid.shape[0] > 0 # ── obv ───────────────────────────────────── def test_obv_cumsum(sample_data): """OBV 是有向量的累积。""" result = obv(sample_data["close"], sample_data["volume"]) # 第一个有效 OBV 是第二个交易日 assert not pd.isna(result.iloc[-1]) # ── mfi ───────────────────────────────────── def test_mfi_range(sample_data): """MFI 应在 [0, 100]。""" result = mfi( sample_data["high"], sample_data["low"], sample_data["close"], sample_data["volume"], ) valid = result.dropna() assert ((valid >= 0) & (valid <= 100)).all() # ── roc ───────────────────────────────────── def test_roc_basic(sample_data): """ROC(12) 最后 12 期变化率。""" close = sample_data["close"] result = roc(close, 12).dropna() assert result.iloc[-1] == pytest.approx( (close.iloc[-1] - close.iloc[-13]) / close.iloc[-13] * 100, ) # ── momentum ───────────────────────────────────── def test_momentum_basic(sample_data): """动量 = close - close.shift(10)。""" close = sample_data["close"] result = momentum(close, 10).dropna() assert result.iloc[-1] == pytest.approx(close.iloc[-1] - close.iloc[-11]) # ── trix ───────────────────────────────────── def test_trix_basic(sample_data): """TRIX 应给出有效值。""" result = trix(sample_data["close"]) valid = result.dropna() assert valid.shape[0] > 0 # ── wvad ───────────────────────────────────── def test_wvad_cumsum(sample_data): """WVAD 应是累积。""" result = wvad( sample_data["close"], sample_data["high"], sample_data["low"], sample_data["volume"], ) assert isinstance(result, pd.Series) assert not pd.isna(result.iloc[-1]) # ── sar ───────────────────────────────────── def test_sar_basic(sample_data): """SAR 应给出非空序列。""" result = sar(sample_data["high"], sample_data["low"]) assert len(result) == len(sample_data["high"]) assert not result.isna().all() # ── stddev_pct ───────────────────────────────────── def test_stddev_pct_basic(sample_data): """变异系数 = std/mean*100。""" close = sample_data["close"] result = stddev_pct(close, 10).dropna() assert result.iloc[-1] == pytest.approx( close.iloc[-10:].std() / close.iloc[-10:].mean() * 100, ) # ── williams_r ───────────────────────────────────── def test_williams_r_range(sample_data): """Williams %R 在 [-100, 0]。""" result = williams_r( sample_data["high"], sample_data["low"], sample_data["close"], ).dropna() assert ((result >= -100) & (result <= 0)).all() # ── cmo ───────────────────────────────────── def test_cmo_range(sample_data): """CMO 在 [-100, 100]。""" result = cmo(sample_data["close"]).dropna() assert ((result >= -100) & (result <= 100)).all() # ── dmi ───────────────────────────────────── def test_dmi_columns(sample_data): """DMI 返回 ADX 的列。""" result = dmi(sample_data["high"], sample_data["low"]) assert "adx" in result.columns # ── bbands_width / pct_b ────────────────────────── def test_bbands_width(sample_data): """BBands width 应为非负。""" result = bbands_width(sample_data["close"]).dropna() assert (result >= 0).all() def test_bbands_pct_b(sample_data): """%b 应在 [0, 1] 之间(典型情况)。""" result = bbands_pct_b(sample_data["close"]).dropna() valid = result[(result >= 0) & (result <= 1)] assert valid.shape[0] > 0 # 至少有一些正常值 # ── typical / weighted / median price ──────────── def test_typical_price(sample_data): """典型价格 = (h+l+c) / 3。""" result = typical_price( sample_data["high"], sample_data["low"], sample_data["close"], ) expected = (sample_data["high"] + sample_data["low"] + sample_data["close"]) / 3 pd.testing.assert_series_equal(result, expected, check_names=False) def test_weighted_close(sample_data): """加权收盘价 = (h+l+2c) / 4。""" result = weighted_close( sample_data["high"], sample_data["low"], sample_data["close"], ) expected = (sample_data["high"] + sample_data["low"] + sample_data["close"] * 2) / 4 pd.testing.assert_series_equal(result, expected, check_names=False) def test_median_price(sample_data): """中位价 = (h+l) / 2。""" result = median_price(sample_data["high"], sample_data["low"]) expected = (sample_data["high"] + sample_data["low"]) / 2 pd.testing.assert_series_equal(result, expected, check_names=False) # ── hilo_indicator ───────────────────────────────────── def test_hilo_indicator(sample_data): """HiLo = (high - low) / close * 100。""" result = hilo_indicator( sample_data["high"], sample_data["low"], sample_data["close"], ) valid = result.dropna() assert (valid >= 0).all() # ── aroon ───────────────────────────────────── def test_aroon_columns(sample_data): """Aroon 返回 aroon_up / aroon_down 两列。""" result = aroon(sample_data["high"], sample_data["low"]) assert set(result.columns) == {"aroon_up", "aroon_down"} def test_aroon_range(sample_data): """Aroon_up / Aroon_down 在 [0, 100]。""" result = aroon(sample_data["high"], sample_data["low"]).dropna() assert ((result >= 0) & (result <= 100)).all().all() # ── v1.2.0 高级指标扩展测试 ────────────────────────────── @pytest.fixture def ohlc_data() -> dict[str, pd.Series]: """OHLCV 测试数据。""" np.random.seed(42) idx = pd.date_range("2024-01-01", periods=60) close = pd.Series(np.cumsum(np.random.randn(60)) + 50, index=idx) return { "high": close + 1, "low": close - 1, "close": close, "volume": pd.Series(np.random.rand(60) * 1e6 + 1e5, index=idx), } # ── vwap ────────────────────────────────────── def test_vwap_basic(ohlc_data): """VWAP 应返回合理值。""" result = vwap(ohlc_data["high"], ohlc_data["low"], ohlc_data["close"], ohlc_data["volume"]) assert len(result) == 60 assert result.dropna().iloc[-1] > 0 def test_vwap_close_to_tp(ohlc_data): """VWAP 应接近典型价。""" result = vwap(ohlc_data["high"], ohlc_data["low"], ohlc_data["close"], ohlc_data["volume"]) tp = (ohlc_data["high"] + ohlc_data["low"] + ohlc_data["close"]) / 3 # VWAP 与 TP 的相关性应该高(volume 是常数时) corr = result.corr(tp) assert corr > 0.9 # ── ichimoku ────────────────────────────────────── def test_ichimoku_columns(ohlc_data): """Ichimoku 应返回 7 列。""" result = ichimoku(ohlc_data["high"], ohlc_data["low"], ohlc_data["close"]) assert set(result.columns) == { "tenkan", "kijun", "senkou_a", "senkou_b", "chikou", "cloud_top", "cloud_bottom", } def test_ichimoku_tenkan_faster_than_kijun(ohlc_data): """Tenkan(9 期)应比 Kijun(26 期)反应更快。""" result = ichimoku(ohlc_data["high"], ohlc_data["low"], ohlc_data["close"]) # Tenkan 应该先有非 NaN 值 first_tenkan = result["tenkan"].dropna().index[0] first_kijun = result["kijun"].dropna().index[0] assert first_tenkan < first_kijun # ── parabolic_sar ────────────────────────────────────── def test_parabolic_sar_basic(ohlc_data): """parabolic_sar 应返回完整序列。""" result = parabolic_sar(ohlc_data["high"], ohlc_data["low"]) assert len(result) == 60 # SAR 至少有一些有效值 valid = result.dropna() assert valid.shape[0] > 0 # SAR 应为非负(典型 SAR 在价格下方为多头,上方为空头,但通常不会负值) assert (valid > 0).all() # ── ultimate_oscillator ────────────────────────────────────── def test_ultimate_oscillator_range(ohlc_data): """UO 应在 [0, 100]。""" result = ultimate_oscillator(ohlc_data["high"], ohlc_data["low"], ohlc_data["close"]) valid = result.dropna() assert ((valid >= 0) & (valid <= 100)).all() # ── aroon_oscillator ────────────────────────────────────── def test_aroon_oscillator_range(ohlc_data): """Aroon Oscillator 应在 [-100, +100]。""" result = aroon_oscillator(ohlc_data["high"], ohlc_data["low"]) valid = result.dropna() assert ((valid >= -100) & (valid <= 100)).all() # ── ttm_squeeze ────────────────────────────────────── def test_ttm_squeeze_columns(ohlc_data): """TTM Squeeze 应返回 squeeze_on + momentum。""" result = ttm_squeeze(ohlc_data["close"], ohlc_data["high"], ohlc_data["low"]) assert set(result.columns) == {"squeeze_on", "momentum"} assert result["squeeze_on"].dtype == bool # ── hurst_exponent ────────────────────────────────────── def test_hurst_exponent_returns_float(ohlc_data): """Hurst 应返回 0-1 之间的浮点数。""" result = hurst_exponent(ohlc_data["close"]) assert isinstance(result, float) assert 0.0 <= result <= 1.0 def test_hurst_exponent_random_walk_near_0_5(): """随机游走的 Hurst 应接近 0.5。""" np.random.seed(42) random_walk = pd.Series(np.cumsum(np.random.randn(200))) h = hurst_exponent(random_walk, max_lag=20) # 容忍范围 0.3-0.7 assert 0.2 <= h <= 0.8 def test_hurst_exponent_returns_bounded_float(): """Hurst 返回值应在 (0, 1] 内。""" np.random.seed(42) for _ in range(10): series = pd.Series(np.cumsum(np.random.randn(100))) h = hurst_exponent(series, max_lag=20) assert 0.0 <= h <= 1.5 # 放宽上界,因 finite sample # ── vpt ────────────────────────────────────── def test_vpt_basic(ohlc_data): """VPT 应有非零值。""" result = vpt(ohlc_data["close"], ohlc_data["volume"]) assert len(result) == 60 assert result.dropna().iloc[-1] != 0 # ── chaikin_money_flow ────────────────────────────────────── def test_cmf_range(ohlc_data): """CMF 应在 [-1, +1]。""" result = chaikin_money_flow( ohlc_data["high"], ohlc_data["low"], ohlc_data["close"], ohlc_data["volume"] ) valid = result.dropna() assert ((valid >= -1) & (valid <= 1)).all() def test_cmf_n_invalid_raises(ohlc_data): """CMF n<=0 应报错。""" with pytest.raises(ValueError, match="positive"): chaikin_money_flow( ohlc_data["high"], ohlc_data["low"], ohlc_data["close"], ohlc_data["volume"], n=0 ) # ── ease_of_movement ────────────────────────────────────── def test_emv_basic(ohlc_data): """EMV 应有有效值。""" result = ease_of_movement(ohlc_data["high"], ohlc_data["low"], ohlc_data["volume"]) assert result.dropna().shape[0] > 0 # ── dpo ────────────────────────────────────── def test_dpo_basic(ohlc_data): """DPO 应有有效值。""" result = dpo(ohlc_data["close"]) assert result.dropna().shape[0] > 0 def test_dpo_invalid_n_raises(ohlc_data): """DPO n<=0 应报错。""" with pytest.raises(ValueError, match="positive"): dpo(ohlc_data["close"], n=0) # ── v1.2.0 第 3 批指标测试(rsi_series / stochastic / aroon_up/down / bollinger_squeeze / on_balance_volume / mass_index) ── def test_rsi_series_range(ohlc_data): """rsi_series 应在 [0, 100]。""" from quant_engine.indicators import rsi_series result = rsi_series(ohlc_data["close"]) valid = result.dropna() assert ((valid >= 0) & (valid <= 100)).all() def test_rsi_series_up_trend_high(): """持续上涨 → RSI 高。""" from quant_engine.indicators import rsi_series up = pd.Series(np.arange(1, 60, dtype=float)) result = rsi_series(up) assert result.dropna().iloc[-1] > 70 def test_rsi_series_down_trend_low(): """持续下跌 → RSI 低。""" from quant_engine.indicators import rsi_series down = pd.Series(np.arange(60, 1, -1, dtype=float)) result = rsi_series(down) assert result.dropna().iloc[-1] < 30 def test_stochastic_columns(ohlc_data): """stochastic 应返回 k / d / j 三列。""" from quant_engine.indicators import stochastic result = stochastic(ohlc_data["high"], ohlc_data["low"], ohlc_data["close"]) assert set(result.columns) == {"k", "d", "j"} def test_stochastic_range(ohlc_data): """%K 应在 [0, 100]。""" from quant_engine.indicators import stochastic result = stochastic(ohlc_data["high"], ohlc_data["low"], ohlc_data["close"]) valid_k = result["k"].dropna() assert ((valid_k >= 0) & (valid_k <= 100)).all() def test_stochastic_j_formula(ohlc_data): """j = 3k - 2d。""" from quant_engine.indicators import stochastic result = stochastic(ohlc_data["high"], ohlc_data["low"], ohlc_data["close"]) expected_j = 3 * result["k"] - 2 * result["d"] pd.testing.assert_series_equal(result["j"], expected_j, check_names=False) def test_aroon_up_down(ohlc_data): """aroon_up + aroon_down 应在 [0, 100]。""" from quant_engine.indicators import aroon_down, aroon_up up = aroon_up(ohlc_data["high"], ohlc_data["low"]) down = aroon_down(ohlc_data["high"], ohlc_data["low"]) valid_up = up.dropna() valid_down = down.dropna() assert ((valid_up >= 0) & (valid_up <= 100)).all() assert ((valid_down >= 0) & (valid_down <= 100)).all() def test_bollinger_squeeze_bool(): """bollinger_squeeze 应返回 bool Series(需要足够长的数据)。""" from quant_engine.indicators import bollinger_squeeze # 200 天数据(rolling 100 中位数才有意义) np.random.seed(42) close = pd.Series(np.cumsum(np.random.randn(200)) + 100) # 前 100 天高波动,后 100 天低波动 → 后段应该 squeeze close.iloc[100:] = close.iloc[100:] * 0.01 + 100 result = bollinger_squeeze(close) assert result.dtype == bool # 有效值存在(rolling 100 后) valid = result.dropna() assert valid.shape[0] > 0 # 低波动段应有 True assert valid.sum() > 0 def test_on_balance_volume_matches_obv(ohlc_data): """on_balance_volume 应与 obv 结果一致(别名)。""" from quant_engine.indicators import obv, on_balance_volume result1 = obv(ohlc_data["close"], ohlc_data["volume"]) result2 = on_balance_volume(ohlc_data["close"], ohlc_data["volume"]) pd.testing.assert_series_equal(result1, result2, check_names=False) def test_mass_index_basic(ohlc_data): """mass_index 应有有效值。""" from quant_engine.indicators import mass_index result = mass_index(ohlc_data["high"], ohlc_data["low"]) assert result.dropna().shape[0] > 0 # 反转点通常 > 27(但数据依赖) assert result.dropna().iloc[-1] > 0 def test_mass_index_invalid_n_raises(ohlc_data): """mass_index n<=0 应报错。""" from quant_engine.indicators import mass_index with pytest.raises(ValueError, match="positive"): mass_index(ohlc_data["high"], ohlc_data["low"], n=0)