# indie:lang_version = 5
from math import sqrt, isnan, nan
from indie import indicator, param, source, color, MainContext, SeriesF
from indie.drawings import LineSegment, LabelAbs, AbsolutePosition, extend_type, callout_position
from indie.math import divide


def calc_slope(src: SeriesF, length: int) -> tuple[float, float, float]:
    sum_x = 0.0
    sum_y = 0.0
    sum_x_sqr = 0.0
    sum_xy = 0.0

    for i in range(length):
        val = src[i]
        per = i + 1.0
        sum_x += per
        sum_y += val
        sum_x_sqr += per * per
        sum_xy += val * per

    slope = divide(length * sum_xy - sum_x * sum_y, length * sum_x_sqr - sum_x * sum_x)
    average = sum_y / length
    intercept = average - slope * sum_x / length + slope

    return slope, average, intercept


def calc_dev(slope: float, average: float, intercept: float,
             src: SeriesF, high: SeriesF, low: SeriesF, length: int) -> tuple[float, float, float, float]:
    up_dev = 0.0
    dn_dev = 0.0
    std_dev_acc = 0.0
    dsxx = 0.0
    dsyy = 0.0
    dsxy = 0.0
    periods = length - 1
    day_y = intercept + slope * periods / 2
    val = intercept

    for j in range(length):
        price = high[j] - val
        if price > up_dev:
            up_dev = price

        price = val - low[j]
        if price > dn_dev:
            dn_dev = price

        price = src[j]
        dxt = price - average
        dyt = val - day_y
        price -= val
        std_dev_acc += price * price
        dsxx += dxt * dxt
        dsyy += dyt * dyt
        dsxy += dxt * dyt
        val += slope

    std_dev = sqrt(std_dev_acc / (periods if periods != 0 else 1))
    pearson_r = divide(dsxy, sqrt(dsxx * dsyy), 0.0)

    return std_dev, pearson_r, up_dev, dn_dev


@indicator('Linear Regression Channel', overlay_main_pane=True)
@param.int('length_input', default=100, min=1, max=5000, title='Length')
@param.source('source_input', default=source.CLOSE, title='Source')
@param.bool('use_upper_dev_input', default=True, title='Upper Deviation')
@param.float('upper_mult_input', default=2.0, title='Upper Multiplier')
@param.bool('use_lower_dev_input', default=True, title='Lower Deviation')
@param.float('lower_mult_input', default=2.0, title='Lower Multiplier')
@param.bool('show_pearson_input', default=True, title="Show Pearson's R")
@param.bool('extend_left_input', default=False, title='Extend Lines Left')
@param.bool('extend_right_input', default=True, title='Extend Lines Right')
class Main(MainContext):
    def __init__(self, extend_left_input, extend_right_input):
        # Determine extend type
        extend_style = extend_type.NONE
        if extend_left_input and extend_right_input:
            extend_style = extend_type.BOTH
        elif extend_left_input:
            extend_style = extend_type.LEFT
        elif extend_right_input:
            extend_style = extend_type.RIGHT

        self._base_line = LineSegment(
            AbsolutePosition(0, 0),
            AbsolutePosition(0, 0),
            extend_type=extend_style,
            color=color.RED,
        )
        self._upper_line = LineSegment(
            AbsolutePosition(0, 0),
            AbsolutePosition(0, 0),
            extend_type=extend_style,
            color=color.BLUE,
        )
        self._lower_line = LineSegment(
            AbsolutePosition(0, 0),
            AbsolutePosition(0, 0),
            extend_type=extend_style,
            color=color.BLUE,
        )
        self._pearson_label = LabelAbs(
            '',
            AbsolutePosition(0, 0),
            text_color=color.BLUE,
            callout_position=callout_position.BOTTOM_LEFT,
            bg_color=color.TRANSPARENT,
        )

    def calc(self,
             length_input, source_input,
             use_upper_dev_input, upper_mult_input,
             use_lower_dev_input, lower_mult_input,
             show_pearson_input):
        if not self.is_last_bar:
            return

        # Calculate regression parameters
        slope, average, intercept = calc_slope(source_input, length_input)

        if isnan(slope):
            return

        start_price = intercept + slope * (length_input - 1)
        end_price = intercept

        # Calculate deviations
        std_dev, pearson_r, up_dev, dn_dev = calc_dev(slope, average, intercept, source_input,
                                                      self.high, self.low, length_input)

        upper_start_price = start_price + (upper_mult_input * std_dev if use_upper_dev_input else up_dev)
        upper_end_price = end_price + (upper_mult_input * std_dev if use_upper_dev_input else up_dev)
        lower_start_price = start_price + (-lower_mult_input * std_dev if use_lower_dev_input else -dn_dev)
        lower_end_price = end_price + (-lower_mult_input * std_dev if use_lower_dev_input else -dn_dev)

        # Draw lines
        if not isnan(start_price) and not isnan(upper_start_price) and not isnan(lower_start_price):
            start_x = self.time[length_input - 1]
            end_x = self.time[0]

            # Update base line coordinates
            self._base_line.point_a = AbsolutePosition(start_x, start_price)
            self._base_line.point_b = AbsolutePosition(end_x, end_price)
            self.chart.draw(self._base_line)

            # Update upper line coordinates
            self._upper_line.point_a = AbsolutePosition(start_x, upper_start_price)
            self._upper_line.point_b = AbsolutePosition(end_x, upper_end_price)
            self.chart.draw(self._upper_line)

            # Update lower line coordinates
            self._lower_line.point_a = AbsolutePosition(start_x, lower_start_price)
            self._lower_line.point_b = AbsolutePosition(end_x, lower_end_price)
            self.chart.draw(self._lower_line)

            # Update Pearson's R label text and position
            if show_pearson_input and not isnan(pearson_r):
                self._pearson_label.text = str(round(pearson_r, 8))
                self._pearson_label.position = AbsolutePosition(start_x, lower_start_price)
                self.chart.draw(self._pearson_label)
