# Copyright (c) 2025 @pavelmedd. All rights reserved.

# This work is licensed under the MIT License.
# For a copy, see <https://opensource.org/licenses/MIT>.

# indie:lang_version = 5
from indie import indicator, MainContext, plot, color, param, MutSeriesF
from indie.algorithms import LinReg, StdDev, Corr, Ema, Atr, Highest, Lowest, Sma
from indie.drawings import LineSegment, LabelAbs, AbsolutePosition, line_segment_style
from indie.color import rgba
from math import isnan, pow

@indicator('UFF Pro', overlay_main_pane=True)
@param.str('mode', default='Moderate', options=['Conservative', 'Moderate', 'Aggressive'], title='Style')
@param.str('asset_class', default='Stocks', options=['Crypto', 'Forex', 'Stocks'], title='Asset Class')
@param.str('proj_type', default='Linear', options=['Linear', 'Flat', 'Parabolic'], title='Projection Type')
@param.bool('use_custom', default=False, title='Use Custom Params')
@param.int('length', default=50, min=10, title='Length')
@param.float('mult', default=2.0, min=0.1, title='Multiplier')
@param.int('fwd', default=20, min=1, title='Forecast')
@param.int('smooth', default=5, min=1, max=20, title='Slope Smoothing')
@param.bool('show_warnings', default=True, title='Show Warnings')
@param.bool('show_mae', default=True, title='Show MAE')
@param.int('curve_segments', default=5, min=2, max=10, title='Curve Segments')
@plot.line('basis', title='Basis')
@plot.line('upper', title='Upper')
@plot.line('lower', title='Lower')
@plot.fill('upper', 'lower', title='Fill')
class Main(MainContext):
    def __init__(self, mode, asset_class, proj_type, use_custom, length, mult, fwd, smooth, show_warnings, show_mae, curve_segments):
        self._mode = mode
        self._asset_class = asset_class
        self._proj_type = proj_type
        self._use_custom = use_custom
        self._length_in = length
        self._mult_in = mult
        self._fwd = fwd
        self._smooth = smooth
        self._show_warnings = show_warnings
        self._show_mae = show_mae
        self._curve_segments = curve_segments
        
        self._proj_label = LabelAbs('', AbsolutePosition(0, 0))
        
        self._lines: list[LineSegment] = []
        for _ in range(30):
            self._lines.append(LineSegment(AbsolutePosition(0, 0), AbsolutePosition(0, 0)))

    def calc(self):
        # Determine multiplier based on trading style
        mode_mult = 2.0
        if self._mode == 'Conservative':
            mode_mult = 1.5
        elif self._mode == 'Aggressive':
            mode_mult = 2.5

        # Default length depending on asset class
        default_len = 50
        if self._asset_class == 'Stocks':
            default_len = 200
        elif self._asset_class == 'Forex':
            default_len = 100
        
        length = self._length_in if self._use_custom else default_len
        mult = self._mult_in if self._use_custom else mode_mult

        # Main regression line (basis)
        lin = LinReg.new(self.close, length, 0)
        basis = lin[0]

        # Volatility-based funnel width
        std = StdDev.new(self.close, length)
        dev_base = std[0] * mult
        
        atr_len = max(10, min(50, int(length / 3)))
        atr = Atr.new(atr_len)
        min_atr = Lowest.new(atr, atr_len)[0]
        max_atr = Highest.new(atr, atr_len)[0]
        
        atr_range = max_atr - min_atr
        norm = 0.0
        if atr_range > 1e-9:
            norm = (atr[0] - min_atr) / atr_range
        
        vol_factor = 0.8 + norm * 0.8
        dev = dev_base * vol_factor

        # Upper and lower forecast bands
        upper = basis + dev
        lower = basis - dev

        # Trend confidence (R²)
        r_val = Corr.new(self.close, lin, length)[0]
        r_val = r_val if not isnan(r_val) else 0.0
        r2 = r_val * r_val
        
        # Mean Absolute Error % for display
        basis_safe = max(1e-9, abs(basis))
        mae_pct_raw = abs(self.close[0] - basis) / basis_safe * 100
        mae_pct_series = MutSeriesF.new(reset=mae_pct_raw)
        mae_pct_sma = Sma.new(mae_pct_series, length)[0]
        mae_pct_sma = mae_pct_sma if not isnan(mae_pct_sma) else 0.0
        
        # MAE minimum threshold = ATR * 0.1 (in percent)
        atr_pct = (atr[0] / basis_safe) * 100 * 0.1
        mae_pct = max(mae_pct_sma, atr_pct)
        
        # Slope and curvature for projection
        slope_raw = lin[0] - lin[1]
        slope_series = MutSeriesF.new(reset=slope_raw)
        slope_ema = Ema.new(slope_series, self._smooth)
        slope_smooth = slope_ema[0]
        
        prev = slope_ema[1]
        if isnan(prev):
            prev = slope_smooth
        curv_raw = slope_smooth - prev
        
        curv_series = MutSeriesF.new(reset=curv_raw)
        curv_smooth_len = self._smooth * 2
        curv_ema = Ema.new(curv_series, curv_smooth_len)
        curvature = curv_ema[0]
        curvature = curvature if not isnan(curvature) else 0.0

        # Adjust curvature weight based on confidence
        curv_weight = 1.0
        if r2 > 0.5:
            curv_weight = 1.0 + (r2 - 0.5) * 2.0
        else:
            curv_weight = r2 * 2.0

        # Normalize curvature relative to forecast band
        curvature_adj = curvature * curv_weight
        
        curvature_norm = curvature_adj / max(1e-8, dev)
        curv_factor = curvature_norm * (length / 100.0)
        curv_factor = max(-0.25, min(0.25, curv_factor))
        
        slope = slope_smooth * (0.3 + r2 * 1.7) * (1.0 + curv_factor)
        
        is_up = slope > 0

        # Colors for trend and funnel
        col_main = rgba(0, 184, 148, 255) if is_up else rgba(255, 107, 107, 255)
        col_line = rgba(0, 184, 148, 200) if is_up else rgba(255, 107, 107, 200)
        col_fill = color.GREEN(0.2) if is_up else color.RED(0.2)
        
        if self.is_last_bar:
            tf_min = self.time_frame.to_minutes()
            tf_step = tf_min * 60
            if self.time[0] > 1e12:
                tf_step = tf_min * 60 * 1000
            
            growth_factor = 1.0 + pow(1.0 - r2, 1.3)
            base_dev = dev
            
            # Determine number of segments in forecast funnel
            n_seg = self._curve_segments
            if tf_min <= 5 and n_seg > 6:
                n_seg = 6
            if self._fwd > 200 and n_seg > 8:
                n_seg = 8
            n_seg = min(n_seg, max(1, self._fwd))
            
            funnel_exp = 2.0
            if n_seg > 4 or length < 30:
                funnel_exp = 1.5
            
            step = float(self._fwd) / n_seg
            max_step_increase = 1.5
            
            prev_t = self.time[0]
            prev_basis = basis
            prev_upper = upper
            prev_lower = lower
            prev_dev = dev
            prev_bars = 0
            
            final_basis = basis
            final_upper = upper
            final_time = self.time[0]
            
            line_idx = 0

            # Draw forecast funnel segments
            for i in range(1, n_seg + 1):
                t_step = step * i
                
                bars_ahead = int(round(t_step))
                bars_ahead = max(prev_bars + 1, bars_ahead)
                if i == n_seg:
                    bars_ahead = self._fwd
                
                curr_t = self.time[0] + bars_ahead * tf_step

                # Calculate projected basis depending on projection type
                curr_basis = basis
                if self._proj_type == 'Flat':
                    curr_basis = basis
                elif self._proj_type == 'Linear':
                    slope_adj = slope * (1.0 + curv_factor * t_step * 0.1)
                    curr_basis = basis + slope_adj * t_step
                elif self._proj_type == 'Parabolic':
                    curr_basis = basis + slope * t_step + 0.5 * curvature_adj * t_step * t_step
                
                # Forecast band expansion
                ratio_i = min(1.0, t_step / max(1, length))
                dev_i = base_dev * growth_factor * pow(ratio_i, funnel_exp)
                dev_i = max(dev, dev_i)
                dev_i = dev_i * (1.0 + abs(curvature_norm) * 0.5)
                dev_i = min(prev_dev * max_step_increase, dev_i)
                
                curr_upper = curr_basis + dev_i
                curr_lower = curr_basis - dev_i

                # Draw basis line segment
                self._lines[line_idx].point_a = AbsolutePosition(prev_t, prev_basis)
                self._lines[line_idx].point_b = AbsolutePosition(curr_t, curr_basis)
                self._lines[line_idx].color = col_main
                self._lines[line_idx].line_width = 2
                self._lines[line_idx].line_style = line_segment_style.DOTTED
                self.chart.draw(self._lines[line_idx])
                line_idx = line_idx + 1
                
                # Draw upper band segment
                self._lines[line_idx].point_a = AbsolutePosition(prev_t, prev_upper)
                self._lines[line_idx].point_b = AbsolutePosition(curr_t, curr_upper)
                self._lines[line_idx].color = col_line
                self._lines[line_idx].line_width = 1
                self._lines[line_idx].line_style = line_segment_style.SOLID
                self.chart.draw(self._lines[line_idx])
                line_idx = line_idx + 1
                
                # Draw lower band segment
                self._lines[line_idx].point_a = AbsolutePosition(prev_t, prev_lower)
                self._lines[line_idx].point_b = AbsolutePosition(curr_t, curr_lower)
                self._lines[line_idx].color = col_line
                self._lines[line_idx].line_width = 1
                self._lines[line_idx].line_style = line_segment_style.SOLID
                self.chart.draw(self._lines[line_idx])
                line_idx = line_idx + 1
                
                # Update previous points for next segment
                prev_t = curr_t
                prev_basis = curr_basis
                prev_upper = curr_upper
                prev_lower = curr_lower
                prev_dev = dev_i
                prev_bars = bars_ahead
                
                final_basis = curr_basis
                final_upper = curr_upper
                final_time = curr_t
            
            # Trend labels
            trend = 'UP' if is_up else 'DOWN'
            conf = int(r2 * 100)
            strength = 'Strong' if r2 > 0.7 else 'Medium' if r2 > 0.4 else 'Weak'
            
            warn = ''
            if self._show_warnings:
                if self._fwd > length * 4:
                    warn = warn + '\n! Horizon >> Len'
                if r2 < 0.18:
                    warn = warn + '\n! Low R2 (noisy)'
            
            mae_str = ''
            if self._show_mae:
                mae_str = '\nMAE: ' + str(round(mae_pct, 2)) + '%'

            # Display projection label on chart
            lbl_text = 'Target: ' + str(round(final_basis, 2)) + '\n' + trend + ' | Conf: ' + str(conf) + '%\n' + strength + mae_str + warn
            
            self._proj_label.text = lbl_text
            self._proj_label.position = AbsolutePosition(final_time, final_upper)
            self._proj_label.bg_color = col_main
            self._proj_label.text_color = color.WHITE
            self._proj_label.font_size = 11
            self.chart.draw(self._proj_label)
        
        # Return main lines and fill for plotting
        return (
            plot.Line(basis, color=col_main),
            plot.Line(upper, color=col_line),
            plot.Line(lower, color=col_line),
            plot.Fill(color=col_fill)
        )
