mirror of
https://ghfast.top/https://github.com/aeroxw/easy-tdx.git
synced 2026-09-12 20:24:16 +08:00
- pyproject.toml: add mypy overrides for pandas/tabulate/matplotlib stubs, disable strict checking for vendored MyTT library - config.py: use cast() for dict[str, Any] .get() returns - beichi.py: widen _calc_bi_force param to BI | XD, import XD - backtest/cli.py: split combo/single strategy into separate typed variables - backtest/combo.py: add bool_array() helper for numpy return types - chanlun/analyser.py: type ignore for pandas row access, fix dict type arg - unified.py: change fields param from object to Any - ex/mac_client.py: add type args to list literals - cli/cmd_offline.py: wrap int market as Market enum before API call - cli/cmd_chanlun.py: fix dict type arg - offline/write_*.py: explicit int() cast for struct.unpack returns - MyTT.py: fix line-too-long comments, UP038 isinstance syntax - tests: fix E712 (==False → ~mask), E741 (noqa), F841, import sorting - ruff format applied across codebase Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
161 lines
5.4 KiB
Python
161 lines
5.4 KiB
Python
"""indicator.py 离线单元测试。"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import warnings
|
|
|
|
import numpy as np
|
|
import pandas as pd
|
|
import pytest
|
|
|
|
from easy_tdx.indicator import _REGISTRY, compute_indicators, list_indicators
|
|
|
|
|
|
def _make_ohlcv(n: int = 200, seed: int = 42) -> pd.DataFrame:
|
|
rng = np.random.default_rng(seed)
|
|
close = 100 + np.cumsum(rng.standard_normal(n) * 0.5)
|
|
high = close + np.abs(rng.standard_normal(n))
|
|
low = close - np.abs(rng.standard_normal(n))
|
|
open_ = low + (high - low) * rng.random(n)
|
|
vol = (rng.random(n) * 1e6).astype(float)
|
|
return pd.DataFrame(
|
|
{
|
|
"datetime": pd.date_range("2024-01-01", periods=n, freq="D"),
|
|
"open": open_,
|
|
"high": high,
|
|
"low": low,
|
|
"close": close,
|
|
"vol": vol,
|
|
"amount": vol * close,
|
|
}
|
|
)
|
|
|
|
|
|
class TestRegistry:
|
|
def test_all_indicators_registered(self):
|
|
assert len(_REGISTRY) >= 22
|
|
|
|
def test_list_indicators_returns_metadata(self):
|
|
info = list_indicators()
|
|
assert len(info) >= 22
|
|
for entry in info:
|
|
assert "name" in entry
|
|
assert "inputs" in entry
|
|
assert "outputs" in entry
|
|
assert "description" in entry
|
|
|
|
|
|
class TestComputeIndicators:
|
|
def test_single_indicator_macd(self):
|
|
df = _make_ohlcv()
|
|
result = compute_indicators(df, ["MACD"])
|
|
assert "MACD_DIF" in result.columns
|
|
assert "MACD_DEA" in result.columns
|
|
assert "MACD_HIST" in result.columns
|
|
assert len(result) == 200
|
|
|
|
def test_multiple_indicators(self):
|
|
df = _make_ohlcv()
|
|
result = compute_indicators(df, ["MACD", "KDJ", "RSI"])
|
|
for col in ["MACD_DIF", "MACD_DEA", "MACD_HIST", "KDJ_K", "KDJ_D", "KDJ_J", "RSI"]:
|
|
assert col in result.columns
|
|
|
|
def test_keep_ohlcv_true(self):
|
|
df = _make_ohlcv()
|
|
result = compute_indicators(df, ["RSI"], keep_ohlcv=True)
|
|
for col in ["open", "high", "low", "close", "vol"]:
|
|
assert col in result.columns
|
|
|
|
def test_keep_ohlcv_false(self):
|
|
df = _make_ohlcv()
|
|
result = compute_indicators(df, ["RSI"], keep_ohlcv=False)
|
|
assert "close" not in result.columns
|
|
assert "RSI" in result.columns
|
|
# datetime 应保留
|
|
assert "datetime" in result.columns
|
|
|
|
def test_keep_ohlcv_false_no_time_cols(self):
|
|
df = _make_ohlcv()
|
|
df = df.drop(columns=["datetime"])
|
|
result = compute_indicators(df, ["RSI"], keep_ohlcv=False)
|
|
assert "close" not in result.columns
|
|
assert "RSI" in result.columns
|
|
|
|
def test_tail_parameter(self):
|
|
df = _make_ohlcv(200)
|
|
result = compute_indicators(df, ["MACD"], tail=30)
|
|
assert len(result) == 30
|
|
assert "MACD_DIF" in result.columns
|
|
|
|
def test_case_insensitive(self):
|
|
df = _make_ohlcv()
|
|
result = compute_indicators(df, ["macd", "kdj"])
|
|
assert "MACD_DIF" in result.columns
|
|
assert "KDJ_K" in result.columns
|
|
|
|
def test_custom_params(self):
|
|
df = _make_ohlcv()
|
|
r1 = compute_indicators(df, ["MACD"])
|
|
r2 = compute_indicators(df, ["MACD"], params={"MACD": {"SHORT": 10}})
|
|
# 不同参数应产生不同结果
|
|
assert not np.allclose(r1["MACD_DIF"].values, r2["MACD_DIF"].values, equal_nan=True)
|
|
|
|
def test_unknown_indicator_raises(self):
|
|
df = _make_ohlcv()
|
|
with pytest.raises(ValueError, match="未知指标"):
|
|
compute_indicators(df, ["FAKE_INDICATOR"])
|
|
|
|
def test_missing_input_columns_raises(self):
|
|
df = pd.DataFrame({"close": np.random.randn(200)})
|
|
with pytest.raises(ValueError, match="缺少必要列"):
|
|
compute_indicators(df, ["KDJ"])
|
|
|
|
def test_empty_dataframe(self):
|
|
df = pd.DataFrame()
|
|
result = compute_indicators(df, ["MACD"])
|
|
assert result.empty
|
|
|
|
def test_short_data_warning(self):
|
|
df = _make_ohlcv(50)
|
|
with warnings.catch_warnings(record=True) as w:
|
|
warnings.simplefilter("always")
|
|
compute_indicators(df, ["MACD"])
|
|
assert any("120" in str(warning.message) for warning in w)
|
|
|
|
def test_rsi_range(self):
|
|
df = _make_ohlcv(200)
|
|
result = compute_indicators(df, ["RSI"])
|
|
rsi = result["RSI"].dropna()
|
|
assert (rsi >= -10).all() and (rsi <= 110).all()
|
|
|
|
def test_boll_bands_order(self):
|
|
df = _make_ohlcv(200)
|
|
result = compute_indicators(df, ["BOLL"])
|
|
valid = result.dropna(subset=["BOLL_UPPER", "BOLL_LOWER"])
|
|
assert (valid["BOLL_UPPER"] >= valid["BOLL_LOWER"]).all()
|
|
|
|
def test_obv_with_volume(self):
|
|
df = _make_ohlcv()
|
|
result = compute_indicators(df, ["OBV"])
|
|
assert "OBV" in result.columns
|
|
|
|
def test_brar_needs_open(self):
|
|
df = _make_ohlcv()
|
|
result = compute_indicators(df, ["BRAR"])
|
|
assert "AR" in result.columns
|
|
assert "BR" in result.columns
|
|
|
|
def test_all_registered_indicators_run(self):
|
|
"""确保所有注册的指标都能无错运行。"""
|
|
df = _make_ohlcv(250)
|
|
for name in _REGISTRY:
|
|
result = compute_indicators(df, [name])
|
|
spec = _REGISTRY[name]
|
|
for col in spec.outputs:
|
|
assert col in result.columns, f"{name} missing output {col}"
|
|
|
|
def test_result_index_reset(self):
|
|
df = _make_ohlcv()
|
|
result = compute_indicators(df, ["RSI"])
|
|
assert list(result.index) == list(range(len(result)))
|