"""ML Predictor - orchestrates feature engineering + ensemble prediction.

v2: Integrates regime detection, multi-timeframe confirmation,
    direction classifier, risk management, and model persistence.

Wraps as a BaseAlgorithm so it integrates into the engine pipeline.
"""
from __future__ import annotations
import logging
import numpy as np
from app.engine.base import BaseAlgorithm
from app.engine.context import AlgorithmContext
from app.engine.ml.feature_eng import FeatureEngineer
from app.engine.ml.ensemble import EnsemblePredictor
from app.engine.ml.regime import RegimeDetector

logger = logging.getLogger(__name__)

# Timeframe to minutes mapping
TF_MINUTES = {
    '1m': 1, '15m': 15, '30m': 30, '1h': 60,
    '4h': 240, '1D': 1440, '1W': 10080,
}


def _cfg(key: str, fallback):
    try:
        from flask import current_app
        return current_app.config.get(key, fallback)
    except RuntimeError:
        return fallback


class MLPredictorAlgorithm(BaseAlgorithm):
    algorithm_id = 'ml.predictor'
    version = '2.0.0'
    category = 'ml'
    display_name = 'ML Ensemble Predictor v2'
    dependencies = ['technical.rsi', 'technical.atr', 'technical.ema']

    # Prediction horizons per timeframe (in candle steps)
    HORIZONS = {
        '1h': {'short': 6, 'medium': 24, 'long': 48},     # 6h, 24h, 48h
        '4h': {'short': 3, 'medium': 6, 'long': 18},       # 12h, 24h, 72h
        '1D': {'short': 3, 'medium': 7, 'long': 14},       # 3d, 7d, 14d
    }

    MIN_TRAINING_SAMPLES = 100  # Raised from 50 for better CV

    def __init__(self):
        self._ensembles: dict[tuple, EnsemblePredictor] = {}
        self._regime_detector = RegimeDetector()

    def _get_ensemble(self, asset_id: str, timeframe: str) -> EnsemblePredictor:
        key = (asset_id, timeframe)
        if key not in self._ensembles:
            self._ensembles[key] = EnsemblePredictor()

            # Try loading from cache
            try:
                loaded = self._ensembles[key].load_models(asset_id, timeframe)
                if loaded:
                    logger.info(f'Loaded cached models for {asset_id}/{timeframe}')
            except Exception as e:
                logger.debug(f'No cached models for {asset_id}/{timeframe}: {e}')

        return self._ensembles[key]

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

        # Batch scan mode: skip heavy ML training to keep scan fast.
        # Set ctx.batch_scan = True from scanner to enable.
        batch_scan = getattr(ctx, 'batch_scan', False)

        try:
            ensemble = self._get_ensemble(ctx.asset_id, ctx.timeframe)

            # Feature engineering (v2 with microstructure, wavelet, Fourier)
            tf_min = TF_MINUTES.get(ctx.timeframe, 60)
            fe = FeatureEngineer(timeframe_minutes=tf_min)
            feature_df = fe.generate(df)

            if len(feature_df) < self.MIN_TRAINING_SAMPLES:
                return ctx

            # v2.1: Add cross-asset correlation features (skip in batch for speed)
            if not batch_scan:
                feature_df = self._add_cross_asset_features(feature_df, ctx.asset_id)

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

            # Regime detection (lightweight, always run)
            regime_result = self._regime_detector.detect(df)
            ctx.ml_predictions['regime'] = regime_result

            # Train if needed (with feature selection enabled)
            if ensemble.needs_training():
                # Batch scan: skip only Optuna hyperopt (slow, optional)
                # but keep full training + feature selection for accuracy
                if batch_scan:
                    use_hyperopt = False
                else:
                    use_hyperopt = len(feature_df) >= 300 and _cfg('ML_OPTUNA_TRIALS', 50) > 0

                ensemble.train(
                    feature_df, closes,
                    use_hyperopt=use_hyperopt,
                    use_feature_selection=True,
                )
                # Save to cache after training
                try:
                    ensemble.save_models(ctx.asset_id, ctx.timeframe)
                except Exception as e:
                    logger.warning(f'Failed to save model cache: {e}')

            # Predict — cap n_steps to configured ML_PREDICT_STEPS
            horizons = self.HORIZONS.get(ctx.timeframe, self.HORIZONS['1h'])
            n_steps = horizons['long']
            max_steps = _cfg('ML_PREDICT_STEPS', 48)
            n_steps = min(n_steps, max_steps)

            predictions = ensemble.predict(feature_df, closes, n_steps=n_steps)

            # Store raw predictions
            ctx.ml_predictions['ensemble'] = predictions
            model_names = list(predictions.get('model_predictions', {}).keys())
            ctx.ml_predictions['model_names'] = model_names or ['GradientBoosting', 'XGBoost', 'LightGBM', 'CatBoost']

            # Direction classifier result
            direction = predictions.get('direction', {})
            ctx.ml_predictions['direction'] = direction

            # Confidence interval
            ctx.ml_predictions['confidence_interval'] = predictions.get('confidence_interval', {})

            # Training metrics
            ctx.ml_predictions['training_metrics'] = predictions.get('training_metrics', {})

            # Ensemble weights
            ctx.ml_predictions['ensemble_weights'] = predictions.get('ensemble_weights', {})

            # Compute horizon summaries
            pred_prices = predictions['predicted_prices']
            current = ctx.current_price

            horizon_results = {}
            for label, steps in horizons.items():
                if steps <= len(pred_prices):
                    target_price = pred_prices[steps - 1]
                    pct_change = (target_price - current) / current * 100
                    horizon_results[label] = {
                        'steps': steps,
                        'predicted_price': ctx.round_price(target_price),
                        'pct_change': round(pct_change, 2),
                        'direction': 'up' if pct_change > 0 else 'down',
                    }

            ctx.ml_predictions['horizons'] = horizon_results
            ctx.ml_predictions['volatility'] = predictions['volatility']

            # Trend direction from predictions
            pred_returns = predictions['predicted_returns']
            short_steps = horizons.get('short', 6)
            avg_return = np.mean(pred_returns[:short_steps]) if pred_returns else 0
            ctx.ml_predictions['short_term_trend'] = 'bullish' if avg_return > 0 else 'bearish'

        except Exception as e:
            logger.error(f'ML Predictor error: {e}', exc_info=True)
            ctx.errors.append({'algorithm': self.algorithm_id, 'error': str(e)})

        return ctx

    def _add_cross_asset_features(self, feature_df, asset_id: str):
        """Add cross-asset correlation features if benchmark data available."""
        try:
            from app.engine.ml.cross_asset import CrossAssetFeatures
            caf = CrossAssetFeatures()
            enriched = caf.add_features(feature_df, asset_id)
            if enriched is not None and len(enriched) > 0:
                return enriched
        except Exception as e:
            logger.debug(f'Cross-asset features skipped: {e}')
        return feature_df

    def get_signal_contribution(self, ctx: AlgorithmContext) -> dict | None:
        horizons = ctx.ml_predictions.get('horizons', {})
        if not horizons:
            return None

        # Use short-term prediction for signal
        short = horizons.get('short', {})
        medium = horizons.get('medium', {})

        if not short:
            return None

        pct = short.get('pct_change', 0)
        med_pct = medium.get('pct_change', 0) if medium else 0

        # v2: Use direction classifier confidence for weight boost
        direction = ctx.ml_predictions.get('direction', {})
        dir_confidence = direction.get('confidence', 0)
        dir_direction = direction.get('direction', 'neutral')

        # v2: Use regime for threshold adjustment
        regime = ctx.ml_predictions.get('regime', {})
        regime_action = regime.get('recommended_action', {})
        threshold_adjust = regime_action.get('threshold_adjust', 0) * 0.01  # Convert to pct

        # Adjusted thresholds based on regime
        buy_threshold = 1.5 + threshold_adjust
        sell_threshold = -(1.5 + threshold_adjust)

        # Signal logic with direction classifier confirmation
        if pct > buy_threshold and med_pct > 0.5:
            signal = 'BUY'
            base_weight = min(0.25, 0.10 + abs(pct) * 0.02)

            # Boost if direction classifier agrees
            if dir_direction == 'up' and dir_confidence > 0.3:
                base_weight = min(0.30, base_weight * (1 + dir_confidence))
            # Penalize if direction classifier disagrees
            elif dir_direction == 'down' and dir_confidence > 0.3:
                base_weight *= 0.5

            reason = (
                f"ML predicts +{pct:.1f}% (short), +{med_pct:.1f}% (medium)"
                f" | Dir: {dir_direction} ({dir_confidence:.0%})"
                f" | Regime: {regime.get('regime', '?')}"
            )

        elif pct < sell_threshold and med_pct < -0.5:
            signal = 'SELL'
            base_weight = min(0.25, 0.10 + abs(pct) * 0.02)

            if dir_direction == 'down' and dir_confidence > 0.3:
                base_weight = min(0.30, base_weight * (1 + dir_confidence))
            elif dir_direction == 'up' and dir_confidence > 0.3:
                base_weight *= 0.5

            reason = (
                f"ML predicts {pct:.1f}% (short), {med_pct:.1f}% (medium)"
                f" | Dir: {dir_direction} ({dir_confidence:.0%})"
                f" | Regime: {regime.get('regime', '?')}"
            )

        else:
            signal = 'HOLD'
            base_weight = 0.05
            reason = (
                f"ML inconclusive: {pct:+.1f}% short, {med_pct:+.1f}% medium"
                f" | Dir: {dir_direction} ({dir_confidence:.0%})"
            )

        return {'signal': signal, 'weight': round(base_weight, 3), 'reason': reason}
