# Ported to Indie from https://www.tradingview.com/script/IYL88A1N-Trendlines-with-Breaks-LuxAlgo/ created by @LuxAlgo

# This source code is licensed under a Attribution-NonCommercial-ShareAlike 4.0 International (CC BY-NC-SA 4.0) 
# To view a copy of this license, visit https://creativecommons.org/licenses/by-nc-sa/4.0/ 

# indie:lang_version = 5
from math import isnan, nan
from indie import indicator, param, color, plot, MainContext, Var, MutSeriesF, MutSeries
from indie.drawings import LineSegment, AbsolutePosition, extend_type, line_segment_style
from indie.algorithms import Atr, StdDev, Corr
from indie.math import divide


@indicator('Trendlines with Breaks', overlay_main_pane=True)
@param.int('length', default=14, min=1, title='Swing Detection Lookback')
@param.float('mult', default=1.0, min=0.0, step=0.1, title='Slope')
@param.str('calc_method', default='Atr', options=['Atr', 'Stdev', 'Linreg'], title='Slope Calculation Method')
@param.bool('backpaint', default=True, title='Backpainting')
@param.bool('show_ext', default=True, title='Show Extended Lines')
@plot.line(color=color.TEAL, title='Upper')
@plot.line(color=color.RED, title='Lower')
@plot.marker('Upper Break', color=color.TEAL, position=plot.marker_position.BELOW, text='B', style=plot.marker_style.LABEL, size=6)
@plot.marker('Lower Break', color=color.RED, position=plot.marker_position.ABOVE, text='B', style=plot.marker_style.LABEL, size=6)
class Main(MainContext):
    def __init__(self):
        # Drawing objects
        self._uptl = LineSegment(AbsolutePosition(0, 0), AbsolutePosition(0, 0), color=color.TEAL,
                                 extend_type=extend_type.RIGHT, line_style=line_segment_style.DASHED)
        self._dntl = LineSegment(AbsolutePosition(0, 0), AbsolutePosition(0, 0), color=color.RED,
                                 extend_type=extend_type.RIGHT, line_style=line_segment_style.DASHED)

    def pivot_high(self, length_left: int, length_right: int) -> float:
        """Find pivot high"""
        if self.bar_count < length_left + length_right + 1:
            return nan

        center_idx = length_right
        center_high = self.high[center_idx]

        # Check left side
        for i in range(length_left):
            if self.high[center_idx + i + 1] >= center_high:
                return nan

        # Check right side
        for i in range(length_right):
            if self.high[i] >= center_high:
                return nan

        return center_high

    def pivot_low(self, length_left: int, length_right: int) -> float:
        """Find pivot low"""
        if self.bar_count < length_left + length_right + 1:
            return nan

        center_idx = length_right
        center_low = self.low[center_idx]

        # Check left side
        for i in range(length_left):
            if self.low[center_idx + i + 1] <= center_low:
                return nan

        # Check right side
        for i in range(length_right):
            if self.low[i] <= center_low:
                return nan

        return center_low

    def calculate_slope(self, length: int, mult: float, calc_method: str) -> float:
        if calc_method == 'Atr':
            atr = Atr.new(length)
            return atr[0] / length * mult
        elif calc_method == 'Stdev':
            stdev = StdDev.new(self.close, length)
            return stdev[0] / length * mult
        else:
            # Linear regression slope calculation
            y = self.close
            x = MutSeriesF.new(self.bar_index)
            dev_x = StdDev.new(x, length)[0]
            dev_y = StdDev.new(y, length)[0]
            corr = Corr.new(x, y, length)[0]
            slope = corr * divide(dev_y, dev_x, 0)
            return abs(slope) / 2.0 * mult
        return 0.0

    def calc(self, length, mult, calc_method, backpaint, show_ext):
        # Find pivots
        ph = self.pivot_high(length, length)
        pl = self.pivot_low(length, length)
        has_ph = not isnan(ph)
        has_pl = not isnan(pl)

        # Calculate slope
        slope = self.calculate_slope(length, mult, calc_method)

        # Update slopes when pivots are found
        slope_ph = Var[float].new(0)
        slope_pl = Var[float].new(0)
        if has_ph:
            slope_ph.set(slope)
        if has_pl:
            slope_pl.set(slope)

        # Calculate trendlines
        upper = Var[float].new(0)
        lower = Var[float].new(0)
        upper.set(ph if has_ph else upper.get() - slope_ph.get())
        lower.set(pl if has_pl else lower.get() + slope_pl.get())

        # Breakout detection levels
        upper_break_level = upper.get() - slope_ph.get() * length
        lower_break_level = lower.get() + slope_pl.get() * length

        # Update breakout states
        upos = MutSeries[int].new(init=0)
        dnos = MutSeries[int].new(init=0)
        upos[0] = 0 if has_ph else 1 if self.close[0] > upper_break_level else upos[0]
        dnos[0] = 0 if has_pl else 1 if self.close[0] < lower_break_level else dnos[0]

        offset = length if backpaint else 0

        # Extended lines drawing
        start_x = self.time[offset]
        end_x = 0.0
        if offset > 0:
            end_x = self.time[offset - 1]
        else:
            end_x = self.time[0] + (self.time[0] - self.time[1])

        if has_ph and show_ext:
            start_y = ph if backpaint else upper_break_level
            end_y = (ph - slope) if backpaint else upper_break_level - slope_ph.get()

            self._uptl.point_a = AbsolutePosition(start_x, start_y)
            self._uptl.point_b = AbsolutePosition(end_x, end_y)

            self.chart.draw(self._uptl)

        if has_pl and show_ext:
            start_y = pl if backpaint else lower_break_level
            end_y = (pl + slope) if backpaint else lower_break_level + slope_pl.get()

            self._dntl.point_a = AbsolutePosition(start_x, start_y)
            self._dntl.point_b = AbsolutePosition(end_x, end_y)

            self.chart.draw(self._dntl)

        # Create Line objects with offset
        upper_plot = plot.Line(
            upper.get() if backpaint else upper_break_level,
            color=color.BLACK(0) if has_ph else None, offset=-offset)
        lower_plot = plot.Line(
            lower.get() if backpaint else lower_break_level,
            color=color.BLACK(0) if has_pl else None, offset=-offset)

        # Breakout markers
        upper_break_marker = plot.Marker(self.low[0] if upos[0] > upos[1] else nan)
        lower_break_marker = plot.Marker(self.high[0] if dnos[0] > dnos[1] else nan)

        return upper_plot, lower_plot, upper_break_marker, lower_break_marker
