"""
Salience Functional Mathematics

Implements the core mathematical framework for the recursive self-optimizing AGI.

The salience functional S'[ω] is defined as:
    S'[ω] = ∫₀^∞ [L_sal(t)] dt

Where L_sal(t) = (w_A·ΔA_t + w_R·R_t + w_M·M_t) · C_t · e^(-λ_t·t) · e^(-k_φ·φ_t) - λ_jerk·||π̈_t||²

Components:
    ΔA_t : Novelty (information gain about world model)
    R_t  : Retention (predictive information about future)
    M_t  : Meaning (instrumental value toward goals)
    C_t  : Continuity (temporal coherence)
    φ_t  : Fatigue (diminishing returns signal)
"""

import numpy as np
from dataclasses import dataclass, field
from typing import Optional, Tuple, List, Dict, Any
from scipy.special import rel_entr
from scipy.stats import entropy


@dataclass
class SalienceComponents:
    """Container for salience component values at time t."""
    delta_A: float = 0.0  # Novelty
    R: float = 0.0        # Retention
    M: float = 0.0        # Meaning
    C: float = 1.0        # Continuity
    phi: float = 0.0      # Fatigue
    jerk: float = 0.0     # Policy acceleration penalty
    
    def to_dict(self) -> Dict[str, float]:
        return {
            'novelty': self.delta_A,
            'retention': self.R,
            'meaning': self.M,
            'continuity': self.C,
            'fatigue': self.phi,
            'jerk': self.jerk
        }


@dataclass 
class SalienceWeights:
    """Hyperparameters for the salience functional."""
    w_A: float = 1.0           # Novelty weight
    w_R: float = 1.0           # Retention weight  
    w_M: float = 1.0           # Meaning weight
    lambda_t: float = 0.01     # Temporal discount
    k_phi: float = 0.1         # Fatigue penalty coefficient
    lambda_jerk: float = 0.01  # Jerk penalty coefficient
    
    def to_array(self) -> np.ndarray:
        return np.array([self.w_A, self.w_R, self.w_M, 
                        self.lambda_t, self.k_phi, self.lambda_jerk])
    
    @classmethod
    def from_array(cls, arr: np.ndarray) -> 'SalienceWeights':
        return cls(
            w_A=float(arr[0]),
            w_R=float(arr[1]),
            w_M=float(arr[2]),
            lambda_t=float(arr[3]),
            k_phi=float(arr[4]),
            lambda_jerk=float(arr[5])
        )
    
    def perturb(self, scale: float = 0.1) -> 'SalienceWeights':
        """Create a perturbed copy for evolution."""
        arr = self.to_array()
        perturbed = arr * (1 + np.random.randn(len(arr)) * scale)
        # Ensure positivity
        perturbed = np.clip(perturbed, 1e-6, 10.0)
        return SalienceWeights.from_array(perturbed)


class SalienceFunctional:
    """
    Implements the Salience Functional S'[ω].
    
    This is the core optimization objective that drives learning.
    The functional rewards:
        - Novelty: New information about the world
        - Retention: Preserving useful predictive knowledge
        - Meaning: Progress toward goals
    
    While penalizing:
        - Fatigue: Diminishing returns from repeated patterns
        - Jerk: Abrupt policy changes
    """
    
    def __init__(self, weights: Optional[SalienceWeights] = None):
        self.weights = weights or SalienceWeights()
        self.history: List[SalienceComponents] = []
        self.cumulative_salience: float = 0.0
        
        # State tracking for temporal computations
        self._prev_belief: Optional[np.ndarray] = None
        self._prev_policy: Optional[np.ndarray] = None
        self._prev_prev_policy: Optional[np.ndarray] = None
        self._observation_buffer: List[np.ndarray] = []
        self._fatigue_accumulator: float = 0.0
        
    def reset(self):
        """Reset the functional for a new trajectory."""
        self.history = []
        self.cumulative_salience = 0.0
        self._prev_belief = None
        self._prev_policy = None
        self._prev_prev_policy = None
        self._observation_buffer = []
        self._fatigue_accumulator = 0.0
        
    def compute_novelty(self, 
                        current_belief: np.ndarray,
                        prev_belief: Optional[np.ndarray] = None) -> float:
        """
        Compute ΔA_t: Information gain about world model parameters.
        
        ΔA_t = D_KL(p(θ|h_t) || p(θ|h_{t-1}))
        
        Approximated as the KL divergence between belief distributions.
        Higher values indicate the agent learned something new.
        """
        if prev_belief is None:
            prev_belief = self._prev_belief
            
        if prev_belief is None:
            # First step: maximum novelty (everything is new)
            return 1.0
        
        # Ensure proper probability distributions
        current = self._to_probability(current_belief)
        previous = self._to_probability(prev_belief)
        
        # Compute KL divergence: D_KL(P || Q)
        # Using relative entropy which handles zeros properly
        kl_div = np.sum(rel_entr(current, previous))
        
        # Clip to prevent infinite values
        return float(np.clip(kl_div, 0, 10))
    
    def compute_retention(self,
                         observations: List[np.ndarray],
                         predictions: List[np.ndarray],
                         horizon: int = 5) -> float:
        """
        Compute R_t: Predictive information about future observations.
        
        R_t = I(o_t; o_{t+1:t+K} | h_{t-1})
        
        Approximated by measuring how well past states predict future states.
        This rewards maintaining useful predictive knowledge.
        """
        if len(observations) < 2:
            return 0.0
            
        # Use prediction accuracy as a proxy for mutual information
        # Higher accuracy = better retention of predictive structure
        
        recent_obs = observations[-min(horizon, len(observations)):]
        recent_pred = predictions[-min(horizon, len(predictions)):] if predictions else []
        
        if not recent_pred:
            return 0.5  # Neutral if no predictions yet
        
        # Compute prediction error (lower is better)
        errors = []
        for obs, pred in zip(recent_obs, recent_pred):
            if obs.shape == pred.shape:
                error = np.mean((obs - pred) ** 2)
                errors.append(error)
        
        if not errors:
            return 0.5
            
        avg_error = np.mean(errors)
        # Convert to retention score (higher is better)
        retention = np.exp(-avg_error)
        
        return float(retention)
    
    def compute_meaning(self,
                       observation: np.ndarray,
                       goal_state: Optional[np.ndarray] = None,
                       reward: float = 0.0) -> float:
        """
        Compute M_t: Instrumental value toward goals.
        
        M_t = I(o_t; R_{future} | h_{t-1})
        
        How much the current observation reduces uncertainty about future rewards.
        This can be approximated by goal proximity or direct reward signal.
        """
        if goal_state is not None:
            # Compute distance to goal (lower = higher meaning)
            distance = np.linalg.norm(observation.flatten() - goal_state.flatten())
            meaning = np.exp(-distance)
        else:
            # Fall back to reward signal
            meaning = np.clip(reward, -1, 1) * 0.5 + 0.5
            
        return float(meaning)
    
    def compute_continuity(self,
                          current_state: np.ndarray,
                          prev_state: Optional[np.ndarray] = None) -> float:
        """
        Compute C_t: Temporal coherence of experience.
        
        Measures how smoothly the trajectory evolves.
        Abrupt jumps in state space indicate discontinuity.
        """
        if prev_state is None:
            return 1.0
            
        # Compute state transition magnitude
        delta = np.linalg.norm(current_state.flatten() - prev_state.flatten())
        
        # Soft threshold: smooth transitions get C_t ≈ 1
        # Large jumps get C_t → 0
        continuity = np.exp(-delta)
        
        return float(continuity)
    
    def compute_fatigue(self,
                       current_belief: np.ndarray,
                       decay: float = 0.95) -> float:
        """
        Compute φ_t: Fatigue/diminishing returns signal.
        
        Accumulates when seeing similar patterns repeatedly.
        Acts as an anti-overfitting mechanism.
        """
        if self._prev_belief is None:
            self._fatigue_accumulator = 0.0
            return 0.0
        
        # Compute similarity to previous belief
        similarity = self._cosine_similarity(current_belief, self._prev_belief)
        
        # Accumulate fatigue (decays over time)
        self._fatigue_accumulator = (
            decay * self._fatigue_accumulator + 
            (1 - decay) * similarity
        )
        
        return float(self._fatigue_accumulator)
    
    def compute_jerk(self,
                    current_policy: np.ndarray,
                    prev_policy: Optional[np.ndarray] = None,
                    prev_prev_policy: Optional[np.ndarray] = None) -> float:
        """
        Compute ||π̈_t||²: Policy acceleration penalty.
        
        Jerk = second derivative of policy parameters.
        Penalizes abrupt changes in learning direction.
        """
        if prev_policy is None:
            prev_policy = self._prev_policy
        if prev_prev_policy is None:
            prev_prev_policy = self._prev_prev_policy
            
        if prev_policy is None or prev_prev_policy is None:
            return 0.0
        
        # First derivative (velocity)
        velocity_current = current_policy - prev_policy
        velocity_prev = prev_policy - prev_prev_policy
        
        # Second derivative (acceleration/jerk)
        acceleration = velocity_current - velocity_prev
        jerk = np.sum(acceleration ** 2)
        
        return float(jerk)
    
    def compute_salience(self,
                        t: int,
                        current_belief: np.ndarray,
                        current_policy: np.ndarray,
                        observation: np.ndarray,
                        predictions: List[np.ndarray],
                        goal_state: Optional[np.ndarray] = None,
                        reward: float = 0.0) -> Tuple[float, SalienceComponents]:
        """
        Compute the full salience value at time t.
        
        S_t = (w_A·ΔA_t + w_R·R_t + w_M·M_t) · C_t · e^(-λ_t·t) · e^(-k_φ·φ_t) - λ_jerk·||π̈_t||²
        """
        # Compute each component
        delta_A = self.compute_novelty(current_belief, self._prev_belief)
        
        self._observation_buffer.append(observation)
        R = self.compute_retention(self._observation_buffer, predictions)
        
        M = self.compute_meaning(observation, goal_state, reward)
        
        prev_obs = self._observation_buffer[-2] if len(self._observation_buffer) > 1 else None
        C = self.compute_continuity(observation, prev_obs)
        
        phi = self.compute_fatigue(current_belief)
        
        jerk = self.compute_jerk(current_policy, self._prev_policy, self._prev_prev_policy)
        
        # Store components
        components = SalienceComponents(
            delta_A=delta_A,
            R=R,
            M=M,
            C=C,
            phi=phi,
            jerk=jerk
        )
        
        # Compute weighted sum
        weighted_sum = (
            self.weights.w_A * delta_A +
            self.weights.w_R * R +
            self.weights.w_M * M
        )
        
        # Apply gating factors
        temporal_discount = np.exp(-self.weights.lambda_t * t)
        fatigue_penalty = np.exp(-self.weights.k_phi * phi)
        jerk_penalty = self.weights.lambda_jerk * jerk
        
        # Final salience
        salience = weighted_sum * C * temporal_discount * fatigue_penalty - jerk_penalty
        
        # Update state
        self._prev_prev_policy = self._prev_policy
        self._prev_policy = current_policy.copy()
        self._prev_belief = current_belief.copy()
        
        # Record history
        self.history.append(components)
        self.cumulative_salience += salience
        
        return float(salience), components
    
    def compute_trajectory_salience(self,
                                   trajectory: List[Dict[str, Any]]) -> float:
        """
        Compute S'[ω] for an entire trajectory.
        
        S'[ω] = ∫₀^T L_sal(t) dt ≈ Σ_t S_t
        """
        self.reset()
        total_salience = 0.0
        
        for t, step in enumerate(trajectory):
            salience, _ = self.compute_salience(
                t=t,
                current_belief=step['belief'],
                current_policy=step['policy'],
                observation=step['observation'],
                predictions=step.get('predictions', []),
                goal_state=step.get('goal'),
                reward=step.get('reward', 0.0)
            )
            total_salience += salience
            
        return total_salience
    
    # --- Utility Methods ---
    
    def _to_probability(self, arr: np.ndarray, eps: float = 1e-10) -> np.ndarray:
        """Convert array to valid probability distribution."""
        arr = arr.flatten()
        arr = np.abs(arr) + eps
        return arr / arr.sum()
    
    def _cosine_similarity(self, a: np.ndarray, b: np.ndarray) -> float:
        """Compute cosine similarity between two vectors."""
        a = a.flatten()
        b = b.flatten()
        
        norm_a = np.linalg.norm(a)
        norm_b = np.linalg.norm(b)
        
        if norm_a < 1e-10 or norm_b < 1e-10:
            return 0.0
            
        return float(np.dot(a, b) / (norm_a * norm_b))
    
    def get_component_history(self) -> Dict[str, List[float]]:
        """Get history of all components as separate lists."""
        return {
            'novelty': [c.delta_A for c in self.history],
            'retention': [c.R for c in self.history],
            'meaning': [c.M for c in self.history],
            'continuity': [c.C for c in self.history],
            'fatigue': [c.phi for c in self.history],
            'jerk': [c.jerk for c in self.history]
        }


class MetaSalienceFunctional(SalienceFunctional):
    """
    Meta-Salience Functional S''[π, θ, L]
    
    Applies the salience framework to the optimization process itself.
    
    S''[π, θ, L] = ∫₀^∞ [(w_A·ΔA_t^meta + w_R·R_t^meta + w_M·M_t^meta) · 
                          C_t^meta · e^(-λ_t·t) · e^(-k_φ·φ_t^meta) - 
                          λ_jerk·||π̈_t||²] dt
    
    Where:
        ΔA_t^meta : Novelty in optimization (how much did π improve?)
        R_t^meta  : Retention of progress (avoiding catastrophic forgetting)
        M_t^meta  : Alignment with objectives (does π align with goals?)
    """
    
    def __init__(self, weights: Optional[SalienceWeights] = None):
        super().__init__(weights)
        self._policy_history: List[np.ndarray] = []
        self._loss_history: List[float] = []
        self._param_history: List[np.ndarray] = []
        
    def compute_meta_novelty(self,
                            current_params: np.ndarray,
                            prev_params: Optional[np.ndarray] = None) -> float:
        """
        ΔA_t^meta = D_KL(p(θ_t) || p(θ_{t-1}))
        
        How much did the model parameters change?
        """
        if prev_params is None and len(self._param_history) > 0:
            prev_params = self._param_history[-1]
            
        if prev_params is None:
            return 1.0
            
        # Parameter divergence (normalized L2)
        delta = np.linalg.norm(current_params - prev_params)
        norm_factor = np.linalg.norm(prev_params) + 1e-8
        
        normalized_delta = delta / norm_factor
        
        # Map to [0, 1] range with soft saturation
        novelty = 1 - np.exp(-normalized_delta * 10)
        
        return float(novelty)
    
    def compute_meta_retention(self,
                              current_loss: float,
                              baseline_loss: Optional[float] = None) -> float:
        """
        R_t^meta = I(θ_t; θ_{t-T})
        
        Did we retain what we learned? Measured by loss stability.
        """
        if baseline_loss is None and len(self._loss_history) > 0:
            # Use loss from T steps ago
            T = min(10, len(self._loss_history))
            baseline_loss = self._loss_history[-T]
        
        if baseline_loss is None:
            return 0.5
            
        # If current loss is close to or better than baseline, good retention
        retention = np.exp(-(current_loss - baseline_loss))
        retention = np.clip(retention, 0, 1)
        
        return float(retention)
    
    def compute_meta_meaning(self,
                            policy_output: np.ndarray,
                            target_distribution: Optional[np.ndarray] = None,
                            alignment_score: float = 0.0) -> float:
        """
        M_t^meta = I(o_t; human_feedback | h_t)
        
        Does the policy align with intended objectives?
        """
        if target_distribution is not None:
            # Measure alignment with target
            current = self._to_probability(policy_output)
            target = self._to_probability(target_distribution)
            
            # Use negative KL divergence (higher = more aligned)
            kl = np.sum(rel_entr(target, current))
            meaning = np.exp(-kl)
        else:
            # Fall back to provided alignment score
            meaning = np.clip(alignment_score, 0, 1)
            
        return float(meaning)
    
    def compute_meta_salience(self,
                             t: int,
                             current_params: np.ndarray,
                             current_loss: float,
                             policy_output: np.ndarray,
                             alignment_score: float = 0.0) -> Tuple[float, SalienceComponents]:
        """
        Compute S'' for the optimization process.
        """
        prev_params = self._param_history[-1] if self._param_history else None
        
        # Meta-components
        delta_A = self.compute_meta_novelty(current_params, prev_params)
        R = self.compute_meta_retention(current_loss)
        M = self.compute_meta_meaning(policy_output, alignment_score=alignment_score)
        
        # Continuity of optimization
        if len(self._loss_history) > 1:
            loss_delta = abs(current_loss - self._loss_history[-1])
            C = np.exp(-loss_delta * 10)
        else:
            C = 1.0
        
        # Meta-fatigue (are we spinning wheels?)
        if len(self._loss_history) > 5:
            recent_variance = np.var(self._loss_history[-5:])
            phi = 1 - np.exp(-recent_variance * 100)
        else:
            phi = 0.0
        
        # Policy jerk
        current_policy = current_params  # Treat params as policy
        jerk = self.compute_jerk(current_policy, self._prev_policy, self._prev_prev_policy)
        
        components = SalienceComponents(
            delta_A=delta_A,
            R=R,
            M=M,
            C=C,
            phi=phi,
            jerk=jerk
        )
        
        # Compute weighted sum
        weighted_sum = (
            self.weights.w_A * delta_A +
            self.weights.w_R * R +
            self.weights.w_M * M
        )
        
        temporal_discount = np.exp(-self.weights.lambda_t * t)
        fatigue_penalty = np.exp(-self.weights.k_phi * phi)
        jerk_penalty = self.weights.lambda_jerk * jerk
        
        salience = weighted_sum * C * temporal_discount * fatigue_penalty - jerk_penalty
        
        # Update history
        self._param_history.append(current_params.copy())
        self._loss_history.append(current_loss)
        self._prev_prev_policy = self._prev_policy
        self._prev_policy = current_params.copy()
        self.history.append(components)
        self.cumulative_salience += salience
        
        return float(salience), components


class SelfSalienceFunctional(SalienceFunctional):
    """
    Self-Salience Functional S'''[S'', S']
    
    Optimizes the self-improvement process itself.
    
    Monitors:
        ΔA_t^self : Convergence of S''
        φ_t^self  : Complexity of S'' (Kolmogorov-like)
        M_t^self  : Alignment of S'' with human values
    """
    
    def __init__(self, weights: Optional[SalienceWeights] = None):
        super().__init__(weights)
        self._s_prime_prime_history: List[float] = []
        self._complexity_history: List[float] = []
        
    def compute_self_novelty(self,
                            current_s_pp: float,
                            prev_s_pp: Optional[float] = None) -> float:
        """
        ΔA_t^self = D_KL(p(S''_t) || p(S''_{t-1}))
        
        How much is the meta-optimization changing?
        """
        if prev_s_pp is None and len(self._s_prime_prime_history) > 0:
            prev_s_pp = self._s_prime_prime_history[-1]
            
        if prev_s_pp is None:
            return 1.0
            
        # Relative change in S''
        delta = abs(current_s_pp - prev_s_pp) / (abs(prev_s_pp) + 1e-8)
        novelty = 1 - np.exp(-delta)
        
        return float(novelty)
    
    def compute_self_complexity(self,
                               meta_params: np.ndarray) -> float:
        """
        φ_t^self = complexity(S''_t)
        
        Approximation of Kolmogorov complexity using parameter entropy.
        """
        # Use entropy of parameter distribution as complexity proxy
        params_normalized = self._to_probability(meta_params)
        complexity = entropy(params_normalized)
        
        return float(complexity)
    
    def compute_self_meaning(self,
                            current_s_pp: float,
                            prev_s_pp: Optional[float] = None,
                            target_improvement: float = 0.0) -> float:
        """
        M_t^self = I(S''_t; human_values)
        
        Is the self-improvement aligned with desired outcomes?
        """
        if prev_s_pp is None and len(self._s_prime_prime_history) > 0:
            prev_s_pp = self._s_prime_prime_history[-1]
            
        if prev_s_pp is None:
            return 0.5
            
        # Did S'' improve? (Higher S'' = better optimization)
        improvement = current_s_pp - prev_s_pp
        
        # Reward improvements, penalize regressions
        meaning = np.tanh(improvement)
        meaning = (meaning + 1) / 2  # Map to [0, 1]
        
        return float(meaning)
    
    def compute_self_salience(self,
                             t: int,
                             current_s_pp: float,
                             meta_params: np.ndarray) -> Tuple[float, SalienceComponents]:
        """
        Compute S''' for the self-improvement process.
        """
        delta_A = self.compute_self_novelty(current_s_pp)
        phi = self.compute_self_complexity(meta_params)
        M = self.compute_self_meaning(current_s_pp)
        
        # Self-continuity (is improvement stable?)
        if len(self._s_prime_prime_history) > 1:
            recent = self._s_prime_prime_history[-3:] if len(self._s_prime_prime_history) >= 3 else self._s_prime_prime_history
            C = np.exp(-np.var(recent))
        else:
            C = 1.0
        
        # Retention (are gains sustained?)
        if len(self._s_prime_prime_history) > 5:
            early = np.mean(self._s_prime_prime_history[:5])
            recent = np.mean(self._s_prime_prime_history[-5:])
            R = 1.0 if recent >= early else np.exp(-(early - recent))
        else:
            R = 0.5
        
        components = SalienceComponents(
            delta_A=delta_A,
            R=R,
            M=M,
            C=C,
            phi=phi,
            jerk=0.0  # Not computed at self level
        )
        
        # Weighted sum with complexity penalty
        weighted_sum = (
            self.weights.w_A * delta_A +
            self.weights.w_R * R +
            self.weights.w_M * M
        )
        
        temporal_discount = np.exp(-self.weights.lambda_t * t)
        complexity_penalty = np.exp(-self.weights.k_phi * phi)
        
        salience = weighted_sum * C * temporal_discount * complexity_penalty
        
        # Update history
        self._s_prime_prime_history.append(current_s_pp)
        self._complexity_history.append(phi)
        self.history.append(components)
        self.cumulative_salience += salience
        
        return float(salience), components
