"""
Structured-Factorial Federated Learning (Fixed)
================================================
Server = Teacher (full knowledge) | Clients = Students (partial knowledge)
"""

import torch
import torch.nn as nn
import torch.optim as optim
from torch.utils.data import DataLoader, Subset
from torchvision import datasets, transforms
import numpy as np
from dataclasses import dataclass, field
from typing import List, Tuple, Dict, Optional
import math


# ============================================================================
# STRUCTURED-FACTORIAL SERVER
# ============================================================================

class StructuredFactorialServer:
    """
    Server = Teacher with comprehensive knowledge
    - Full training set
    - Computes structured embeddings (topologically constrained)
    - Governs multiple client models
    """
    
    def __init__(self, n_layers: int, embedding_dim: int, gamma: float = 1.5):
        self.n_layers = n_layers
        self.embedding_dim = embedding_dim
        self.gamma = gamma
        
        # Structured-Factorial
        self.F_S = 1
        self.n_factorial = math.factorial(n_layers)
        self.confinement_ratio = self.F_S / self.n_factorial
        
        # Work parameters (high for server orchestration)
        self.lambda_base = 10.0
        self.kappa = 2.0
        
        # Damping factors per layer
        self.damping_factors = [1.0 / (1 + i * 0.5) for i in range(n_layers)]
        
        # Client registry
        self.registered_clients = {}
        
        print(f"Server: F_S={self.F_S}, R={self.confinement_ratio:.6f}, λ_base={self.lambda_base}")
    
    def compute_lambda(self, H_AS: float = 0.0) -> float:
        confinement = (self.n_factorial / self.F_S) ** self.gamma
        return self.lambda_base + self.kappa * H_AS * confinement
    
    def compute_work(self) -> float:
        return math.log(self.n_factorial) - math.log(self.F_S)
    
    def register_client(self, client_id: int, client_data_size: int) -> Dict:
        """Register client and compute its structured configuration"""
        self.registered_clients[client_id] = {
            'data_size': client_data_size,
            'structured_config': self._compute_client_config(client_data_size)
        }
        return self.registered_clients[client_id]
    
    def _compute_client_config(self, client_data_size: int) -> Dict:
        """Compute structured config based on data size (smaller = more guidance)"""
        scarcity_factor = 1.0 / (1 + math.log(1 + client_data_size))
        
        return {
            'embedding_scale': scarcity_factor,
            'damping_weights': [d * scarcity_factor for d in self.damping_factors],
            # Keep guidance width fixed so server embeddings match client input layers.
            'structured_dim': self.embedding_dim,
            'work_allocation': self.compute_work() * scarcity_factor
        }
    
    def generate_structured_embedding(self, features: torch.Tensor, 
                                      client_id: int) -> torch.Tensor:
        """Generate structured embedding for client (server's work)"""
        config = self.registered_clients.get(client_id, 
                                             self.registered_clients[0])['structured_config']
        lambda_sec = self.compute_lambda()
        
        batch_size = features.shape[0]
        structured_dim = config['structured_dim']
        
        # Project to structured space
        proj = nn.Linear(features.shape[-1], structured_dim).to(features.device)
        embedded = torch.relu(proj(features))
        
        # Apply damping
        damped = embedded.clone()
        for i in range(min(self.n_layers, embedded.shape[-1] // structured_dim)):
            damping = config['damping_weights'][i]
            start, end = i * structured_dim, (i + 1) * structured_dim
            if end <= embedded.shape[-1]:
                damped[:, start:end] = embedded[:, start:end] * damping
        
        damped = damped / (1 + lambda_sec * 0.05)
        return damped
    
    def distill_guidance(self, features: torch.Tensor, 
                        labels: torch.Tensor, client_id: int) -> Dict[str, torch.Tensor]:
        """Server distills knowledge into structured guidance"""
        config = self.registered_clients.get(client_id,
                                             self.registered_clients[0])['structured_config']
        structured_embedding = self.generate_structured_embedding(features, client_id).detach()
        
        return {
            'structured_embedding': structured_embedding,
            'damping_weights': torch.tensor(config['damping_weights']).to(features.device),
            'lambda_sec': self.compute_lambda(),
            'work_done': config['work_allocation']
        }


# ============================================================================
# CLIENT MODEL
# ============================================================================

@dataclass
class ClientConfig:
    """Configuration for a client"""
    client_id: int
    client_data_size: int  # Now properly included here
    input_size: int = 784
    num_classes: int = 10
    hidden_dims: List[int] = field(default_factory=lambda: [256, 128])
    n_layers: int = 2
    structured_dim: int = 32


class StructuredClient(nn.Module):
    """Client = Student with partial knowledge, receives structured guidance"""
    
    def __init__(self, config: ClientConfig, server: StructuredFactorialServer):
        super().__init__()
        self.config = config
        self.server = server
        
        # Build structured layers
        self.input_proj = nn.Linear(config.input_size, config.structured_dim)
        
        layer_dims = [config.structured_dim] + config.hidden_dims
        self.structured_layers = nn.ModuleList([
            nn.Linear(layer_dims[i], layer_dims[i + 1])
            for i in range(len(config.hidden_dims))
        ])
        
        self.classifier = nn.Linear(config.hidden_dims[-1], config.num_classes)
        
        # Get damping from server
        self.damping_weights = server.registered_clients[config.client_id]['structured_config']['damping_weights']
        
        self._init_layers()
    
    def _init_layers(self):
        for i, layer in enumerate(self.structured_layers):
            damping = self.damping_weights[i] if i < len(self.damping_weights) else 1.0
            nn.init.xavier_uniform_(layer.weight)
            layer.weight.data *= damping
            nn.init.zeros_(layer.bias)
    
    def forward(self, x: torch.Tensor, 
                use_guidance: bool = False,
                guidance: Optional[Dict] = None) -> Tuple[torch.Tensor, float]:
        """Forward with optional server guidance"""
        if use_guidance and guidance is not None:
            h = guidance['structured_embedding']
        else:
            h = self.input_proj(x)
            h = torch.relu(h)
        
        for i, layer in enumerate(self.structured_layers):
            h = layer(h)
            h = torch.relu(h)
            
            if use_guidance and guidance is not None:
                damping = guidance['damping_weights'][i].item()
                h = h * damping
        
        logits = self.classifier(h)
        work = guidance.get('work_done', 0.0) if guidance else 0.0
        
        return logits, work


# ============================================================================
# FEDERATED PIPELINE
# ============================================================================

@dataclass
class FederatedConfig:
    """Global configuration"""
    input_size: int = 784
    num_classes: int = 10
    hidden_dims: List[int] = field(default_factory=lambda: [256, 128])
    n_layers: int = 2
    embedding_dim: int = 64
    structured_dim: int = 32
    gamma: float = 1.5
    n_clients: int = 5
    client_data_fraction: float = 0.1
    batch_size: int = 64
    server_rounds: int = 8
    client_epochs: int = 3
    learning_rate: float = 0.001


class FederatedPipeline:
    def __init__(self, config: FederatedConfig):
        self.config = config
        self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
        self.data_root = "../../data"
        
        # Initialize server
        self.server = StructuredFactorialServer(
            n_layers=config.n_layers,
            embedding_dim=config.structured_dim,
            gamma=config.gamma
        )
        
        # Initialize clients
        self.clients = []
        self.client_loaders = []
        
        self._prepare_clients()
        
        self.history = {
            'round': [], 'server_work': [],
            'client_train_acc': [], 'client_test_acc': []
        }
    
    def _prepare_clients(self):
        """Prepare client datasets"""
        transform = transforms.Compose([
            transforms.ToTensor(),
            transforms.Normalize((0.1307,), (0.3081,))
        ])
        
        full_dataset = datasets.MNIST(root=self.data_root, train=True, download=True, transform=transform)
        n_total = len(full_dataset)
        
        n_per_client = int(n_total * self.config.client_data_fraction)
        indices = np.random.permutation(n_total)
        
        for i in range(self.config.n_clients):
            start, end = i * n_per_client, (i + 1) * n_per_client
            client_indices = indices[start:end]
            
            client_dataset = Subset(full_dataset, client_indices)
            client_loader = DataLoader(client_dataset, batch_size=self.config.batch_size, shuffle=True)
            self.client_loaders.append(client_loader)
            
            # Register client with server
            self.server.register_client(i, len(client_dataset))
            
            # Create client config
            client_config = ClientConfig(
                client_id=i,
                client_data_size=len(client_dataset),
                input_size=self.config.input_size,
                num_classes=self.config.num_classes,
                hidden_dims=self.config.hidden_dims,
                n_layers=self.config.n_layers,
                structured_dim=self.config.structured_dim
            )
            
            # Create client model
            client = StructuredClient(client_config, self.server).to(self.device)
            self.clients.append(client)
            
            print(f"Client {i}: {len(client_dataset)} samples")
        
        print(f"Server has {n_total} samples (teacher's full book)")
    
    def _get_test_loader(self):
        transform = transforms.Compose([
            transforms.ToTensor(),
            transforms.Normalize((0.1307,), (0.3081,))
        ])
        return DataLoader(
            datasets.MNIST(root=self.data_root, train=False, download=True, transform=transform),
            batch_size=self.config.batch_size, shuffle=False
        )
    
    def client_train(self, client_id: int) -> Tuple[float, float]:
        """Train single client with server guidance"""
        client = self.clients[client_id]
        client.train()
        
        optimizer = optim.Adam(client.parameters(), lr=self.config.learning_rate)
        total_loss, correct, total = 0, 0, 0
        
        for data, target in self.client_loaders[client_id]:
            data, target = data.to(self.device), target.to(self.device)
            data = data.view(-1, self.config.input_size)
            
            # Get guidance from server (teacher's work)
            guidance = self.server.distill_guidance(data, target, client_id)
            
            optimizer.zero_grad()
            output, work = client.forward(data, use_guidance=True, guidance=guidance)
            
            loss = nn.CrossEntropyLoss()(output, target) + 0.01 * work
            loss.backward()
            optimizer.step()
            
            total_loss += loss.item()
            correct += (output.argmax(1) == target).sum().item()
            total += target.size(0)
        
        return total_loss / len(self.client_loaders[client_id]), 100. * correct / total
    
    @torch.no_grad()
    def evaluate_all(self) -> float:
        """Evaluate all clients on test set"""
        test_loader = self._get_test_loader()
        
        all_correct, all_total = 0, 0
        for client in self.clients:
            client.eval()
            for data, target in test_loader:
                data, target = data.to(self.device), target.to(self.device)
                data = data.view(-1, self.config.input_size)
                
                output, _ = client.forward(data, use_guidance=True)
                all_correct += (output.argmax(1) == target).sum().item()
                all_total += target.size(0)
        
        return 100. * all_correct / all_total
    
    def federated_train(self):
        """Full federated learning loop"""
        print("\n" + "="*65)
        print("FEDERATED STRUCTURED-LEARNING")
        print("Server (Teacher) → Structured Guidance → Clients (Students)")
        print("="*65)
        
        for round_num in range(1, self.config.server_rounds + 1):
            # Server orchestrates (does work)
            server_work = self.server.compute_work()
            lambda_sec = self.server.compute_lambda()
            
            # Each client learns with guidance
            round_accs = []
            for client_id in range(self.config.n_clients):
                loss, acc = self.client_train(client_id)
                round_accs.append(acc)
            
            # Evaluate
            test_acc = self.evaluate_all()
            
            # Store history
            self.history['round'].append(round_num)
            self.history['server_work'].append(server_work)
            self.history['client_train_acc'].append(np.mean(round_accs))
            self.history['client_test_acc'].append(test_acc)
            
            print(f"Round {round_num:2d}/{self.config.server_rounds} | "
                  f"Client Avg: {np.mean(round_accs):.1f}% | "
                  f"Test: {test_acc:.1f}% | "
                  f"λ_sec: {lambda_sec:.2f}")
        
        self._summary()
        return self.history
    
    def _summary(self):
        print("\n" + "="*65)
        print("SUMMARY")
        print("="*65)
        print(f"Clients: {self.config.n_clients} (each with {100*self.config.client_data_fraction:.0f}% of data)")
        print(f"Server work: {sum(self.history['server_work']):.2f} total")
        print(f"Final Test Acc: {self.history['client_test_acc'][-1]:.2f}%")
    
    def ablation_study(self):
        """Compare WITH vs WITHOUT guidance"""
        print("\n" + "="*65)
        print("ABLATION: With Guidance vs Without")
        print("="*65)
        
        # Client WITH guidance
        client_id = 0
        loss, acc_with = self.client_train(client_id)
        print(f"Client WITH server guidance: {acc_with:.1f}%")
        
        # Reset and train WITHOUT guidance
        client_config = self.clients[0].config
        client_no_guide = StructuredClient(client_config, self.server).to(self.device)
        client_no_guide.train()
        
        optimizer = optim.Adam(client_no_guide.parameters(), lr=self.config.learning_rate)
        
        for data, target in self.client_loaders[0]:
            data, target = data.to(self.device), target.to(self.device)
            data = data.view(-1, self.config.input_size)
            
            optimizer.zero_grad()
            output, _ = client_no_guide.forward(data, use_guidance=False)
            loss = nn.CrossEntropyLoss()(output, target)
            loss.backward()
            optimizer.step()
        
        # Quick eval
        test_loader = self._get_test_loader()
        correct, total = 0, 0
        with torch.no_grad():
            for data, target in test_loader:
                data, target = data.to(self.device), target.to(self.device)
                data = data.view(-1, self.config.input_size)
                output, _ = client_no_guide.forward(data, use_guidance=False)
                correct += (output.argmax(1) == target).sum().item()
                total += target.size(0)
        
        acc_without = 100. * correct / total
        print(f"Client WITHOUT server guidance: {acc_without:.1f}%")
        print(f"Improvement: +{acc_with - acc_without:.1f}%")


# ============================================================================
# MAIN
# ============================================================================

def main():
    print("\n" + "="*70)
    print("STRUCTURED-FACTORIAL FEDERATED LEARNING")
    print("Server=Teacher (full book) | Clients=Students (partial books)")
    print("="*70)
    
    config = FederatedConfig(
        n_clients=5,
        client_data_fraction=0.1,
        server_rounds=8,
        client_epochs=3,
        hidden_dims=[256, 128],
        n_layers=2,
        structured_dim=32
    )
    
    pipeline = FederatedPipeline(config)
    history = pipeline.federated_train()
    pipeline.ablation_study()
    
    return pipeline, history


if __name__ == "__main__":
    pipeline, history = main()
