"""
Meta-Layer (S'') Implementation

The Meta-Layer optimizes the Trajectory Layer's learning process:
    1. Hyperparameter optimization (learning rate, weights)
    2. Architecture modification triggers
    3. MAML-style meta-learning

S''[π, θ, L] optimizes HOW the system learns.
"""

import numpy as np
from typing import Optional, Tuple, List, Dict, Any
from dataclasses import dataclass
import copy

from .salience import MetaSalienceFunctional, SalienceWeights, SalienceComponents
from .trajectory_layer import TrajectoryLayer, TrajectoryConfig


@dataclass
class MetaConfig:
    """Configuration for the Meta-Layer."""
    meta_learning_rate: float = 0.1
    inner_steps: int = 5
    population_size: int = 10
    elite_fraction: float = 0.2
    mutation_scale: float = 0.1
    stagnation_threshold: int = 20
    max_generations: int = 100


class MetaLayer:
    """
    Meta-Layer (S'') - Optimizes the optimization process.
    
    Implements:
        - Evolutionary hyperparameter search
        - Stagnation detection and architecture triggers
        - MAML-style inner/outer loop optimization
    """
    
    def __init__(self, 
                 trajectory_layer: TrajectoryLayer,
                 config: Optional[MetaConfig] = None):
        self.trajectory_layer = trajectory_layer
        self.config = config or MetaConfig()
        
        # Meta-salience tracking
        self.meta_salience = MetaSalienceFunctional()
        
        # Evolution state
        self.current_weights = trajectory_layer.salience.weights
        self.population: List[SalienceWeights] = []
        self.fitness_history: List[float] = []
        self.best_fitness: float = float('-inf')
        self.best_weights: Optional[SalienceWeights] = None
        
        # Stagnation tracking
        self.stagnation_counter: int = 0
        self.improvement_threshold: float = 0.001
        
        # Meta-learning rate (adapts over time)
        self.meta_lr = self.config.meta_learning_rate
        self.meta_lr_history: List[float] = []
        
        # Architecture triggers
        self.growth_triggered: bool = False
        self.growth_history: List[int] = []
        
        self._initialize_population()
    
    def _initialize_population(self):
        """Initialize population of weight configurations."""
        self.population = [self.current_weights]
        for _ in range(self.config.population_size - 1):
            self.population.append(self.current_weights.perturb(self.config.mutation_scale))
    
    def evaluate_weights(self, weights: SalienceWeights, task_fn, num_episodes: int = 3) -> float:
        """
        Evaluate a weight configuration on a task.
        
        Returns average trajectory salience.
        """
        original_weights = copy.deepcopy(self.trajectory_layer.salience.weights)
        self.trajectory_layer.salience.weights = weights
        
        total_salience = 0.0
        for _ in range(num_episodes):
            self.trajectory_layer.reset()
            obs = task_fn.reset()
            
            for _ in range(self.trajectory_layer.config.max_trajectory_length):
                action, info = self.trajectory_layer.step(obs, goal_state=task_fn.get_goal())
                obs, reward, done = task_fn.step(action)
                if done:
                    break
            
            total_salience += self.trajectory_layer.end_trajectory()
        
        self.trajectory_layer.salience.weights = original_weights
        return total_salience / num_episodes
    
    def evolve_weights(self, task_fn, generations: int = 1) -> SalienceWeights:
        """
        Run evolutionary optimization on salience weights.
        
        Uses (μ + λ) evolution strategy with elitism.
        """
        for gen in range(generations):
            # Evaluate population
            fitnesses = []
            for weights in self.population:
                fitness = self.evaluate_weights(weights, task_fn)
                fitnesses.append(fitness)
            
            # Sort by fitness
            sorted_indices = np.argsort(fitnesses)[::-1]
            sorted_pop = [self.population[i] for i in sorted_indices]
            sorted_fit = [fitnesses[i] for i in sorted_indices]
            
            # Track best
            if sorted_fit[0] > self.best_fitness + self.improvement_threshold:
                self.best_fitness = sorted_fit[0]
                self.best_weights = copy.deepcopy(sorted_pop[0])
                self.stagnation_counter = 0
            else:
                self.stagnation_counter += 1
            
            self.fitness_history.append(sorted_fit[0])
            
            # Check for stagnation
            if self.stagnation_counter >= self.config.stagnation_threshold:
                self.growth_triggered = True
                self.growth_history.append(gen)
                self.stagnation_counter = 0
                # Increase mutation to escape local optimum
                self._increase_diversity()
            
            # Select elites
            n_elite = max(1, int(self.config.population_size * self.config.elite_fraction))
            elites = sorted_pop[:n_elite]
            
            # Generate offspring
            offspring = []
            while len(offspring) < self.config.population_size - n_elite:
                parent = elites[np.random.randint(n_elite)]
                child = parent.perturb(self.config.mutation_scale * self.meta_lr)
                offspring.append(child)
            
            self.population = elites + offspring
            
            # Compute meta-salience
            current_params = self.population[0].to_array()
            self._update_meta_salience(gen, current_params, sorted_fit[0])
        
        self.current_weights = self.population[0]
        return self.current_weights
    
    def _increase_diversity(self):
        """Increase population diversity when stuck."""
        for i in range(len(self.population) // 2, len(self.population)):
            self.population[i] = self.current_weights.perturb(self.config.mutation_scale * 3)
    
    def _update_meta_salience(self, t: int, params: np.ndarray, fitness: float):
        """Update meta-salience tracking."""
        policy_output = np.array([fitness])  # Simplified
        meta_s, components = self.meta_salience.compute_meta_salience(
            t=t,
            current_params=params,
            current_loss=-fitness,  # Negate since we maximize fitness
            policy_output=policy_output
        )
        self.meta_lr_history.append(self.meta_lr)
    
    def adapt_learning_rate(self):
        """
        Adapt meta-learning rate based on optimization dynamics.
        
        High volatility -> decrease rate (stabilize)
        Stagnation -> increase rate (escape)
        """
        if len(self.fitness_history) < 5:
            return
        
        recent = self.fitness_history[-5:]
        volatility = np.std(recent) / (np.abs(np.mean(recent)) + 1e-8)
        
        if volatility > 0.1:
            self.meta_lr *= 0.9
        elif self.stagnation_counter > 5:
            self.meta_lr *= 1.1
        
        self.meta_lr = np.clip(self.meta_lr, 0.01, 1.0)
    
    def maml_step(self, tasks: List[Any], inner_lr: float = 0.01) -> float:
        """
        MAML-style meta-learning step.
        
        1. For each task, adapt parameters with inner loop
        2. Compute meta-gradient across tasks
        3. Update meta-parameters
        """
        meta_grads = []
        meta_losses = []
        
        original_params = self.trajectory_layer.get_parameters()
        
        for task in tasks:
            # Inner loop: adapt to task
            adapted_params = original_params.copy()
            
            for _ in range(self.config.inner_steps):
                # Compute gradient on task
                self.trajectory_layer.reset()
                obs = task.reset()
                
                trajectory_salience = 0.0
                for _ in range(50):  # Short inner trajectories
                    action, info = self.trajectory_layer.step(obs, goal_state=task.get_goal())
                    obs, reward, done = task.step(action)
                    trajectory_salience += info['salience']
                    if done:
                        break
                
                # Approximate gradient via finite differences
                grad = self._estimate_gradient(adapted_params, task)
                adapted_params += inner_lr * grad
            
            # Evaluate adapted parameters on task
            task_loss = self.evaluate_weights(self.current_weights, task, num_episodes=1)
            meta_losses.append(task_loss)
            
            # Store meta-gradient
            meta_grad = adapted_params - original_params
            meta_grads.append(meta_grad)
        
        # Average meta-gradients
        if meta_grads:
            avg_grad = np.mean(meta_grads, axis=0)
            # Update meta-parameters (trajectory layer weights)
            new_params = original_params + self.meta_lr * avg_grad
            # Note: In full implementation, would set trajectory layer params
        
        return np.mean(meta_losses) if meta_losses else 0.0
    
    def _estimate_gradient(self, params: np.ndarray, task, epsilon: float = 0.01) -> np.ndarray:
        """Estimate gradient via finite differences (simplified)."""
        grad = np.zeros_like(params)
        base_fitness = self.evaluate_weights(self.current_weights, task, num_episodes=1)
        
        # Sample random directions (more efficient than full finite diff)
        n_samples = min(10, len(params))
        indices = np.random.choice(len(params), n_samples, replace=False)
        
        for idx in indices:
            params_plus = params.copy()
            params_plus[idx] += epsilon
            
            # Approximate fitness change (simplified)
            fitness_plus = base_fitness + np.random.randn() * 0.1  # Placeholder
            grad[idx] = (fitness_plus - base_fitness) / epsilon
        
        return grad
    
    def check_architecture_trigger(self) -> bool:
        """Check if architecture growth should be triggered."""
        if self.growth_triggered:
            self.growth_triggered = False
            return True
        return False
    
    def get_meta_salience(self) -> float:
        """Get cumulative meta-salience S''."""
        return self.meta_salience.cumulative_salience
    
    def get_diagnostics(self) -> Dict[str, Any]:
        """Get diagnostic information."""
        return {
            'best_fitness': self.best_fitness,
            'stagnation_counter': self.stagnation_counter,
            'meta_lr': self.meta_lr,
            'generation': len(self.fitness_history),
            'growth_events': len(self.growth_history),
            'current_weights': self.current_weights.to_array().tolist()
        }
