# indie:lang_version = 5
from math import isnan, nan
from indie import Optional, color, Color, Context, MainContext, Var, Algorithm, indicator, param
from indie.drawings import LineSegment, LabelAbs, AbsolutePosition, Chart, callout_position
from indie.algorithms import Sum, ZigZag


class Settings:
    def __init__(self,
                 line_color: Color,
                 dev_threshold: float = 5.0,
                 depth: int = 10,
                 extend_last: bool = True,
                 display_reversal_price: bool = True,
                 display_cumulative_volume: bool = True,
                 display_reversal_price_change: bool = True,
                 difference_price_mode: str = 'Absolute',
                 draw: bool = True,
                 allow_zig_zag_on_one_bar: bool = True):
        self.dev_threshold = dev_threshold
        self.depth = depth
        self.extend_last = extend_last
        self.display_reversal_price = display_reversal_price
        self.display_cumulative_volume = display_cumulative_volume
        self.display_reversal_price_change = display_reversal_price_change
        self.difference_price_mode = difference_price_mode
        self.draw = draw
        self.allow_zig_zag_on_one_bar = allow_zig_zag_on_one_bar
        self.line_color = line_color
        self.price_precision = 2


class Pivot:
    def __init__(self,
                 start: Optional[AbsolutePosition],
                 end: AbsolutePosition,
                 vol: float,
                 is_high: bool,
                 chart: Chart,
                 settings: Settings):
        self.ln = Optional[LineSegment]()
        self.lb = Optional[LabelAbs]()
        self.is_high = is_high
        self.vol = vol
        self.start = start
        self.end = end
        self.chart = chart

        if settings.draw and start is not None:
            self.ln = LineSegment(start.value(), end, color=settings.line_color, line_width=2)
            self.lb = make_pivot_label(is_high, end, settings)
        self.update_pivot(end, vol, settings)


    def update_pivot(self, end: AbsolutePosition,
                     vol: float, settings: Settings) -> None:
        self.end = end
        self.vol = vol
        if self.lb is not None:
            self.lb.value().position = self.end
            self.lb.value().text = price_rotation_aggregate(self.start.value().price, self.end.price,
                                                            self.vol, settings)
            self.chart.draw(self.lb.value())
        if self.ln is not None:
            self.ln.value().point_b = self.end
            self.chart.draw(self.ln.value())

    def delete(self) -> None:
        if self.ln is not None:
            self.chart.erase(self.ln.value())
        if self.lb is not None:
            self.chart.erase(self.lb.value())


class ZigZagPainter(Algorithm):
    def __init__(self, ctx: Context):
        super().__init__(ctx)
        self.last_pivot = ctx.new_var(Optional[Pivot]())
        self.sum_vol = ctx.new_var(0.0)

    def calc(self, chart: Chart, settings: Settings) -> None:
        ctx = self.ctx
        length = max(2, settings.depth // 2)
        if not isnan(ctx.volume[length]):
            self.sum_vol.set(self.sum_vol.get() + ctx.volume[length])

        new_high, upd_high, new_low, upd_low = ZigZag.new(
            length, length, settings.dev_threshold, settings.allow_zig_zag_on_one_bar)
        if new_high or upd_high:
            self._handle_zigzag_event(ctx.time[length], ctx.high[length],
                                      new_high, upd_high, True, chart, settings)
        if new_low or upd_low:
            self._handle_zigzag_event(ctx.time[length], ctx.low[length],
                                      new_low, upd_low, False, chart, settings)

        if settings.extend_last:
            zigzag_updated = new_high or upd_high or new_low or upd_low
            self._draw_extend(zigzag_updated, length, chart, settings)

    def _handle_zigzag_event(self, time: float, price: float,
                             new_pivot: bool, upd_pivot: bool, is_high: bool,
                             chart: Chart, settings: Settings) -> None:
        point = AbsolutePosition(time, price)
        if new_pivot and self.last_pivot.get() is None:
            self.last_pivot.set(Pivot(None, point, nan, is_high, chart, settings))
            self.sum_vol.set(0)
            return
        last_pivot = self.last_pivot.get().value()
        if upd_pivot:
            last_pivot.update_pivot(point, last_pivot.vol + self.sum_vol.get(), settings)
            self.sum_vol.set(0)
        elif new_pivot:
            self.last_pivot.set(Pivot(last_pivot.end, point, self.sum_vol.get(), is_high,
                                      chart, settings))
            self.sum_vol.set(0)

    def _draw_extend(self, zigzag_updated: bool, length: int,
                     chart: Chart, settings: Settings) -> None:
        ctx = self.ctx
        last_pivot = self.last_pivot.get()
        extend = Var[Optional[Pivot]].new(None)

        rem_vol = Sum.new(ctx.volume, length)[0]
        if ctx.is_last_bar and last_pivot is not None:
            is_high = not last_pivot.value().is_high
            cur_series = ctx.high if is_high else ctx.low
            end = AbsolutePosition(ctx.time[0], cur_series[0])
            extend_vol = self.sum_vol.get() + rem_vol
            if extend.get() is None:  # there was no extend before, create a new one
                extend.set(Pivot(last_pivot.value().end, end, extend_vol, is_high,
                                 chart, settings))
            elif zigzag_updated:  # existing extend is obsolete, recreate it
                extend.get().value().delete()
                extend.set(Pivot(last_pivot.value().end, end, extend_vol, is_high,
                                 chart, settings))
            else:  # update right point of the extend
                extend.get().value().update_pivot(end, extend_vol, settings)


def price_rotation_diff(start: float, end: float, settings: Settings) -> str:
    diff = end - start
    sign = '+' if diff > 0 else ''
    diff_str = ''
    if settings.difference_price_mode == 'Absolute':
        diff_str = str(round(diff, settings.price_precision))
    else:
        diff_str = str(round(diff * 100 / start, 2)) + '%'
    return '(' + sign + diff_str + ')'


def to_vol_format(v: float) -> str:
    if v > 1000000:
        return str(round(v / 1000000, 3)) + 'M'
    if v > 1000:
        return str(round(v / 1000, 3)) + 'K'
    if v == 0:
        return '0'
    return str(v)


def price_rotation_aggregate(start: float, end: float, vol: float, settings: Settings) -> str:
    s = ''
    if settings.display_reversal_price:
        s += str(round(end, settings.price_precision)) + ' '
    if settings.display_reversal_price_change:
        s += price_rotation_diff(start, end, settings) + ' '
    if settings.display_cumulative_volume:
        s += '\n' + to_vol_format(vol)
    return s


def make_pivot_label(is_high: bool, point: AbsolutePosition,
                     settings: Settings) -> Optional[LabelAbs]:
    if (not settings.display_reversal_price and
            not settings.display_reversal_price_change and
            not settings.display_cumulative_volume):
        return None
    txt_color = color.RED
    callout_pos = callout_position.BOTTOM_RIGHT  # TODO: callout_position.BOTTOM
    if is_high:
        txt_color = color.GREEN
        callout_pos = callout_position.TOP_RIGHT  # TODO: callout_position.TOP
    return LabelAbs(text='', position=point, text_color=txt_color, callout_position=callout_pos,
                    font_size=11, bg_color=color.TRANSPARENT)  # TODO: text_align=text_align.CENTER


@indicator('Zig Zag', overlay_main_pane=True)
@param.float('deviation_input', default=5.0, min=0.00001, max=100.0, step=0.5,
             title='Price deviation for reversals (%)')
@param.int('depth_input', default=10, min=2, title='Pivot legs')
@param.color('line_color', default=color.BLUE, title='Line color')
@param.bool('extend_input', default=True, title='Extend to last bar')
@param.bool('show_price_input', default=True, title='Display reversal price')
@param.bool('show_vol_input', default=True, title='Display cumulative volume')
@param.bool('show_chg_input', default=True, title='Display reversal price change')
@param.str('price_diff_input', default='Absolute', options=['Absolute', 'Percent'],
           title='Price Reversal')
class Main(MainContext):
    def __init__(self, deviation_input, depth_input, line_color, extend_input,
                 show_price_input, show_vol_input, show_chg_input, price_diff_input):
        self._settings = Settings(
            line_color, deviation_input,
            depth_input, extend_input,
            show_price_input, show_vol_input,
            show_chg_input, price_diff_input,
        )

    def pre_calc(self):
        self._settings.price_precision = self.info.price_precision

    def calc(self):
        ZigZagPainter.new(self.chart, self._settings)
