70 lines
2.1 KiB
Python
70 lines
2.1 KiB
Python
"""Boundary tests for TDX-style selection indicators."""
|
|
|
|
import numpy as np
|
|
import pandas as pd
|
|
import pytest
|
|
|
|
from zhixing_server.modules.selection.domain.indicators import (
|
|
BARSLAST,
|
|
COUNT,
|
|
CROSS,
|
|
EVERY,
|
|
HHVBARS,
|
|
MA,
|
|
REF,
|
|
compute_amplitude_params,
|
|
compute_kdj,
|
|
compute_rsi,
|
|
)
|
|
|
|
|
|
def test_rolling_primitives_use_trading_rows_and_keep_every_warmup() -> None:
|
|
values = pd.Series([1.0, 2.0, 3.0, 2.0])
|
|
|
|
assert MA(values, 3).tolist() == [1.0, 1.5, 2.0, 7 / 3]
|
|
assert REF(values, 1).isna().iloc[0]
|
|
assert EVERY(pd.Series([True, True, True]), 3).tolist() == [False, False, True]
|
|
assert COUNT(pd.Series([True, False, True]), 2).tolist() == [1.0, 1.0, 1.0]
|
|
|
|
|
|
def test_hhvbars_and_barslast_are_stable_for_ties_and_missing_prefix() -> None:
|
|
values = pd.Series([1.0, 3.0, 3.0, 2.0])
|
|
assert HHVBARS(values, 3).tolist() == [0.0, 0.0, 0.0, 1.0]
|
|
bars_last = BARSLAST(pd.Series([False, True, False, True]))
|
|
assert np.isnan(bars_last.iloc[0])
|
|
assert bars_last.iloc[1:].tolist() == [0.0, 1.0, 0.0]
|
|
|
|
|
|
def test_cross_does_not_match_without_a_previous_complete_row() -> None:
|
|
assert CROSS(pd.Series([1.0, 3.0, 2.0]), pd.Series([2.0, 2.0, 2.0])).tolist() == [
|
|
False,
|
|
True,
|
|
False,
|
|
]
|
|
|
|
|
|
def test_zero_range_and_zero_rsi_denominator_do_not_create_finite_signals() -> None:
|
|
frame = pd.DataFrame(
|
|
{
|
|
"low": [10.0, 10.0, 10.0],
|
|
"high": [10.0, 10.0, 10.0],
|
|
"close": [10.0, 10.0, 10.0],
|
|
}
|
|
)
|
|
kdj = compute_kdj(frame, 3)
|
|
rsi = compute_rsi(frame["close"], 3)
|
|
assert kdj["J"].isna().all()
|
|
assert rsi.isna().all()
|
|
|
|
|
|
def test_amplitude_parameters_cover_wide_prefix_and_historical_wide_move() -> None:
|
|
close = pd.Series([10.0, 10.0, 11.6, 11.0])
|
|
assert compute_amplitude_params("688001", close) == (8.0, 0.9)
|
|
assert compute_amplitude_params("000001", close) == (8.0, 0.9)
|
|
assert compute_amplitude_params("000001", pd.Series([10.0, 10.1])) == (5.0, 1.0)
|
|
|
|
|
|
def test_invalid_indicator_windows_fail_loudly() -> None:
|
|
with pytest.raises(ValueError):
|
|
MA(pd.Series([1.0]), 0)
|