"""
Advanced Statistical Analysis - Fractal Compass v4.0
=====================================================

Journal-tier statistical methods for uncertainty quantification,
hypothesis testing, and scientific inference.

Components:
- Bootstrap error estimation
- Bayesian parameter inference
- Hypothesis testing with multiple comparisons
- Information-theoretic measures
- Non-parametric statistics
- Time series analysis

References:
[1] Efron & Tibshirani, "An Introduction to the Bootstrap", 1993
[2] Gelman et al., "Bayesian Data Analysis", 2013
[3] Hastie et al., "The Elements of Statistical Learning", 2009
"""

import numpy as np
import scipy.stats as stats
from scipy import optimize
from scipy.signal import periodogram, welch
from scipy.stats import entropy
from typing import Dict, List, Tuple, Optional, Any, Callable
import warnings
from dataclasses import dataclass

@dataclass
class StatisticalResult:
    """Container for statistical analysis results."""
    statistic: float
    p_value: float
    confidence_interval: Tuple[float, float]
    method: str
    interpretation: str
    effect_size: Optional[float] = None

class AdvancedStatistics:
    """
    Advanced statistical methods for scientific analysis.

    Implements rigorous uncertainty quantification and hypothesis testing
    appropriate for publication in peer-reviewed journals.
    """

    def __init__(self, confidence_level: float = 0.95, n_bootstrap: int = 10000):
        self.confidence_level = confidence_level
        self.n_bootstrap = n_bootstrap
        self.alpha = 1 - confidence_level

    def bootstrap_confidence_interval(self,
                                    data: np.ndarray,
                                    statistic_func: Callable,
                                    method: str = 'percentile') -> Tuple[float, float, np.ndarray]:
        """
        Bootstrap confidence interval estimation.

        Args:
            data: Input data array
            statistic_func: Function to compute statistic
            method: 'percentile', 'bias_corrected', or 'bca'

        Returns:
            Tuple of (lower_bound, upper_bound, bootstrap_samples)
        """
        n = len(data)
        bootstrap_stats = []

        # Generate bootstrap samples
        for _ in range(self.n_bootstrap):
            bootstrap_sample = np.random.choice(data, size=n, replace=True)
            bootstrap_stat = statistic_func(bootstrap_sample)
            bootstrap_stats.append(bootstrap_stat)

        bootstrap_stats = np.array(bootstrap_stats)

        if method == 'percentile':
            # Simple percentile method
            lower = np.percentile(bootstrap_stats, 100 * self.alpha / 2)
            upper = np.percentile(bootstrap_stats, 100 * (1 - self.alpha / 2))

        elif method == 'bias_corrected':
            # Bias-corrected percentile method
            original_stat = statistic_func(data)
            bias_correction = stats.norm.ppf(np.mean(bootstrap_stats <= original_stat))

            # Adjust percentiles for bias
            alpha_1 = stats.norm.cdf(2 * bias_correction + stats.norm.ppf(self.alpha / 2))
            alpha_2 = stats.norm.cdf(2 * bias_correction + stats.norm.ppf(1 - self.alpha / 2))

            lower = np.percentile(bootstrap_stats, 100 * alpha_1)
            upper = np.percentile(bootstrap_stats, 100 * alpha_2)

        elif method == 'bca':
            # Bias-corrected and accelerated (BCa) method
            original_stat = statistic_func(data)

            # Bias correction
            bias_correction = stats.norm.ppf(np.mean(bootstrap_stats <= original_stat))

            # Acceleration via jackknife
            jackknife_stats = []
            for i in range(n):
                jackknife_sample = np.delete(data, i)
                jackknife_stat = statistic_func(jackknife_sample)
                jackknife_stats.append(jackknife_stat)

            jackknife_mean = np.mean(jackknife_stats)
            acceleration = np.sum((jackknife_mean - np.array(jackknife_stats))**3) / (
                6 * (np.sum((jackknife_mean - np.array(jackknife_stats))**2))**(3/2)
            )

            # BCa adjusted percentiles
            z_alpha_2 = stats.norm.ppf(self.alpha / 2)
            z_1_alpha_2 = stats.norm.ppf(1 - self.alpha / 2)

            alpha_1 = stats.norm.cdf(bias_correction + (bias_correction + z_alpha_2) /
                                   (1 - acceleration * (bias_correction + z_alpha_2)))
            alpha_2 = stats.norm.cdf(bias_correction + (bias_correction + z_1_alpha_2) /
                                   (1 - acceleration * (bias_correction + z_1_alpha_2)))

            lower = np.percentile(bootstrap_stats, 100 * alpha_1)
            upper = np.percentile(bootstrap_stats, 100 * alpha_2)

        return lower, upper, bootstrap_stats

    def bayesian_parameter_estimation(self,
                                    data: np.ndarray,
                                    likelihood_func: Callable,
                                    prior_params: Dict[str, Tuple[float, float]],
                                    n_samples: int = 50000) -> Dict[str, Any]:
        """
        Bayesian parameter estimation using Metropolis-Hastings MCMC.

        Args:
            data: Observed data
            likelihood_func: Likelihood function
            prior_params: Dictionary of parameter names to (mean, std) priors
            n_samples: Number of MCMC samples

        Returns:
            Dictionary with posterior samples and statistics
        """
        # Initialize parameters
        param_names = list(prior_params.keys())
        n_params = len(param_names)

        # Starting values (prior means)
        current_params = {name: params[0] for name, params in prior_params.items()}

        # MCMC chain storage
        chain = {name: [] for name in param_names}
        log_likelihood_chain = []

        # Proposal covariance (adaptive)
        proposal_cov = np.eye(n_params) * 0.1
        acceptance_count = 0

        for i in range(n_samples):
            # Propose new parameters
            param_vector = np.array([current_params[name] for name in param_names])
            proposed_vector = np.random.multivariate_normal(param_vector, proposal_cov)
            proposed_params = {name: val for name, val in zip(param_names, proposed_vector)}

            # Calculate log likelihood
            try:
                current_log_like = likelihood_func(data, current_params)
                proposed_log_like = likelihood_func(data, proposed_params)

                # Calculate log prior
                current_log_prior = sum([
                    stats.norm.logpdf(current_params[name], prior_params[name][0], prior_params[name][1])
                    for name in param_names
                ])
                proposed_log_prior = sum([
                    stats.norm.logpdf(proposed_params[name], prior_params[name][0], prior_params[name][1])
                    for name in param_names
                ])

                # Metropolis acceptance criterion
                log_acceptance_ratio = (proposed_log_like + proposed_log_prior) - (current_log_like + current_log_prior)

                if log_acceptance_ratio > np.log(np.random.random()):
                    current_params = proposed_params
                    current_log_like = proposed_log_like
                    acceptance_count += 1

            except:
                # Reject proposal if likelihood calculation fails
                pass

            # Store samples
            for name in param_names:
                chain[name].append(current_params[name])
            log_likelihood_chain.append(current_log_like)

            # Adaptive proposal tuning
            if i > 0 and i % 1000 == 0:
                acceptance_rate = acceptance_count / 1000
                if acceptance_rate < 0.2:
                    proposal_cov *= 0.8
                elif acceptance_rate > 0.5:
                    proposal_cov *= 1.2
                acceptance_count = 0

        # Calculate posterior statistics
        burn_in = n_samples // 4  # Remove first 25% as burn-in
        posterior_stats = {}

        for name in param_names:
            samples = np.array(chain[name][burn_in:])
            posterior_stats[name] = {
                'mean': np.mean(samples),
                'std': np.std(samples),
                'median': np.median(samples),
                'credible_interval': np.percentile(samples, [2.5, 97.5]),
                'effective_sample_size': self._calculate_ess(samples)
            }

        return {
            'posterior_stats': posterior_stats,
            'chains': {name: chain[name][burn_in:] for name in param_names},
            'log_likelihood': log_likelihood_chain[burn_in:],
            'acceptance_rate': acceptance_count / 1000
        }

    def _calculate_ess(self, samples: np.ndarray) -> float:
        """Calculate effective sample size using autocorrelation."""
        n = len(samples)
        if n < 10:
            return n

        # Autocorrelation function
        autocorr = np.correlate(samples - np.mean(samples), samples - np.mean(samples), mode='full')
        autocorr = autocorr[autocorr.size // 2:]
        autocorr = autocorr / autocorr[0]

        # Integrated autocorrelation time
        tau_int = 1 + 2 * np.sum(autocorr[1:autocorr.size//4])  # Sum up to lag n/4

        # Effective sample size
        return n / (2 * tau_int + 1)

    def multiple_hypothesis_correction(self,
                                     p_values: List[float],
                                     method: str = 'fdr_bh') -> Tuple[List[bool], List[float]]:
        """
        Multiple hypothesis testing correction.

        Args:
            p_values: List of uncorrected p-values
            method: 'bonferroni', 'holm', 'fdr_bh' (Benjamini-Hochberg)

        Returns:
            Tuple of (rejected_hypotheses, corrected_p_values)
        """
        p_values = np.array(p_values)
        n = len(p_values)

        if method == 'bonferroni':
            corrected_p = p_values * n
            rejected = corrected_p <= self.alpha

        elif method == 'holm':
            sorted_indices = np.argsort(p_values)
            corrected_p = np.zeros_like(p_values)
            rejected = np.zeros_like(p_values, dtype=bool)

            for i, idx in enumerate(sorted_indices):
                corrected_p[idx] = p_values[idx] * (n - i)
                rejected[idx] = corrected_p[idx] <= self.alpha
                if not rejected[idx]:
                    break  # Stop at first non-rejection

        elif method == 'fdr_bh':
            sorted_indices = np.argsort(p_values)
            corrected_p = np.zeros_like(p_values)
            rejected = np.zeros_like(p_values, dtype=bool)

            for i in range(n-1, -1, -1):
                idx = sorted_indices[i]
                corrected_p[idx] = p_values[idx] * n / (i + 1)
                rejected[idx] = corrected_p[idx] <= self.alpha
                if rejected[idx]:
                    # Accept all more significant hypotheses
                    for j in range(i):
                        rejected[sorted_indices[j]] = True
                    break

        return rejected.tolist(), corrected_p.tolist()

    def power_analysis(self,
                      effect_size: float,
                      sample_size: int,
                      alpha: float = None,
                      test_type: str = 't_test') -> Dict[str, float]:
        """
        Statistical power analysis.

        Args:
            effect_size: Cohen's d for t-test, other measures for other tests
            sample_size: Sample size for analysis
            alpha: Significance level (uses self.alpha if None)
            test_type: 't_test', 'chi_square', 'correlation'

        Returns:
            Dictionary with power analysis results
        """
        if alpha is None:
            alpha = self.alpha

        if test_type == 't_test':
            # Cohen's power analysis for t-test
            critical_t = stats.t.ppf(1 - alpha/2, df=sample_size-1)
            ncp = effect_size * np.sqrt(sample_size)  # Non-centrality parameter
            power = 1 - stats.nct.cdf(critical_t, df=sample_size-1, nc=ncp)

        elif test_type == 'correlation':
            # Power for correlation test
            z_critical = stats.norm.ppf(1 - alpha/2)
            z_effect = 0.5 * np.log((1 + effect_size) / (1 - effect_size))  # Fisher's z
            z_power = z_effect * np.sqrt(sample_size - 3) - z_critical
            power = 1 - stats.norm.cdf(z_power)

        else:
            power = np.nan

        return {
            'power': power,
            'effect_size': effect_size,
            'sample_size': sample_size,
            'alpha': alpha,
            'beta': 1 - power
        }

    def time_series_stationarity_test(self, data: np.ndarray) -> StatisticalResult:
        """
        Augmented Dickey-Fuller test for time series stationarity.

        Args:
            data: Time series data

        Returns:
            StatisticalResult with test statistics and interpretation
        """
        from scipy.stats import linregress

        n = len(data)

        # First difference
        y_diff = np.diff(data)
        y_lag = data[:-1]

        # Simple ADF test without lagged differences (for stability)
        X = y_lag[:-1] if len(y_lag) > 1 else y_lag
        y = y_diff

        # Regression: Δy_t = α + β*y_{t-1} + ε_t
        try:
            # Simple regression for coefficient
            slope, intercept, r_value, p_value, std_err = linregress(X, y)

            # ADF statistic approximation
            adf_stat = slope / std_err if std_err > 0 else 0

            # Critical values (approximate)
            critical_values = {0.01: -3.43, 0.05: -2.86, 0.10: -2.57}

            # Interpretation
            if adf_stat < critical_values[0.05]:
                interpretation = "Time series is stationary (rejects unit root)"
            else:
                interpretation = "Time series may have unit root (non-stationary)"

        except:
            adf_stat = 0
            p_value = 1.0
            interpretation = "Test failed - insufficient data"

        return StatisticalResult(
            statistic=adf_stat,
            p_value=p_value,
            confidence_interval=(-np.inf, np.inf),
            method="Augmented Dickey-Fuller Test",
            interpretation=interpretation
        )

    def spectral_analysis(self, data: np.ndarray, sampling_rate: float = 1.0) -> Dict[str, Any]:
        """
        Power spectral density analysis for frequency domain insights.

        Args:
            data: Time series data
            sampling_rate: Sampling rate in Hz

        Returns:
            Dictionary with spectral analysis results
        """
        # Welch's method for PSD estimation
        frequencies, psd = welch(data, fs=sampling_rate, nperseg=min(256, len(data)//4))

        # Peak frequency
        peak_idx = np.argmax(psd)
        peak_frequency = frequencies[peak_idx]

        # Spectral centroid (center of mass)
        spectral_centroid = np.sum(frequencies * psd) / np.sum(psd)

        # Spectral spread (second moment)
        spectral_spread = np.sqrt(np.sum(((frequencies - spectral_centroid)**2) * psd) / np.sum(psd))

        # Shannon entropy of power spectrum
        psd_normalized = psd / np.sum(psd)
        psd_normalized = psd_normalized[psd_normalized > 0]  # Remove zeros
        spectral_entropy = -np.sum(psd_normalized * np.log2(psd_normalized))

        return {
            'frequencies': frequencies,
            'power_spectral_density': psd,
            'peak_frequency': peak_frequency,
            'spectral_centroid': spectral_centroid,
            'spectral_spread': spectral_spread,
            'spectral_entropy': spectral_entropy,
            'total_power': np.sum(psd)
        }

    def information_theoretic_measures(self, data1: np.ndarray, data2: np.ndarray = None) -> Dict[str, float]:
        """
        Information-theoretic measures for data analysis.

        Args:
            data1: First dataset
            data2: Second dataset (optional, for mutual information)

        Returns:
            Dictionary with information measures
        """
        # Discretize data for entropy calculation
        bins = min(50, int(np.sqrt(len(data1))))
        hist1, _ = np.histogram(data1, bins=bins, density=True)
        hist1 = hist1[hist1 > 0]

        # Shannon entropy
        shannon_entropy = entropy(hist1, base=2)

        # Differential entropy approximation
        differential_entropy = shannon_entropy + np.log2(np.std(data1) * np.sqrt(2 * np.pi * np.e))

        measures = {
            'shannon_entropy': shannon_entropy,
            'differential_entropy': differential_entropy
        }

        if data2 is not None:
            # Joint and mutual information
            hist2, _ = np.histogram(data2, bins=bins, density=True)
            hist2 = hist2[hist2 > 0]

            # Joint histogram
            joint_hist, _, _ = np.histogram2d(data1, data2, bins=bins, density=True)
            joint_hist = joint_hist[joint_hist > 0]

            # Mutual information
            joint_entropy = entropy(joint_hist.flatten(), base=2)
            entropy2 = entropy(hist2, base=2)
            mutual_information = shannon_entropy + entropy2 - joint_entropy

            measures.update({
                'entropy_data2': entropy2,
                'joint_entropy': joint_entropy,
                'mutual_information': mutual_information,
                'normalized_mutual_information': 2 * mutual_information / (shannon_entropy + entropy2)
            })

        return measures
