Coverage for src / monte_neo / indicators / rsi.py: 67%

42 statements  

« prev     ^ index     » next       coverage.py v7.13.1, created at 2026-01-28 16:27 +0200

1from __future__ import annotations 

2 

3import numpy as np 

4import pandas as pd 

5 

6from monte_neo.indicators.base import BaseIndicator, IndicatorConfig 

7from monte_neo.indicators.numba_funcs import rsi_signals_numba 

8from monte_neo.indicators.technical_lib import TechnicalIndicators 

9 

10 

11class RSIIndicator(BaseIndicator): 

12 """RSI indicator with overbought/oversold signals.""" 

13 

14 def __init__(self, config: IndicatorConfig | None = None) -> None: 

15 super().__init__(config) 

16 self._parameters.setdefault("period", 14) 

17 self._parameters.setdefault("overbought", 70) 

18 self._parameters.setdefault("oversold", 30) 

19 

20 def calculate(self, data: pd.DataFrame) -> pd.DataFrame: 

21 result = data.copy() 

22 result["rsi"] = TechnicalIndicators.rsi( 

23 data["close"], self._parameters["period"] 

24 ) 

25 return result 

26 

27 def generate_signals(self, data: pd.DataFrame) -> pd.DataFrame: 

28 # High-performance Numba-based RSI signals 

29 close = data["close"].to_numpy() 

30 period = int(round(self._parameters["period"])) 

31 oversold = float(self._parameters["oversold"]) 

32 overbought = float(self._parameters["overbought"]) 

33 

34 # Minimum period is 2 

35 period = max(2, period) 

36 

37 sig_vals = rsi_signals_numba(close, period, oversold, overbought) 

38 return pd.DataFrame({"signal": sig_vals}, index=data.index) 

39 

40 def generate_signals_fast(self, data: pd.DataFrame | np.ndarray) -> np.ndarray: 

41 if isinstance(data, pd.DataFrame): 

42 close = data["close"].to_numpy() 

43 else: 

44 close = data[:, 3] if data.ndim > 1 else data 

45 

46 period = int(round(self._parameters["period"])) 

47 period = max(2, period) 

48 oversold = float(self._parameters["oversold"]) 

49 overbought = float(self._parameters["overbought"]) 

50 

51 return rsi_signals_numba(close, period, oversold, overbought) 

52 

53 def get_formula(self) -> str: 

54 p = self._parameters["period"] 

55 ob = self._parameters["overbought"] 

56 os = self._parameters["oversold"] 

57 return f"RSI({p}) [Buy < {os}, Sell > {ob}]" 

58 

59 def get_min_periods(self) -> int: 

60 return self._parameters["period"] + 1 

61 

62 def get_metal_params(self, commission_bps: float = 0.0, slippage_bps: float = 0.0) -> list[float] | None: 

63 """Return parameters for native Metal kernel.""" 

64 # Layout: [type, p1, p2, p3, atr_period, sl_mult, tp_mult, ts_mult, commission, slippage] 

65 # type 1: RSI 

66 return [ 

67 1.0, # type 

68 float(self._parameters.get("period", 14)), 

69 float(self._parameters.get("overbought", 70)), 

70 float(self._parameters.get("oversold", 30)), 

71 14.0, # ATR 

72 1.5, # SL 

73 3.0, # TP 

74 2.0, # TS 

75 commission_bps, 

76 slippage_bps 

77 ]