# indie:lang_version = 5
from math import nan, sqrt, isnan
from indie import indicator, param, MainContext, Var, MutSeriesF, SeriesF, color, plot, algorithm, line_style
from indie.algorithms import Ema, Sma, Wma, Atr, StdDev


# =========================
# Hull Moving Average
# HMA(n) = WMA( 2*WMA(n/2) - WMA(n), sqrt(n) )
# =========================

@algorithm
def Hma(self, src: SeriesF, length: int) -> SeriesF:
    '''Hull Moving Average'''
    half     = max(1, length // 2)
    sqrt_len = max(1, int(sqrt(length)))
    wma_half = Wma.new(src, half)
    wma_full = Wma.new(src, length)
    diff     = MutSeriesF.new(2.0 * wma_half[0] - wma_full[0])
    return Wma.new(diff, sqrt_len)


# =========================
# Decorator order (top → bottom) = return tuple (index 0 → last):
#   [0] fast_ma             @plot.line   orange
#   [1] mid_ma              @plot.line   blue
#   [2] slow_ma             @plot.line   red
#   [3] long_entry_marker   @plot.marker green circle, below bar  (text = "score/max")
#   [4] short_entry_marker  @plot.marker red cross, above bar     (text = "score/max")
#   [5] long_exit_marker    @plot.marker red cross, above bar     (text = "EXIT")
#   [6] short_exit_marker   @plot.marker green circle, below bar  (text = "EXIT")
#   [7] long_stop_plot      @plot.line   gray dashed
#   [8] short_stop_plot     @plot.line   gray dashed
#
# NOTE: Score is NOT plotted as a line (it would break price-scale on overlay).
#       Score is shown as text directly on entry markers, e.g. "4/7".
# =========================

@indicator('Adaptive Triple MA Crossover', overlay_main_pane=True)
# MA lengths
@param.int('fast_len',        default=20,  min=1,            title='Fast MA Length')
@param.int('mid_len',         default=50,  min=1,            title='Mid MA Length')
@param.int('slow_len',        default=200, min=1,            title='Slow MA Length')
# ATR risk engine
@param.int('atr_len',         default=14,  min=1,            title='ATR Length')
@param.float('atr_mult',      default=2.0, min=0.1,          title='ATR Multiplier')
# Score filter parameters
@param.float('pullback_atr',  default=1.5, min=1.0, max=3.0, title='Pullback Distance (ATR)')
@param.float('impulse_atr',   default=0.5, min=0.1, max=2.0, title='Impulse Threshold (ATR)')
@param.int('score_threshold', default=3,   min=1,   max=7,   title='Score Threshold (1-7)')
@param.int('persist_bars',    default=5,   min=1,            title='Persistence Bars')
@param.int('max_bars_trade',  default=80,  min=1,            title='Max Bars in Trade')
# Filter toggles
@param.bool('use_pullback',   default=True,                   title='Pullback Filter (+2)')
@param.bool('use_impulse',    default=True,                   title='Impulse Filter (+2)')
@param.bool('use_volume',     default=True,                   title='Volume Filter (+1)')
@param.bool('use_volatility', default=True,                   title='Volatility Filter (+1)')
@param.bool('use_breakeven',  default=True,                   title='Breakeven Logic')
# Plots — order must match return tuple exactly
@plot.line(title='Fast MA (HMA)',  color=color.ORANGE)
@plot.line(title='Mid MA (EMA)',   color=color.BLUE)
@plot.line(title='Slow MA (EMA)',  color=color.RED)
@plot.marker(title='Long Entry',   color=color.GREEN,
             style=plot.marker_style.CIRCLE, position=plot.marker_position.BELOW)
@plot.marker(title='Short Entry',  color=color.RED,
             style=plot.marker_style.CROSS,  position=plot.marker_position.ABOVE)
@plot.marker(title='Long Exit',    color=color.RED,
             style=plot.marker_style.CROSS,  position=plot.marker_position.ABOVE)
@plot.marker(title='Short Exit',   color=color.GREEN,
             style=plot.marker_style.CIRCLE, position=plot.marker_position.BELOW)
@plot.line(title='Long Stop',      color=color.GRAY,   line_style=line_style.DASHED)
@plot.line(title='Short Stop',     color=color.GRAY,   line_style=line_style.DASHED)
class Main(MainContext):

    def calc(self, fast_len, mid_len, slow_len, atr_len, atr_mult,
             pullback_atr, impulse_atr, score_threshold, persist_bars, max_bars_trade,
             use_pullback, use_impulse, use_volume, use_volatility, use_breakeven):

        close = self.close

        # =========================
        # Persistent state
        # Var.new() is syntactic sugar — must be in calc(), not __init__()
        # =========================
        _in_long     = Var[bool].new(init=False)
        _in_short    = Var[bool].new(init=False)
        _entry_price = Var[float].new(init=nan)
        _entry_bar   = Var[int].new(init=0)
        _long_stop   = Var[float].new(init=nan)
        _short_stop  = Var[float].new(init=nan)

        in_long     = _in_long.get()
        in_short    = _in_short.get()
        entry_price = _entry_price.get()
        entry_bar   = _entry_bar.get()
        long_stop   = _long_stop.get()
        short_stop  = _short_stop.get()

        # =========================
        # Moving averages
        # =========================
        fast_ma = Hma.new(close, fast_len)
        mid_ma  = Ema.new(close, mid_len)
        slow_ma = Ema.new(close, slow_len)

        # =========================
        # ATR
        # =========================
        atr     = Atr.new(atr_len)
        atr_val = atr[0]

        # =========================
        # CORE TREND FILTER (required gate — no trend, no entry allowed)
        # =========================
        trend_long  = fast_ma[0] > mid_ma[0] and mid_ma[0] > slow_ma[0]
        trend_short = fast_ma[0] < mid_ma[0] and mid_ma[0] < slow_ma[0]

        # =========================
        # SCORE SYSTEM
        # All variables pre-declared before if-blocks (Indie block-level scoping).
        # Scores are directional — each filter only credits the side it confirms.
        # =========================
        score_long  = 0
        score_short = 0

        # Max possible score from all enabled filters:
        #   persistence = +1  (always on)
        #   pullback    = +2  (if use_pullback)
        #   impulse     = +2  (if use_impulse)
        #   volume      = +1  (if use_volume)
        #   volatility  = +1  (if use_volatility)
        max_score = 1
        if use_pullback:
            max_score = max_score + 2
        if use_impulse:
            max_score = max_score + 2
        if use_volume:
            max_score = max_score + 1
        if use_volatility:
            max_score = max_score + 1

        # --- Factor 1: Trend persistence (+1, directional, always active) ---
        if mid_ma[0] > mid_ma[persist_bars]:
            score_long = score_long + 1
        if mid_ma[0] < mid_ma[persist_bars]:
            score_short = score_short + 1

        # --- Factor 2: Pullback zone (+2, directional) ---
        # Price must be pulling back toward mid MA from the correct side.
        # distance = abs(close - midMA) / ATR; want < pullback_atr (default 1.5)
        distance: float = 999.0
        if atr_val > 0.0:
            distance = abs(close[0] - mid_ma[0]) / atr_val
        # Long pullback: price is above mid MA but close enough (returning from above)
        if use_pullback and distance < pullback_atr and close[0] > mid_ma[0]:
            score_long = score_long + 2
        # Short pullback: price is below mid MA but close enough (returning from below)
        if use_pullback and distance < pullback_atr and close[0] < mid_ma[0]:
            score_short = score_short + 2

        # --- Factor 3: Impulse breakout (+2, directional, always asymmetric) ---
        rng           = self.high[0] - self.low[0]
        impulse_long  = close[0] > self.high[1] and rng > atr_val * impulse_atr
        impulse_short = close[0] < self.low[1]  and rng > atr_val * impulse_atr
        if use_impulse and impulse_long:
            score_long = score_long + 2
        if use_impulse and impulse_short:
            score_short = score_short + 2

        # --- Factor 4: Volume confirmation (+1, directional by bar close) ---
        vol_ma    = Sma.new(self.volume, 30)
        vol_spike = self.volume[0] > vol_ma[0] * 1.2
        # Volume confirms the direction of the current bar's close
        if use_volume and vol_spike and close[0] > close[1]:
            score_long = score_long + 1
        if use_volume and vol_spike and close[0] < close[1]:
            score_short = score_short + 1

        # --- Factor 5: Volatility / ATR Z-score (+1, symmetric — confirms market activity) ---
        atr_mean = Sma.new(atr, 100)[0]
        atr_std  = StdDev.new(atr, 100)[0]
        atr_z: float = 0.0
        if atr_std > 0.0:
            atr_z = (atr_val - atr_mean) / atr_std
        if use_volatility and atr_z > 0.0:
            score_long  = score_long  + 1
            score_short = score_short + 1

        # =========================
        # ENTRY: trend required + score >= threshold + anti-conflict guard
        # Anti-conflict: if both sides somehow score >= threshold, take the stronger one only.
        # =========================
        long_signal  = trend_long  and score_long  >= score_threshold and score_long  > score_short
        short_signal = trend_short and score_short >= score_threshold and score_short > score_long

        enter_long  = long_signal  and not in_long and not in_short
        enter_short = short_signal and not in_long and not in_short

        if enter_long:
            in_long     = True
            entry_price = close[0]
            entry_bar   = self.bar_index
            long_stop   = close[0] - atr_val * atr_mult

        if enter_short:
            in_short    = True
            entry_price = close[0]
            entry_bar   = self.bar_index
            short_stop  = close[0] + atr_val * atr_mult

        # =========================
        # ATR RISK ENGINE: trailing stop + optional breakeven
        # Pre-declare profit/trail before if-blocks (Indie block-level scoping)
        # =========================
        profit: float = 0.0
        trail: float  = 0.0

        if in_long and not isnan(entry_price):
            profit = close[0] - entry_price
            if use_breakeven and profit > atr_val and long_stop < entry_price:
                long_stop = entry_price          # promote stop to breakeven
            trail     = close[0] - atr_val * atr_mult
            long_stop = max(long_stop, trail)    # ratchet upward

        if in_short and not isnan(entry_price):
            profit = entry_price - close[0]
            if use_breakeven and profit > atr_val and short_stop > entry_price:
                short_stop = entry_price         # promote stop to breakeven
            trail      = close[0] + atr_val * atr_mult
            short_stop = min(short_stop, trail)  # ratchet downward

        # =========================
        # EXIT: stop hit or max bars elapsed
        # bars_in_trade is gated by in_long / in_short so entry_bar=0 harmless when flat
        # =========================
        bars_in_trade = self.bar_index - entry_bar
        time_exit     = bars_in_trade > max_bars_trade

        exit_long  = in_long  and (close[0] < long_stop  or time_exit)
        exit_short = in_short and (close[0] > short_stop or time_exit)

        # Capture exit booleans BEFORE resetting state (needed for exit markers below)
        show_long_exit  = exit_long
        show_short_exit = exit_short

        if exit_long:
            in_long     = False
            entry_price = nan
            entry_bar   = 0
            long_stop   = nan

        if exit_short:
            in_short    = False
            entry_price = nan
            entry_bar   = 0
            short_stop  = nan

        # =========================
        # Persist state for next bar
        # =========================
        _in_long.set(in_long)
        _in_short.set(in_short)
        _entry_price.set(entry_price)
        _entry_bar.set(entry_bar)
        _long_stop.set(long_stop)
        _short_stop.set(short_stop)

        # =========================
        # BUILD OUTPUTS
        # =========================

        # Entry markers with score text e.g. "4/7"
        score_text_long  = str(score_long)  + '/' + str(max_score)
        score_text_short = str(score_short) + '/' + str(max_score)

        long_entry_marker  = plot.Marker(close[0], text=score_text_long)  if enter_long  else plot.Marker(nan)
        short_entry_marker = plot.Marker(close[0], text=score_text_short) if enter_short else plot.Marker(nan)

        # Exit markers
        long_exit_marker  = plot.Marker(close[0], text='EXIT') if show_long_exit  else plot.Marker(nan)
        short_exit_marker = plot.Marker(close[0], text='EXIT') if show_short_exit else plot.Marker(nan)

        # Stop lines — visible only while in a trade
        long_stop_plot  = long_stop  if in_long  else nan
        short_stop_plot = short_stop if in_short else nan

        return (
            fast_ma[0],           # [0] Fast MA (HMA)          — orange line
            mid_ma[0],            # [1] Mid MA  (EMA)          — blue line
            slow_ma[0],           # [2] Slow MA (EMA)          — red line
            long_entry_marker,    # [3] green circle below bar — long entry  (text="score/max")
            short_entry_marker,   # [4] red cross above bar    — short entry (text="score/max")
            long_exit_marker,     # [5] red cross above bar    — long exit   (text="EXIT")
            short_exit_marker,    # [6] green circle below bar — short exit  (text="EXIT")
            long_stop_plot,       # [7] gray dashed line       — long stop
            short_stop_plot,      # [8] gray dashed line       — short stop
        )
