# Ported to Indie from TradingView built-in https://www.tradingview.com/support/solutions/43000729030-trading-sessions/

# indie:lang_version = 5
from math import isnan
from datetime import time
from indie import IndieError, indicator, param, color, MainContext, Context, Color, Optional, TimeFrame, time_frame_unit, Algorithm
from indie.schedule import ScheduleRule, Schedule
from indie.drawings import LineSegment, LabelAbs, AbsolutePosition, callout_position, line_segment_style, Chart


def make_schedule(time_str: str, timezone: str) -> Schedule:
    if len(time_str) != 9 or time_str[4] != '-':
        raise IndieError('time must be of pattern "hhmm-hhmm", got: ' + time_str)
    h_start = int(time_str[0]) * 10 + int(time_str[1])
    m_start = int(time_str[2]) * 10 + int(time_str[3])
    h_end = int(time_str[5]) * 10 + int(time_str[6])
    m_end = int(time_str[7]) * 10 + int(time_str[8])
    schedule_rule = ScheduleRule(start=time(hour=h_start, minute=m_start), end=time(hour=h_end, minute=m_end))
    return Schedule(rules=[schedule_rule], timezone=timezone)


class SessionDisplay:
    def __init__(self, session_color: Color):
        # Box elements (4 lines to replace box)
        self._session_box_top = LineSegment(AbsolutePosition(0, 0), AbsolutePosition(0, 0), color=session_color)
        self._session_box_bottom = LineSegment(AbsolutePosition(0, 0), AbsolutePosition(0, 0), color=session_color)
        self._session_box_left = LineSegment(AbsolutePosition(0, 0), AbsolutePosition(0, 0), color=session_color)
        self._session_box_right = LineSegment(AbsolutePosition(0, 0), AbsolutePosition(0, 0), color=session_color)

        self._session_label = LabelAbs('', AbsolutePosition(0, 0), font_size=12, text_color=session_color,
                                       callout_position=callout_position.BOTTOM_RIGHT, bg_color=color.BLACK(0.0))

        self._open_line = LineSegment(AbsolutePosition(0, 0), AbsolutePosition(0, 0),
                                      color=session_color, line_style=line_segment_style.DASHED)
        self._close_line = LineSegment(AbsolutePosition(0, 0), AbsolutePosition(0, 0),
                                       color=session_color, line_style=line_segment_style.DASHED)
        self._avg_line = LineSegment(AbsolutePosition(0, 0), AbsolutePosition(0, 0),
                                     color=session_color, line_style=line_segment_style.DOTTED, line_width=2)

        self.session_high = 0.0
        self.session_low = 0.0
        self.session_open = 0.0
        self.session_close = 0.0
        self.start_time = 0.0
        self.end_time = 0.0

    def set_name(self, name: str, sum_close: float, num_of_bars: int,
                 show_session_names: bool, show_session_tick_range: bool,
                 show_session_average: bool, tick_size: float, price_precision: int) -> None:
        box_text: list[str] = []
        if show_session_tick_range:
            tick_range = (self.session_high - self.session_low) / tick_size
            box_text.append("Range: " + str(round(tick_range, price_precision)))
        if show_session_average and num_of_bars > 0:
            avg = sum_close / num_of_bars
            box_text.append("Avg: " + str(round(avg, price_precision)))
        if show_session_names:
            box_text.append(name)

        self._session_label.position = AbsolutePosition(self.start_time, self.session_low)
        self._session_label.text = "\n".join(box_text)

    def update_box_coordinates(self) -> None:
        self._session_box_top.point_a = AbsolutePosition(self.start_time, self.session_high)
        self._session_box_top.point_b = AbsolutePosition(self.end_time, self.session_high)

        self._session_box_bottom.point_a = AbsolutePosition(self.start_time, self.session_low)
        self._session_box_bottom.point_b = AbsolutePosition(self.end_time, self.session_low)

        self._session_box_left.point_a = AbsolutePosition(self.start_time, self.session_high)
        self._session_box_left.point_b = AbsolutePosition(self.start_time, self.session_low)

        self._session_box_right.point_a = AbsolutePosition(self.end_time, self.session_high)
        self._session_box_right.point_b = AbsolutePosition(self.end_time, self.session_low)

    def update_lines(self, sum_close: float, num_of_bars: int, show_session_oc: bool, show_session_average: bool) -> None:
        if show_session_oc:
            self._open_line.point_a = AbsolutePosition(self.start_time, self.session_open)
            self._open_line.point_b = AbsolutePosition(self.end_time, self.session_open)

            self._close_line.point_a = AbsolutePosition(self.start_time, self.session_close)
            self._close_line.point_b = AbsolutePosition(self.end_time, self.session_close)

        if show_session_average and num_of_bars > 0:
            avg = sum_close / num_of_bars
            self._avg_line.point_a = AbsolutePosition(self.start_time, avg)
            self._avg_line.point_b = AbsolutePosition(self.end_time, avg)

    def draw_all(self, chart: Chart, show_session_oc: bool, show_session_average: bool) -> None:
        # Draw box
        chart.draw(self._session_box_top)
        chart.draw(self._session_box_bottom)
        chart.draw(self._session_box_left)
        chart.draw(self._session_box_right)

        # Draw lines
        if show_session_oc:
            chart.draw(self._open_line)
            chart.draw(self._close_line)

        if show_session_average:
            chart.draw(self._avg_line)

        # Draw label
        chart.draw(self._session_label)


class SessionInfo(Algorithm):
    def __init__(self, ctx: Context, session_color: Color, name: str, schedule: Schedule):
        super().__init__(ctx)
        self._color = session_color
        self._name = name
        self._schedule = schedule
        self._active = ctx.new_var(Optional[SessionDisplay]())
        self._sum_close = ctx.new_var(0.0)
        self._num_of_bars = ctx.new_var(1)

    def create_session_display(self) -> None:
        disp = SessionDisplay(self._color)
        disp.start_time = self.ctx.time[0]
        disp.end_time = self.ctx.time[0]
        disp.session_high = self.ctx.high[0]
        disp.session_low = self.ctx.low[0]
        disp.session_open = self.ctx.open[0]
        disp.session_close = self.ctx.close[0]

        self._active.set(disp)
        self._sum_close.set(self.ctx.close[0])
        self._num_of_bars.set(1)

    def update_session_display(self) -> None:
        session_disp = self._active.get().value()
        session_disp.session_high = max(session_disp.session_high, self.ctx.high[0])
        session_disp.session_low = min(session_disp.session_low, self.ctx.low[0])
        session_disp.session_close = self.ctx.close[0]
        session_disp.end_time = self.ctx.time[0]

        self._sum_close.set(self._sum_close.get() + self.ctx.close[0])
        self._num_of_bars.set(self._num_of_bars.get() + 1)

    def calc(self, chart: Chart, is_change: bool, show_session_names: bool,
               show_session_oc: bool, show_session_tick_range: bool, show_session_average: bool,
               tick_size: float, price_precision: int) -> None:
        in_session = self.ctx.time[0] in self._schedule

        if in_session:
            if self._active.get() is None or is_change:
                self.create_session_display()
            else:
                self.update_session_display()

            # Update coordinates and draw
            active_display = self._active.get().value()
            active_display.update_box_coordinates()
            active_display.update_lines(self._sum_close.get(), self._num_of_bars.get(), show_session_oc, show_session_average)
            active_display.set_name(self._name, self._sum_close.get(), self._num_of_bars.get(),
                                    show_session_names, show_session_tick_range,
                                    show_session_average, tick_size, price_precision)
            active_display.draw_all(chart, show_session_oc, show_session_average)

        elif self._active.get() is not None:
            self._active.set(None)


@indicator('3 Trading Sessions', overlay_main_pane=True)
@param.bool('show_session_names', default=True, title='Show session names')
@param.bool('show_session_oc', default=True, title='Draw session open and close lines')
@param.bool('show_session_tick_range', default=True, title='Show tick range for each session')
@param.bool('show_session_average', default=True, title='Show average price per session')
@param.bool('show_first', default=True, title='Show first session')
@param.str('first_session_name', default='Tokyo', title='First session name')
@param.str('first_session_time', default='0900-1500', title='First session time')
@param.str('first_session_tz', default='Asia/Tokyo', title='First session timezone')
@param.bool('show_second', default=True, title='Show second session')
@param.str('second_session_name', default='London', title='Second session name')
@param.str('second_session_time', default='0830-1630', title='Second session time')
@param.str('second_session_tz', default='Europe/London', title='Second session timezone')
@param.bool('show_third', default=True, title='Show third session')
@param.str('third_session_name', default='New York', title='Third session name')
@param.str('third_session_time', default='0930-1600', title='Third session time')
@param.str('third_session_tz', default='America/New_York', title='Third session timezone')
class Main(MainContext):
    def __init__(self,
                 show_first, first_session_name, first_session_time, first_session_tz,
                 show_second, second_session_name, second_session_time, second_session_tz,
                 show_third, third_session_name, third_session_time, third_session_tz):
        if self.time_frame.to_minutes() >= TimeFrame(1, time_frame_unit.DAY).to_minutes():
            raise IndieError('This indicator can only be used on intraday timeframes.')

        self._session_infos: list[SessionInfo] = []
        if show_first:
            self._session_infos.append(SessionInfo(self, color.BLUE, first_session_name,
                                                   make_schedule(first_session_time, first_session_tz)))
        if show_second:
            self._session_infos.append(SessionInfo(self, color.YELLOW, second_session_name,
                                                   make_schedule(second_session_time, second_session_tz)))
        if show_third:
            self._session_infos.append(SessionInfo(self, color.GREEN, third_session_name,
                                                   make_schedule(third_session_time, third_session_tz)))

    def calc(self, show_session_names, show_session_oc, show_session_tick_range, show_session_average):
        # Check for new day
        is_change = (
            not isnan(self.time[1]) and
            not self.trading_session.is_same_period(self.time[0], self.time[1])
        )

        # Update each session
        for info in self._session_infos:
            info.calc(self.chart, is_change, show_session_names, show_session_oc,
                      show_session_tick_range, show_session_average,
                      self.info.tick_size, self.info.price_precision)
