"""
Recursive Self-Optimizing AGI Framework
Main Entry Point

Demonstrates the three-layer recursive optimization:
    S'   - Trajectory Layer (behavior optimization)
    S''  - Meta Layer (learning optimization)  
    S''' - Self Layer (self-improvement optimization)

AGI = argmax_{π, θ, L} [ S'[ω] + S''[π, θ, L] + S'''[S'', S'] ]
"""

import numpy as np
import matplotlib.pyplot as plt
from typing import Dict, List, Any

# Local imports
from networks.dynamic_net import DynamicNeuralNetwork, SalienceGuidedNetwork, DynamicNetConfig
from networks.sparse_mesh import SparseSynapticMesh, SparseMeshConfig
from environments.tasks import TwoMoonsTask, ConcentricCirclesTask, SpiralTask
from utils.visualization import AGIVisualizer
from utils.metrics import compute_metrics


def run_growing_brain_experiment(epochs: int = 1000, verbose: bool = True):
    """
    Experiment 1: Self-Growing Neural Core
    
    Network starts small and grows its architecture when optimization stagnates.
    Demonstrates S'' triggering structural changes.
    """
    print("\n" + "="*60)
    print("EXPERIMENT 1: SELF-GROWING NEURAL CORE")
    print("="*60)
    
    # Setup
    task = TwoMoonsTask()
    X, y = task.get_data()
    
    config = DynamicNetConfig(input_dim=2, initial_hidden=2, output_dim=1, max_hidden=50)
    brain = SalienceGuidedNetwork(config)
    
    # Tracking
    losses = []
    neuron_counts = []
    growth_events = []
    lr = 0.1
    lr_history = []
    
    print(f"Initial Brain Size: {brain.hidden_dim} neurons")
    print(f"Task: Two Moons Classification ({len(X)} samples)")
    print("-"*60)
    
    for epoch in range(epochs):
        # Self-Layer (S'''): Adapt learning rate based on volatility
        if len(losses) >= 10:
            recent_var = np.var(losses[-10:])
            if recent_var > 0.01:
                lr *= 0.95  # Stabilize
            elif brain.stagnation_counter > 5:
                lr *= 1.05  # Accelerate
            lr = np.clip(lr, 0.001, 1.0)
        lr_history.append(lr)
        
        # Trajectory Layer (S'): Train
        loss, grew = brain.train_step_with_salience(X, y, lr)
        losses.append(loss)
        neuron_counts.append(brain.hidden_dim)
        
        if grew:
            growth_events.append(epoch)
            lr = 0.2  # Boost LR to integrate new neurons
            if verbose:
                print(f"  [Epoch {epoch}] GROWTH -> {brain.hidden_dim} neurons (loss: {loss:.4f})")
        
        if verbose and epoch % 200 == 0:
            print(f"  [Epoch {epoch}] Loss: {loss:.4f}, Neurons: {brain.hidden_dim}, LR: {lr:.4f}")
    
    # Final evaluation
    predictions = brain.forward(X)
    metrics = compute_metrics(predictions, y)
    
    print("-"*60)
    print(f"FINAL RESULTS:")
    print(f"  Loss: {losses[-1]:.4f}")
    print(f"  Accuracy: {metrics['accuracy']*100:.1f}%")
    print(f"  Final Brain Size: {brain.hidden_dim} neurons")
    print(f"  Growth Events: {len(growth_events)}")
    print(f"  LR Adjustments: {len(lr_history)}")
    
    # Visualize
    viz = AGIVisualizer(save_dir="plots")
    viz.plot_comprehensive_dashboard(losses, brain, X, y, neuron_counts, lr_history, 
                                     save_name="growing_brain_dashboard.png")
    plt.show()
    
    return brain, losses, metrics


def run_sparse_topology_experiment(epochs: int = 2000, verbose: bool = True):
    """
    Experiment 2: Bio-Mimetic Sparse Core
    
    Network maintains fixed metabolic budget via pruning and neurogenesis.
    Demonstrates efficient learning with ~95% sparsity.
    """
    print("\n" + "="*60)
    print("EXPERIMENT 2: BIO-MIMETIC SPARSE TOPOLOGY")
    print("="*60)
    
    # Setup
    task = ConcentricCirclesTask()
    X, y = task.get_data()
    
    config = SparseMeshConfig(
        input_dim=2,
        hidden_dim=1000,  # Large capacity
        output_dim=1,
        density=0.05     # Only 5% active
    )
    brain = SparseSynapticMesh(config)
    
    print(f"Network: {config.hidden_dim} neurons, {config.density*100:.0f}% density")
    print(f"Task: Concentric Circles ({len(X)} samples)")
    print("-"*60)
    
    losses = []
    
    for epoch in range(epochs):
        loss = brain.train_step(X, y, lr=0.05)
        losses.append(loss)
        
        # Metabolic cycle every 50 epochs
        if epoch % 50 == 0 and epoch > 0:
            results = brain.sleep_and_rewire(pruning_rate=0.2)
            
            if verbose and epoch % 200 == 0:
                stats = brain.get_topology_stats()
                print(f"  [Epoch {epoch}] Loss: {loss:.4f}, "
                      f"Active: {stats['active_neurons']}/{stats['total_neurons']}, "
                      f"Sparsity: {stats['sparsity']*100:.1f}%")
    
    # Final evaluation
    predictions = brain.forward(X)
    metrics = compute_metrics(predictions, y)
    stats = brain.get_topology_stats()
    
    print("-"*60)
    print(f"FINAL RESULTS:")
    print(f"  Loss: {losses[-1]:.4f}")
    print(f"  Accuracy: {metrics['accuracy']*100:.1f}%")
    print(f"  Active Connections: {stats['total_connections']}")
    print(f"  Sparsity: {stats['sparsity']*100:.1f}%")
    print(f"  Active Neurons: {stats['active_neurons']}/{stats['total_neurons']}")
    
    # Visualize
    viz = AGIVisualizer(save_dir="plots")
    viz.plot_sparse_topology(
        brain.connection_history,
        brain.sparsity_history,
        brain.neuron_salience,
        save_name="sparse_topology.png"
    )
    viz.plot_decision_boundary(brain, X, y, 
                               title=f"Sparse Network (Sparsity: {stats['sparsity']*100:.1f}%)",
                               save_name="sparse_decision_boundary.png")
    plt.show()
    
    return brain, losses, metrics


def run_full_agi_demo():
    """
    Full three-layer recursive AGI demonstration.
    """
    print("\n" + "="*60)
    print("FULL AGI CORE DEMONSTRATION")
    print("="*60)
    print("Running both experiments to demonstrate recursive optimization...")
    
    # Run experiments
    print("\n[1/2] Growing Brain Experiment...")
    brain1, losses1, metrics1 = run_growing_brain_experiment(epochs=500, verbose=False)
    
    print("\n[2/2] Sparse Topology Experiment...")
    brain2, losses2, metrics2 = run_sparse_topology_experiment(epochs=1000, verbose=False)
    
    print("\n" + "="*60)
    print("SUMMARY")
    print("="*60)
    print(f"Growing Brain:  {metrics1['accuracy']*100:.1f}% accuracy, {brain1.hidden_dim} neurons")
    print(f"Sparse Network: {metrics2['accuracy']*100:.1f}% accuracy, {brain2.get_sparsity()*100:.1f}% sparse")
    print("\nPlots saved to ./plots/")


if __name__ == "__main__":
    import sys
    
    print("="*60)
    print("RECURSIVE SELF-OPTIMIZING AGI FRAMEWORK")
    print("="*60)
    print("\nAvailable experiments:")
    print("  1. Growing Brain (Self-modifying architecture)")
    print("  2. Sparse Topology (Bio-mimetic efficiency)")
    print("  3. Full Demo (Both experiments)")
    
    if len(sys.argv) > 1:
        choice = sys.argv[1]
    else:
        choice = input("\nSelect experiment (1/2/3) or press Enter for full demo: ").strip()
    
    if choice == "1":
        run_growing_brain_experiment()
    elif choice == "2":
        run_sparse_topology_experiment()
    else:
        run_full_agi_demo()
