"""Market Regime Detection.

Detects current market regime to adapt model behavior and thresholds:
- trending_up: strong uptrend, momentum strategies work best
- trending_down: strong downtrend, avoid buying
- ranging: sideways, mean-reversion strategies work best
- volatile: high uncertainty, reduce position sizes

Uses: Hurst exponent, ADX, volatility percentile ranking.
"""
from __future__ import annotations
import logging
import numpy as np
import pandas as pd

logger = logging.getLogger(__name__)


class RegimeDetector:
    """Detect market regime from OHLCV data."""

    def detect(self, df: pd.DataFrame) -> dict:
        """Detect current market regime.

        Args:
            df: OHLCV DataFrame with columns: open, high, low, close, volume

        Returns:
            dict with regime, confidence, metrics
        """
        if df is None or len(df) < 50:
            return self._default_regime()

        closes = df['close'].values.astype(float)

        hurst = self._hurst_exponent(closes)
        adx = self._compute_adx(df)
        vol_rank = self._volatility_rank(closes)
        trend_strength = self._trend_strength(closes)

        # Regime classification logic
        regime = self._classify_regime(hurst, adx, vol_rank, trend_strength)

        return {
            'regime': regime,
            'hurst_exponent': round(hurst, 4),
            'adx': round(adx, 2),
            'volatility_rank': round(vol_rank, 2),
            'trend_strength': round(trend_strength, 4),
            'description': self._regime_description(regime),
            'recommended_action': self._regime_action(regime),
        }

    def _hurst_exponent(self, series: np.ndarray, max_lag: int = 20) -> float:
        """Compute Hurst exponent to detect mean-reversion vs trending.

        H < 0.5: mean-reverting (ranging)
        H ≈ 0.5: random walk
        H > 0.5: trending
        """
        if len(series) < max_lag * 2:
            return 0.5

        lags = range(2, max_lag + 1)
        tau = []
        for lag in lags:
            diffs = series[lag:] - series[:-lag]
            std = np.std(diffs)
            if std > 0:
                tau.append(std)
            else:
                tau.append(1e-10)

        if len(tau) < 2:
            return 0.5

        log_lags = np.log(list(lags[:len(tau)]))
        log_tau = np.log(tau)

        # Linear regression
        try:
            coeffs = np.polyfit(log_lags, log_tau, 1)
            return float(coeffs[0])
        except Exception:
            return 0.5

    def _compute_adx(self, df: pd.DataFrame, period: int = 14) -> float:
        """Compute Average Directional Index."""
        if len(df) < period * 2:
            return 25.0

        high = df['high'].values.astype(float)
        low = df['low'].values.astype(float)
        close = df['close'].values.astype(float)

        # True Range
        tr = np.maximum(
            high[1:] - low[1:],
            np.maximum(
                np.abs(high[1:] - close[:-1]),
                np.abs(low[1:] - close[:-1])
            )
        )

        # +DM and -DM
        plus_dm = np.maximum(high[1:] - high[:-1], 0)
        minus_dm = np.maximum(low[:-1] - low[1:], 0)

        # Zero out where one is larger
        mask = plus_dm > minus_dm
        minus_dm[mask] = 0
        plus_dm[~mask] = 0

        # Smoothed averages
        atr = self._ema(tr, period)
        plus_di = 100 * self._ema(plus_dm, period) / np.maximum(atr, 1e-10)
        minus_di = 100 * self._ema(minus_dm, period) / np.maximum(atr, 1e-10)

        dx = 100 * np.abs(plus_di - minus_di) / np.maximum(plus_di + minus_di, 1e-10)
        adx = self._ema(dx, period)

        return float(adx[-1]) if len(adx) > 0 else 25.0

    def _ema(self, data: np.ndarray, period: int) -> np.ndarray:
        """Simple EMA computation."""
        result = np.zeros_like(data, dtype=float)
        if len(data) == 0:
            return result
        result[0] = data[0]
        alpha = 2.0 / (period + 1)
        for i in range(1, len(data)):
            result[i] = alpha * data[i] + (1 - alpha) * result[i - 1]
        return result

    def _volatility_rank(self, closes: np.ndarray, lookback: int = 20) -> float:
        """Rank current volatility vs historical (0-100 percentile)."""
        if len(closes) < lookback * 3:
            return 50.0

        returns = np.diff(closes) / closes[:-1]
        current_vol = np.std(returns[-lookback:])

        # Historical volatility windows
        vols = []
        for i in range(lookback, len(returns), lookback):
            window = returns[max(0, i - lookback):i]
            if len(window) >= lookback // 2:
                vols.append(np.std(window))

        if not vols:
            return 50.0

        # Percentile rank
        rank = sum(1 for v in vols if current_vol > v) / len(vols) * 100
        return float(rank)

    def _trend_strength(self, closes: np.ndarray, period: int = 20) -> float:
        """Measure trend strength via linear regression slope."""
        if len(closes) < period:
            return 0.0

        recent = closes[-period:]
        x = np.arange(period)
        try:
            coeffs = np.polyfit(x, recent, 1)
            slope = coeffs[0]
            # Normalize by price level
            return float(slope / np.mean(recent))
        except Exception:
            return 0.0

    def _classify_regime(
        self, hurst: float, adx: float, vol_rank: float, trend: float
    ) -> str:
        """Classify market regime from indicators."""
        # High volatility overrides everything
        if vol_rank > 85:
            return 'volatile'

        # Strong trend
        if adx > 30 and hurst > 0.55:
            if trend > 0:
                return 'trending_up'
            else:
                return 'trending_down'

        # Mean-reverting / ranging
        if hurst < 0.45 or adx < 20:
            return 'ranging'

        # Moderate trend
        if adx > 20 and trend > 0.001:
            return 'trending_up'
        elif adx > 20 and trend < -0.001:
            return 'trending_down'

        return 'ranging'

    def _regime_description(self, regime: str) -> str:
        descriptions = {
            'trending_up': 'Strong uptrend — momentum strategies favored',
            'trending_down': 'Strong downtrend — avoid new long positions',
            'ranging': 'Sideways market — mean-reversion strategies favored',
            'volatile': 'High volatility — reduce position sizes, widen stops',
        }
        return descriptions.get(regime, 'Unknown regime')

    def _regime_action(self, regime: str) -> dict:
        """Recommended adjustments per regime."""
        actions = {
            'trending_up': {
                'signal_bias': 'buy',
                'threshold_adjust': -5,    # Lower buy threshold
                'position_scale': 1.0,
                'stop_multiplier': 1.2,    # Wider stops in trends
            },
            'trending_down': {
                'signal_bias': 'sell',
                'threshold_adjust': 5,     # Higher buy threshold
                'position_scale': 0.5,     # Smaller positions
                'stop_multiplier': 1.0,
            },
            'ranging': {
                'signal_bias': 'neutral',
                'threshold_adjust': 10,    # Higher thresholds (need more conviction)
                'position_scale': 0.7,
                'stop_multiplier': 0.8,    # Tighter stops in range
            },
            'volatile': {
                'signal_bias': 'neutral',
                'threshold_adjust': 15,    # Much higher thresholds
                'position_scale': 0.3,     # Much smaller positions
                'stop_multiplier': 1.5,    # Much wider stops
            },
        }
        return actions.get(regime, actions['ranging'])

    def _default_regime(self) -> dict:
        return {
            'regime': 'ranging',
            'hurst_exponent': 0.5,
            'adx': 25.0,
            'volatility_rank': 50.0,
            'trend_strength': 0.0,
            'description': 'Insufficient data for regime detection',
            'recommended_action': self._regime_action('ranging'),
        }
