"""Mean Reversion Detection - identifies overextended moves likely to snap back.

Uses z-score and deviation from moving averages to detect mean reversion setups.
"""
from __future__ import annotations
import warnings
import numpy as np
from app.engine.base import BaseAlgorithm
from app.engine.context import AlgorithmContext


class MeanReversionAlgorithm(BaseAlgorithm):
    algorithm_id = 'velocity.mean_reversion'
    version = '1.0.0'
    category = 'velocity'
    display_name = 'Mean Reversion'
    dependencies = ['technical.sma', 'technical.bollinger']

    def __init__(self, zscore_threshold: float = 2.0, lookback: int = 20):
        self.zscore_threshold = zscore_threshold
        self.lookback = lookback

    def compute(self, ctx: AlgorithmContext) -> AlgorithmContext:
        df = ctx.ohlcv
        if df is None or len(df) < self.lookback:
            return ctx

        prices = df['close'].tail(self.lookback).values
        price = ctx.current_price

        # Z-score: how many standard deviations from mean
        mean_price = float(np.mean(prices))
        std_price = float(np.std(prices))
        zscore = (price - mean_price) / std_price if std_price > 0 else 0

        # Deviation from SMA 20
        sma20 = ctx.indicators.get('sma_20_latest')
        sma_dev_pct = ((price - sma20) / sma20 * 100) if sma20 and sma20 > 0 else 0

        # Bollinger Band position
        bb_upper = ctx.indicators.get('bb_upper_latest', 0)
        bb_lower = ctx.indicators.get('bb_lower_latest', 0)
        bb_mid = ctx.indicators.get('bb_middle_latest', 0)
        bb_width = bb_upper - bb_lower if bb_upper and bb_lower else 0
        bb_position = ((price - bb_lower) / bb_width * 100) if bb_width > 0 else 50

        # Half-life estimation (how fast price reverts to mean)
        # Using autocorrelation lag-1
        if len(prices) > 5:
            returns = np.diff(prices) / prices[:-1]
            if len(returns) > 1:
                with warnings.catch_warnings():
                    warnings.simplefilter('ignore', RuntimeWarning)
                    lag1_corr = float(np.corrcoef(returns[:-1], returns[1:])[0, 1])
                if np.isnan(lag1_corr) or np.isinf(lag1_corr):
                    lag1_corr = 0
                if lag1_corr < 0:
                    half_life = -np.log(2) / np.log(abs(lag1_corr)) if lag1_corr != 0 else 999
                else:
                    half_life = 999  # trending, no mean reversion
            else:
                lag1_corr = 0
                half_life = 999
        else:
            lag1_corr = 0
            half_life = 999

        # Mean reversion signal
        if zscore < -self.zscore_threshold:
            mr_signal = 'OVERSOLD'
            reversion_target = mean_price
            expected_move_pct = (mean_price - price) / price * 100
        elif zscore > self.zscore_threshold:
            mr_signal = 'OVERBOUGHT'
            reversion_target = mean_price
            expected_move_pct = (mean_price - price) / price * 100
        else:
            mr_signal = 'NEUTRAL'
            reversion_target = mean_price
            expected_move_pct = (mean_price - price) / price * 100

        # Confidence in mean reversion
        is_mean_reverting = lag1_corr < -0.1 and half_life < 20
        confidence = 'HIGH' if is_mean_reverting and abs(zscore) > 2 else \
                     'MEDIUM' if abs(zscore) > 1.5 else 'LOW'

        ctx.velocity['mean_reversion'] = {
            'zscore': round(zscore, 3),
            'mean_price': ctx.round_price(mean_price),
            'std_dev': round(std_price, 2),
            'sma_deviation_pct': round(sma_dev_pct, 2),
            'bb_position': round(bb_position, 1),
            'signal': mr_signal,
            'reversion_target': ctx.round_price(reversion_target),
            'expected_move_pct': round(expected_move_pct, 2),
            'lag1_autocorrelation': round(lag1_corr, 4),
            'half_life_candles': round(min(half_life, 999), 1),
            'is_mean_reverting': bool(is_mean_reverting),
            'confidence': confidence,
        }
        return ctx

    def get_signal_contribution(self, ctx: AlgorithmContext):
        mr = ctx.velocity.get('mean_reversion', {})
        if not mr:
            return None

        signal = mr.get('signal', 'NEUTRAL')
        conf = mr.get('confidence', 'LOW')
        zscore = mr.get('zscore', 0)

        weight = 0.08 if conf == 'HIGH' else 0.05 if conf == 'MEDIUM' else 0.03

        if signal == 'OVERSOLD':
            return {'signal': 'BUY', 'weight': weight,
                    'reason': f'Mean reversion: oversold (z={zscore:.1f})'}
        elif signal == 'OVERBOUGHT':
            return {'signal': 'SELL', 'weight': weight,
                    'reason': f'Mean reversion: overbought (z={zscore:.1f})'}
        return {'signal': 'HOLD', 'weight': 0.02,
                'reason': f'Mean reversion neutral (z={zscore:.1f})'}
