"""
Self-Layer (S''') Implementation

The Self-Layer optimizes the self-improvement process itself:
    1. Monitors convergence of S''
    2. Prevents divergence via complexity penalties
    3. Aligns optimization with goals
    4. Triggers structural modifications

S'''[S'', S'] is the outermost recursive loop.
"""

import numpy as np
from typing import Optional, Tuple, List, Dict, Any, Callable
from dataclasses import dataclass

from .salience import SelfSalienceFunctional, SalienceWeights
from .meta_layer import MetaLayer


@dataclass
class SelfConfig:
    """Configuration for the Self-Layer."""
    stability_window: int = 10
    divergence_threshold: float = 10.0
    complexity_budget: float = 1000.0
    alignment_weight: float = 0.5
    adaptation_rate: float = 0.05


class SelfLayer:
    """
    Self-Layer (S''') - Optimizes the self-improvement process.
    
    Monitors:
        - Meta-optimization convergence
        - System complexity (Kolmogorov-like proxy)
        - Goal alignment
    
    Actions:
        - Adjust meta-learning rate
        - Trigger architecture changes
        - Reset unstable components
    """
    
    def __init__(self,
                 meta_layer: MetaLayer,
                 config: Optional[SelfConfig] = None):
        self.meta_layer = meta_layer
        self.config = config or SelfConfig()
        
        # Self-salience tracking
        self.self_salience = SelfSalienceFunctional()
        
        # Monitoring history
        self.s_pp_history: List[float] = []  # S'' values
        self.complexity_history: List[float] = []
        self.stability_history: List[float] = []
        self.alignment_history: List[float] = []
        
        # Control signals
        self.meta_lr_multiplier: float = 1.0
        self.architecture_growth_enabled: bool = True
        self.emergency_reset_triggered: bool = False
        
        # Goal tracking
        self.goal_function: Optional[Callable] = None
        self.goal_progress: List[float] = []
    
    def set_goal_function(self, goal_fn: Callable[[Any], float]):
        """Set the goal function for alignment tracking."""
        self.goal_function = goal_fn
    
    def step(self, t: int) -> Dict[str, Any]:
        """
        Execute one step of self-optimization.
        
        1. Observe meta-layer state
        2. Compute self-salience
        3. Apply control actions
        """
        # Get current meta-state
        s_pp = self.meta_layer.get_meta_salience()
        meta_diagnostics = self.meta_layer.get_diagnostics()
        meta_params = self.meta_layer.current_weights.to_array()
        
        # Compute self-salience
        self_s, components = self.self_salience.compute_self_salience(
            t=t,
            current_s_pp=s_pp,
            meta_params=meta_params
        )
        
        self.s_pp_history.append(s_pp)
        
        # Compute system complexity
        complexity = self._compute_complexity(meta_params)
        self.complexity_history.append(complexity)
        
        # Compute stability
        stability = self._compute_stability()
        self.stability_history.append(stability)
        
        # Compute alignment (if goal function set)
        alignment = self._compute_alignment()
        self.alignment_history.append(alignment)
        
        # Apply control actions
        actions_taken = self._apply_control(stability, complexity, alignment)
        
        return {
            'self_salience': self_s,
            'components': components.to_dict(),
            'complexity': complexity,
            'stability': stability,
            'alignment': alignment,
            'actions': actions_taken,
            'meta_lr_multiplier': self.meta_lr_multiplier
        }
    
    def _compute_complexity(self, params: np.ndarray) -> float:
        """
        Approximate system complexity.
        
        Uses parameter entropy and effective dimensionality.
        """
        # Normalize parameters
        params_abs = np.abs(params) + 1e-10
        params_norm = params_abs / np.sum(params_abs)
        
        # Entropy-based complexity
        entropy = -np.sum(params_norm * np.log(params_norm))
        
        # Effective dimensionality (number of significant parameters)
        threshold = 0.01 * np.max(params_abs)
        effective_dim = np.sum(params_abs > threshold)
        
        complexity = entropy * np.log(effective_dim + 1)
        return float(complexity)
    
    def _compute_stability(self) -> float:
        """
        Compute optimization stability.
        
        Low variance in S'' -> high stability.
        """
        if len(self.s_pp_history) < self.config.stability_window:
            return 1.0
        
        recent = self.s_pp_history[-self.config.stability_window:]
        variance = np.var(recent)
        mean_abs = np.abs(np.mean(recent)) + 1e-8
        
        cv = np.sqrt(variance) / mean_abs  # Coefficient of variation
        stability = np.exp(-cv)
        
        return float(stability)
    
    def _compute_alignment(self) -> float:
        """Compute alignment with goal function."""
        if self.goal_function is None:
            return 0.5
        
        try:
            alignment = self.goal_function(self.meta_layer)
            alignment = np.clip(alignment, 0, 1)
            self.goal_progress.append(alignment)
            return float(alignment)
        except:
            return 0.5
    
    def _apply_control(self, stability: float, complexity: float, alignment: float) -> List[str]:
        """
        Apply control actions based on system state.
        """
        actions = []
        
        # Stability control
        if stability < 0.3:
            self.meta_lr_multiplier *= 0.8
            actions.append('decreased_meta_lr')
            
            if stability < 0.1:
                self.emergency_reset_triggered = True
                actions.append('emergency_reset_triggered')
        elif stability > 0.9 and self.meta_layer.stagnation_counter > 5:
            self.meta_lr_multiplier *= 1.2
            actions.append('increased_meta_lr')
        
        # Complexity control
        if complexity > self.config.complexity_budget:
            self.architecture_growth_enabled = False
            actions.append('disabled_growth')
        else:
            self.architecture_growth_enabled = True
        
        # Alignment control
        if alignment < 0.3:
            # Increase meaning weight to prioritize goal alignment
            self.meta_layer.current_weights.w_M *= 1.1
            actions.append('boosted_meaning_weight')
        
        # Apply meta-LR multiplier
        self.meta_layer.meta_lr *= self.meta_lr_multiplier
        self.meta_layer.meta_lr = np.clip(self.meta_layer.meta_lr, 0.001, 1.0)
        self.meta_lr_multiplier = 1.0  # Reset multiplier
        
        return actions
    
    def should_grow_architecture(self) -> bool:
        """Check if architecture growth is allowed and needed."""
        if not self.architecture_growth_enabled:
            return False
        return self.meta_layer.check_architecture_trigger()
    
    def emergency_reset(self):
        """
        Emergency reset when system is diverging.
        
        Resets meta-layer to known good state.
        """
        if self.emergency_reset_triggered:
            # Reset to best known weights
            if self.meta_layer.best_weights is not None:
                self.meta_layer.current_weights = self.meta_layer.best_weights
            
            # Reset meta learning rate
            self.meta_layer.meta_lr = self.meta_layer.config.meta_learning_rate
            
            # Clear stagnation
            self.meta_layer.stagnation_counter = 0
            
            self.emergency_reset_triggered = False
            return True
        return False
    
    def get_self_salience(self) -> float:
        """Get cumulative self-salience S'''."""
        return self.self_salience.cumulative_salience
    
    def get_diagnostics(self) -> Dict[str, Any]:
        """Get diagnostic information."""
        return {
            'self_salience': self.get_self_salience(),
            'complexity': self.complexity_history[-1] if self.complexity_history else 0,
            'stability': self.stability_history[-1] if self.stability_history else 1,
            'alignment': self.alignment_history[-1] if self.alignment_history else 0.5,
            'growth_enabled': self.architecture_growth_enabled,
            'emergency_reset': self.emergency_reset_triggered
        }


class RecursiveAGICore:
    """
    Complete three-layer recursive AGI system.
    
    Integrates S', S'', S''' into a unified optimization loop.
    
    AGI = argmax_{π, θ, L} [ S'[ω] + S''[π, θ, L] + S'''[S'', S'] ]
    """
    
    def __init__(self,
                 trajectory_config: Optional[Any] = None,
                 meta_config: Optional[Any] = None,
                 self_config: Optional[SelfConfig] = None):
        from .trajectory_layer import TrajectoryLayer, TrajectoryConfig
        from .meta_layer import MetaConfig
        
        # Build layers
        self.trajectory_layer = TrajectoryLayer(trajectory_config or TrajectoryConfig())
        self.meta_layer = MetaLayer(self.trajectory_layer, meta_config or MetaConfig())
        self.self_layer = SelfLayer(self.meta_layer, self_config or SelfConfig())
        
        # Unified tracking
        self.total_steps: int = 0
        self.s_prime_history: List[float] = []
        self.s_double_prime_history: List[float] = []
        self.s_triple_prime_history: List[float] = []
        self.agi_score_history: List[float] = []
    
    def step(self, task) -> Dict[str, Any]:
        """Execute one full AGI step."""
        # Trajectory layer step
        obs = task.get_observation()
        action, traj_info = self.trajectory_layer.step(obs, goal_state=task.get_goal())
        obs_next, reward, done = task.step(action)
        
        s_prime = traj_info['salience']
        self.s_prime_history.append(s_prime)
        
        # Meta layer step (periodic)
        s_double_prime = self.meta_layer.get_meta_salience()
        self.s_double_prime_history.append(s_double_prime)
        
        # Self layer step (periodic)
        self_info = self.self_layer.step(self.total_steps)
        s_triple_prime = self_info['self_salience']
        self.s_triple_prime_history.append(s_triple_prime)
        
        # Compute AGI score
        agi_score = s_prime + s_double_prime + s_triple_prime
        self.agi_score_history.append(agi_score)
        
        # Check for architecture growth
        if self.self_layer.should_grow_architecture():
            self._trigger_growth()
        
        # Check for emergency reset
        self.self_layer.emergency_reset()
        
        self.total_steps += 1
        
        return {
            's_prime': s_prime,
            's_double_prime': s_double_prime,
            's_triple_prime': s_triple_prime,
            'agi_score': agi_score,
            'action': action,
            'done': done,
            'self_diagnostics': self_info
        }
    
    def _trigger_growth(self):
        """Trigger architecture growth event."""
        # This would trigger neural architecture growth
        # Implementation depends on the network type used
        pass
    
    def evolve(self, task, generations: int = 10) -> float:
        """Run evolutionary optimization."""
        self.meta_layer.evolve_weights(task, generations)
        return self.meta_layer.best_fitness
    
    def get_total_salience(self) -> Tuple[float, float, float]:
        """Get total salience from all layers."""
        return (
            self.trajectory_layer.get_cumulative_salience(),
            self.meta_layer.get_meta_salience(),
            self.self_layer.get_self_salience()
        )
