"""Cross-Asset Correlation Features.

Generates correlation features between a asset/stock and market benchmarks.
For crypto: BTC correlation, ETH correlation, market dominance
For stocks: sector correlation, market beta, VIX correlation
"""
from __future__ import annotations
import logging
import numpy as np
import pandas as pd
from typing import Optional

logger = logging.getLogger(__name__)


class CrossAssetFeatures:
    """Generate cross-asset correlation features."""

    def __init__(self):
        self._benchmark_cache: dict[str, pd.Series] = {}

    def add_features(self, df: pd.DataFrame, asset_id: str) -> pd.DataFrame:
        """Convenience wrapper: auto-loads benchmark and generates features.

        Tries to find a benchmark from DB (BTC for crypto, market index for stock).
        Falls back to self-correlation features if no benchmark found.
        """
        benchmark = None
        try:
            # Try BTC as default benchmark for crypto
            benchmark = self.load_benchmark('bitcoin', 'crypto')
            if benchmark is None:
                # Try market index
                benchmark = self.load_benchmark('NASDAQ.IXIC', 'stock')
        except Exception:
            pass

        return self.generate(df, asset_id, benchmark_closes=benchmark)

    def generate(
        self, df: pd.DataFrame, asset_id: str,
        benchmark_closes: Optional[pd.Series] = None,
        asset_type: str = 'crypto',
    ) -> pd.DataFrame:
        """Add cross-asset correlation features to DataFrame.

        Args:
            df: OHLCV DataFrame for the target asset
            asset_id: Identifier of the asset
            benchmark_closes: Close prices of benchmark (BTC for crypto, SPY for stocks)
            asset_type: 'crypto' or 'stock'

        Returns:
            DataFrame with correlation features added
        """
        df = df.copy()

        if benchmark_closes is not None and len(benchmark_closes) >= 20:
            df = self._add_correlation_features(df, benchmark_closes)
        else:
            df = self._add_self_correlation_features(df)

        return df

    def _add_correlation_features(
        self, df: pd.DataFrame, benchmark: pd.Series
    ) -> pd.DataFrame:
        """Add rolling correlation with benchmark."""
        asset_returns = df['close'].pct_change()

        # Align lengths
        min_len = min(len(asset_returns), len(benchmark))
        if min_len < 20:
            return self._add_self_correlation_features(df)

        bench_returns = benchmark.pct_change().iloc[-min_len:]
        asset_ret = asset_returns.iloc[-min_len:]

        # Rolling correlations
        for window in [10, 20, 40]:
            if min_len >= window:
                corr = asset_ret.rolling(window).corr(bench_returns)
                df[f'benchmark_corr_{window}'] = np.nan
                df.iloc[-min_len:, df.columns.get_loc(f'benchmark_corr_{window}')] = corr.values

        # Beta (rolling regression slope)
        if min_len >= 20:
            rolling_cov = asset_ret.rolling(20).cov(bench_returns)
            rolling_var = bench_returns.rolling(20).var()
            beta = rolling_cov / rolling_var.replace(0, np.nan)
            df['market_beta'] = np.nan
            df.iloc[-min_len:, df.columns.get_loc('market_beta')] = beta.values

        return df

    def _add_self_correlation_features(self, df: pd.DataFrame) -> pd.DataFrame:
        """When no benchmark available, add autocorrelation features."""
        returns = df['close'].pct_change()

        # Serial autocorrelation at different lags
        for lag in [1, 5, 10]:
            if len(returns) > lag + 10:
                autocorr = returns.rolling(20).apply(
                    lambda x: x.autocorr(lag=lag) if len(x) > lag else 0,
                    raw=False
                )
                df[f'autocorr_lag_{lag}'] = autocorr

        # Volatility clustering (GARCH-like)
        squared_returns = returns ** 2
        df['vol_clustering'] = squared_returns.rolling(10).mean()

        return df

    def load_benchmark(
        self, benchmark_id: str, asset_type: str = 'crypto'
    ) -> Optional[pd.Series]:
        """Load benchmark close prices from database.

        Args:
            benchmark_id: 'bitcoin' for crypto, index ticker for stocks
            asset_type: 'crypto' or 'stock'

        Returns:
            pandas Series of close prices, or None if not available
        """
        if benchmark_id in self._benchmark_cache:
            return self._benchmark_cache[benchmark_id]

        try:
            from app.models.ohlcv import OHLCVData
            from app.extensions import db

            query = db.session.query(
                OHLCVData.timestamp, OHLCVData.close
            ).filter(
                OHLCVData.asset_id == benchmark_id,
                OHLCVData.timeframe == '1D',
            ).order_by(OHLCVData.timestamp.asc()).all()

            if not query:
                return None

            series = pd.Series(
                [float(r.close) for r in query],
                index=[r.timestamp for r in query],
                name=benchmark_id,
            )
            self._benchmark_cache[benchmark_id] = series
            return series

        except Exception as e:
            logger.warning(f'Failed to load benchmark {benchmark_id}: {e}')
            return None
