#!/usr/bin/env python3
"""
TiltDetector.py
Detects tilted or fragile market conditions where state estimation is unreliable.

Part of the Lacefield Portable Laboratory v1.0
Location: new/detectors/TiltDetector.py

This module identifies conditions where the kernel's diagnosis should be treated
with lower confidence (high tilt = fragile / unreliable diagnosis).
"""

import numpy as np
from dataclasses import dataclass
from typing import Dict, Any, List, Tuple, Optional


@dataclass
class TiltSignal:
    is_tilted: bool
    tilt_severity: float          # 0.0 - 1.0
    primary_cause: str
    secondary_causes: List[str]
    fragility_score: float        # Likelihood of imminent regime flip
    recommendations: List[str]
    metadata: Dict[str, Any]


class TiltDetector:
    """
    Identifies conditions where state estimation is unreliable.
    
    High tilt = low confidence in current diagnosis
    Low tilt = reliable diagnosis
    """

    def __init__(self, verbosity: int = 0):
        self.verbosity = verbosity

    def detect(self, channels: List[Any]) -> TiltSignal:
        """
        Analyze channels for tilted conditions.
        channels: List of objects with .coherence, .aggregate_value, .vectors
        """
        if not channels:
            return TiltSignal(
                is_tilted=True,
                tilt_severity=1.0,
                primary_cause="No channels provided",
                secondary_causes=[],
                fragility_score=1.0,
                recommendations=["Provide valid channels"],
                metadata={}
            )

        signals = {
            "coherence_breaks": self._check_coherence_breaks(channels),
            "extreme_divergence": self._check_extreme_divergence(channels),
            "single_factor_domination": self._check_single_factor_domination(channels),
            "correlation_shocks": self._check_correlation_shocks(channels),
        }

        active_signals = [(k, v) for k, v in signals.items() if v.get("active", False)]

        if not active_signals:
            return TiltSignal(
                is_tilted=False,
                tilt_severity=0.0,
                primary_cause="No tilt detected",
                secondary_causes=[],
                fragility_score=0.0,
                recommendations=[],
                metadata=signals
            )

        active_signals.sort(key=lambda x: x[1].get("severity", 0), reverse=True)
        primary = active_signals[0][0]
        secondary = [s[0] for s in active_signals[1:]]

        tilt_severity = self._aggregate_severity([s for _, s in active_signals])
        fragility = self._estimate_fragility(channels, signals, tilt_severity)
        recommendations = self._generate_recommendations(primary, secondary, channels)

        return TiltSignal(
            is_tilted=tilt_severity > 0.3,
            tilt_severity=tilt_severity,
            primary_cause=primary,
            secondary_causes=secondary,
            fragility_score=fragility,
            recommendations=recommendations,
            metadata=signals
        )

    # --- Internal Checks ---

    def _check_coherence_breaks(self, channels: List[Any]) -> Dict[str, Any]:
        coherence_values = [getattr(ch, 'coherence', 1.0) for ch in channels]
        mean_coherence = np.mean(coherence_values)
        min_coherence = np.min(coherence_values)

        severe = any(c < 0.3 for c in coherence_values)
        moderate = mean_coherence < 0.5

        severity = 0.8 if severe else (0.4 if moderate else 0.0)

        return {
            "active": severity > 0.0,
            "severity": severity,
            "mean_coherence": float(mean_coherence),
            "min_coherence": float(min_coherence),
        }

    def _check_extreme_divergence(self, channels: List[Any]) -> Dict[str, Any]:
        values = [getattr(ch, 'aggregate_value', 0.0) for ch in channels]
        std_val = np.std(values) if len(values) > 1 else 0.0

        severe = std_val > 0.6
        moderate = std_val > 0.4

        severity = 0.7 if severe else (0.35 if moderate else 0.0)

        return {
            "active": severity > 0.0,
            "severity": severity,
            "std_value": float(std_val),
        }

    def _check_single_factor_domination(self, channels: List[Any]) -> Dict[str, Any]:
        all_vectors = []
        for ch in channels:
            for vec in getattr(ch, 'vectors', []):
                val = getattr(vec, 'value', 0.0)
                all_vectors.append(abs(val))

        if not all_vectors:
            return {"active": False, "severity": 0.0}

        sorted_vals = sorted(all_vectors, reverse=True)
        if len(sorted_vals) > 1 and sorted_vals[1] > 0:
            ratio = sorted_vals[0] / sorted_vals[1]
            severe = ratio > 3.0
            moderate = ratio > 2.0
        else:
            severe = True
            moderate = False

        severity = 0.65 if severe else (0.3 if moderate else 0.0)

        return {
            "active": severity > 0.0,
            "severity": severity,
            "domination_ratio": float(ratio) if 'ratio' in locals() else None,
        }

    def _check_correlation_shocks(self, channels: List[Any]) -> Dict[str, Any]:
        correlations = []
        for ch in channels:
            if hasattr(ch, 'vector_corr_summary') and ch.vector_corr_summary:
                correlations.extend(ch.vector_corr_summary.values())

        if not correlations:
            return {"active": False, "severity": 0.0}

        std_corr = np.std(correlations)
        severe = std_corr > 0.7
        moderate = std_corr > 0.5

        severity = 0.6 if severe else (0.3 if moderate else 0.0)

        return {
            "active": severity > 0.0,
            "severity": severity,
            "std_correlation": float(std_corr),
        }

    def _aggregate_severity(self, signals: List[Dict[str, Any]]) -> float:
        if not signals:
            return 0.0
        severities = [s.get("severity", 0) for s in signals]
        return float(np.mean(severities))

    def _estimate_fragility(self, channels, signals, tilt_severity):
        coherence_values = [getattr(ch, 'coherence', 1.0) for ch in channels]
        mean_coherence = np.mean(coherence_values)
        fragility = tilt_severity
        if mean_coherence < 0.4:
            fragility += 0.3
        if signals.get("correlation_shocks", {}).get("active", False):
            fragility += 0.2
        return min(fragility, 1.0)

    def _generate_recommendations(self, primary, secondary, channels):
        recommendations = []
        if primary == "coherence_breaks":
            recommendations.append("Check vector data quality and noise levels")
            recommendations.append("Consider regime has changed; update baseline expectations")
        if primary == "extreme_divergence":
            recommendations.append("Channels pointing in opposite directions — conflicting evidence")
            recommendations.append("Require additional observation before committing to action")
        if primary == "single_factor_domination":
            recommendations.append("Diagnosis depends too heavily on one factor")
            recommendations.append("Seek independent confirmation from other channels")
        if primary == "correlation_shocks":
            recommendations.append("Inter-vector correlations are breaking down")
            recommendations.append("Regime transition likely — increase sampling frequency")
        return recommendations
