"""
Sparse Synaptic Mesh

Bio-mimetic neural network with sparse connectivity:
    - Fixed metabolic budget (connection limit)
    - Synaptic pruning (remove weak connections)
    - Neurogenesis (grow connections where needed)
    - Salience-guided topology evolution

Mimics brain efficiency: massive capacity, sparse activation.
"""

import numpy as np
from typing import Optional, Tuple, List, Dict, Any
from dataclasses import dataclass


@dataclass  
class SparseMeshConfig:
    """Configuration for sparse mesh network."""
    input_dim: int = 2
    hidden_dim: int = 1000  # Large capacity
    output_dim: int = 1
    density: float = 0.05   # Only 5% connections active
    pruning_rate: float = 0.2
    regrowth_rate: float = 0.1


class SparseSynapticMesh:
    """
    Bio-mimetic sparse neural network.
    
    Key principles:
        - Metabolic budget: Fixed number of active connections
        - Synaptic plasticity: Connections strengthen/weaken based on use
        - Structural plasticity: Topology changes via pruning/growth
        - Homeostasis: System self-regulates to maintain stability
    """
    
    def __init__(self, config: Optional[SparseMeshConfig] = None):
        self.config = config or SparseMeshConfig()
        
        self.input_dim = self.config.input_dim
        self.hidden_dim = self.config.hidden_dim
        self.output_dim = self.config.output_dim
        self.density = self.config.density
        
        # Initialize dense weights (will be masked)
        self.W1 = np.random.randn(self.input_dim, self.hidden_dim) * np.sqrt(2/self.input_dim)
        self.W2 = np.random.randn(self.hidden_dim, self.output_dim) * np.sqrt(2/self.hidden_dim)
        
        # Connectivity masks (the connectome)
        self.mask1 = np.random.rand(self.input_dim, self.hidden_dim) < self.density
        self.mask2 = np.random.rand(self.hidden_dim, self.output_dim) < self.density
        
        # Neuron-level salience tracking
        self.neuron_salience = np.zeros(self.hidden_dim)
        self.connection_strength = np.zeros_like(self.W1)
        
        # Metabolic tracking
        self.active_connections: List[int] = []
        self.connection_history: List[int] = []
        self.sparsity_history: List[float] = []
        
    def forward(self, X: np.ndarray) -> np.ndarray:
        """Forward pass with sparse connectivity."""
        X = X.reshape(-1, self.input_dim) if X.ndim == 1 else X
        
        # Apply masks (enforce sparsity)
        W1_active = self.W1 * self.mask1
        W2_active = self.W2 * self.mask2
        
        # Forward propagation
        self.z1 = np.dot(X, W1_active)
        self.a1 = np.maximum(0, self.z1)  # ReLU
        self.z2 = np.dot(self.a1, W2_active)
        self.a2 = 1 / (1 + np.exp(-np.clip(self.z2, -500, 500)))  # Sigmoid
        
        return self.a2
    
    def train_step(self, X: np.ndarray, y: np.ndarray, lr: float = 0.01) -> float:
        """Train with gradient masking."""
        X = X.reshape(-1, self.input_dim) if X.ndim == 1 else X
        y = y.reshape(-1, self.output_dim) if y.ndim == 1 else y
        m = X.shape[0]
        
        # Forward
        pred = self.forward(X)
        
        # Loss
        eps = 1e-8
        loss = -np.mean(y * np.log(pred + eps) + (1-y) * np.log(1 - pred + eps))
        
        # Backward
        W1_active = self.W1 * self.mask1
        W2_active = self.W2 * self.mask2
        
        d_z2 = pred - y
        d_W2 = np.dot(self.a1.T, d_z2) / m
        
        d_a1 = np.dot(d_z2, W2_active.T)
        d_z1 = d_a1 * (self.z1 > 0)  # ReLU gradient
        d_W1 = np.dot(X.T, d_z1) / m
        
        # Update only active connections
        self.W1 -= lr * (d_W1 * self.mask1)
        self.W2 -= lr * (d_W2 * self.mask2)
        
        # Track salience
        batch_salience = np.mean(np.abs(self.a1 * d_a1), axis=0)
        self.neuron_salience = 0.9 * self.neuron_salience + 0.1 * batch_salience
        
        # Track connection strength (Hebbian-like)
        self.connection_strength = 0.95 * self.connection_strength + 0.05 * np.abs(d_W1)
        
        return float(loss)
    
    def sleep_and_rewire(self, pruning_rate: Optional[float] = None) -> Dict[str, Any]:
        """
        Metabolic cycle: prune weak connections, grow new ones.
        
        Mimics synaptic homeostasis during sleep:
            1. Prune: Remove connections below threshold
            2. Regrow: Add connections where salience is high
            3. Balance: Maintain fixed metabolic budget
        """
        if pruning_rate is None:
            pruning_rate = self.config.pruning_rate
            
        results = {'pruned': 0, 'regrown': 0, 'active': 0}
        
        # === PRUNING ===
        # Find weak connections
        active_weights = np.abs(self.W1[self.mask1])
        if len(active_weights) > 0:
            threshold = np.percentile(active_weights, pruning_rate * 100)
            weak_connections = np.abs(self.W1) < threshold
            
            # Prune (remove from mask)
            pruned_mask = self.mask1 & weak_connections
            results['pruned'] = np.sum(pruned_mask)
            self.mask1[pruned_mask] = False
        
        # === REGROWTH ===
        # Calculate growth potential per neuron
        current_connectivity = np.sum(self.mask1, axis=0)  # Connections per neuron
        
        # Growth potential = salience / connectivity
        # High salience + low connectivity = needs more connections
        growth_potential = self.neuron_salience / (current_connectivity + 1e-5)
        
        # Target connection count (metabolic budget)
        target_connections = int(self.input_dim * self.hidden_dim * self.density)
        current_connections = np.sum(self.mask1)
        vacancies = target_connections - current_connections
        
        if vacancies > 0:
            # Probabilistic regrowth based on growth potential
            probs = growth_potential / (np.sum(growth_potential) + 1e-8)
            
            # Find empty connection slots
            empty_slots = np.where(~self.mask1)
            if len(empty_slots[0]) > 0:
                # Sample slots to fill
                n_regrow = min(vacancies, len(empty_slots[0]))
                
                # Weight by neuron salience
                slot_weights = probs[empty_slots[1]]
                slot_weights = slot_weights / (np.sum(slot_weights) + 1e-8)
                
                indices = np.random.choice(
                    len(empty_slots[0]), 
                    size=n_regrow, 
                    replace=False,
                    p=slot_weights if np.sum(slot_weights) > 0 else None
                )
                
                for idx in indices:
                    row, col = empty_slots[0][idx], empty_slots[1][idx]
                    self.mask1[row, col] = True
                    self.W1[row, col] = np.random.randn() * 0.1  # Small initial weight
                
                results['regrown'] = n_regrow
        
        results['active'] = np.sum(self.mask1)
        self.connection_history.append(results['active'])
        self.sparsity_history.append(1 - results['active'] / (self.input_dim * self.hidden_dim))
        
        return results
    
    def get_active_neurons(self) -> int:
        """Count neurons with at least one active connection."""
        return np.sum(np.sum(self.mask1, axis=0) > 0)
    
    def get_sparsity(self) -> float:
        """Get current sparsity level."""
        total_possible = self.input_dim * self.hidden_dim
        active = np.sum(self.mask1)
        return 1 - active / total_possible
    
    def get_parameters(self) -> np.ndarray:
        """Get only active parameters."""
        active_w1 = self.W1[self.mask1]
        active_w2 = self.W2[self.mask2]
        return np.concatenate([active_w1, active_w2])
    
    def get_topology_stats(self) -> Dict[str, Any]:
        """Get network topology statistics."""
        connectivity = np.sum(self.mask1, axis=0)
        
        return {
            'total_neurons': self.hidden_dim,
            'active_neurons': self.get_active_neurons(),
            'total_connections': np.sum(self.mask1) + np.sum(self.mask2),
            'sparsity': self.get_sparsity(),
            'avg_connectivity': np.mean(connectivity),
            'max_connectivity': np.max(connectivity),
            'min_connectivity': np.min(connectivity[connectivity > 0]) if np.any(connectivity > 0) else 0
        }


class HomeostaticSparseMesh(SparseSynapticMesh):
    """
    Sparse mesh with homeostatic regulation.
    
    Implements biological homeostasis:
        - Activity normalization
        - Synaptic scaling
        - Intrinsic plasticity
    """
    
    def __init__(self, config: Optional[SparseMeshConfig] = None):
        super().__init__(config)
        
        self.target_activity = 0.1  # Target firing rate
        self.activity_history = []
        self.scaling_factor = np.ones(self.hidden_dim)
        
    def forward(self, X: np.ndarray) -> np.ndarray:
        """Forward with homeostatic scaling."""
        X = X.reshape(-1, self.input_dim) if X.ndim == 1 else X
        
        W1_active = self.W1 * self.mask1
        W2_active = self.W2 * self.mask2
        
        self.z1 = np.dot(X, W1_active) * self.scaling_factor
        self.a1 = np.maximum(0, self.z1)
        
        # Track activity
        activity = np.mean(self.a1 > 0, axis=0)
        self.activity_history.append(np.mean(activity))
        
        self.z2 = np.dot(self.a1, W2_active)
        self.a2 = 1 / (1 + np.exp(-np.clip(self.z2, -500, 500)))
        
        return self.a2
    
    def homeostatic_update(self):
        """
        Update scaling factors to maintain target activity.
        
        If neuron is too active -> scale down
        If neuron is too quiet -> scale up
        """
        if not self.activity_history:
            return
            
        recent_activity = np.mean(self.activity_history[-10:]) if len(self.activity_history) >= 10 else self.activity_history[-1]
        
        # Compute error from target
        error = self.target_activity - recent_activity
        
        # Update scaling (slow dynamics)
        self.scaling_factor *= (1 + 0.01 * error)
        self.scaling_factor = np.clip(self.scaling_factor, 0.1, 10.0)
