"""
Reproducibility utilities.

Ensures deterministic training for scientific reproducibility.
"""

import torch
import numpy as np
import random
from typing import Optional
from dataclasses import dataclass


def set_seed(seed: int):
    """
    Set random seeds for reproducibility.

    Args:
        seed: Random seed
    """
    random.seed(seed)
    np.random.seed(seed)
    torch.manual_seed(seed)
    torch.cuda.manual_seed_all(seed)


def make_deterministic(seed: int = 42, warn_only: bool = False):
    """
    Make PyTorch operations deterministic.

    Args:
        seed: Random seed
        warn_only: If True, only warn about non-deterministic operations
    """
    set_seed(seed)

    # PyTorch deterministic settings
    torch.backends.cudnn.deterministic = True
    torch.backends.cudnn.benchmark = False

    # Set deterministic algorithms
    if hasattr(torch, 'use_deterministic_algorithms'):
        try:
            torch.use_deterministic_algorithms(True, warn_only=warn_only)
        except Exception as e:
            print(f"Warning: Could not enable deterministic algorithms: {e}")

    # Environment variables for additional reproducibility
    import os
    os.environ['PYTHONHASHSEED'] = str(seed)
    os.environ['CUBLAS_WORKSPACE_CONFIG'] = ':4096:8'


@dataclass
class ReproducibilityConfig:
    """
    Configuration for reproducibility.

    Stores all settings needed to reproduce a training run.
    """
    seed: int = 42
    torch_version: str = torch.__version__
    cuda_available: bool = torch.cuda.is_available()
    cuda_version: Optional[str] = torch.version.cuda if torch.cuda.is_available() else None
    cudnn_version: Optional[int] = torch.backends.cudnn.version() if torch.cuda.is_available() else None
    device: str = 'cuda' if torch.cuda.is_available() else 'cpu'

    def apply(self):
        """Apply reproducibility settings."""
        make_deterministic(self.seed)

    def save(self, path: str):
        """Save config to file."""
        import json

        data = {
            'seed': self.seed,
            'torch_version': self.torch_version,
            'cuda_available': self.cuda_available,
            'cuda_version': self.cuda_version,
            'cudnn_version': self.cudnn_version,
            'device': self.device,
        }

        with open(path, 'w') as f:
            json.dump(data, f, indent=2)

    @classmethod
    def load(cls, path: str) -> 'ReproducibilityConfig':
        """Load config from file."""
        import json

        with open(path, 'r') as f:
            data = json.load(f)

        return cls(**data)

    def __str__(self) -> str:
        return (
            f"ReproducibilityConfig(\n"
            f"  seed={self.seed},\n"
            f"  torch_version={self.torch_version},\n"
            f"  cuda_available={self.cuda_available},\n"
            f"  cuda_version={self.cuda_version},\n"
            f"  cudnn_version={self.cudnn_version},\n"
            f"  device={self.device}\n"
            f")"
        )


class DeterministicDataLoader:
    """
    Wrapper for DataLoader that ensures deterministic iteration.

    Uses worker_init_fn to set seeds for each worker.
    """

    def __init__(self, dataloader: torch.utils.data.DataLoader, seed: int = 42):
        self.dataloader = dataloader
        self.seed = seed

        # Set worker_init_fn
        def worker_init_fn(worker_id):
            worker_seed = seed + worker_id
            np.random.seed(worker_seed)
            random.seed(worker_seed)

        self.dataloader.worker_init_fn = worker_init_fn

    def __iter__(self):
        return iter(self.dataloader)

    def __len__(self):
        return len(self.dataloader)


def get_device(device: Optional[str] = None) -> torch.device:
    """
    Get torch device with proper error handling.

    Args:
        device: Device string ('cuda', 'cpu', 'cuda:0', etc.)

    Returns:
        torch_device: PyTorch device
    """
    if device is None:
        device = 'cuda' if torch.cuda.is_available() else 'cpu'

    try:
        torch_device = torch.device(device)

        # Test device availability
        if torch_device.type == 'cuda':
            # Try to allocate a small tensor
            test_tensor = torch.zeros(1, device=torch_device)
            del test_tensor
            torch.cuda.empty_cache()

        return torch_device

    except Exception as e:
        print(f"Warning: Could not use device '{device}': {e}")
        print("Falling back to CPU")
        return torch.device('cpu')


def print_system_info():
    """Print system information for debugging."""
    print("=" * 60)
    print("System Information")
    print("=" * 60)
    print(f"PyTorch version: {torch.__version__}")
    print(f"CUDA available: {torch.cuda.is_available()}")

    if torch.cuda.is_available():
        print(f"CUDA version: {torch.version.cuda}")
        print(f"cuDNN version: {torch.backends.cudnn.version()}")
        print(f"Number of GPUs: {torch.cuda.device_count()}")

        for i in range(torch.cuda.device_count()):
            print(f"GPU {i}: {torch.cuda.get_device_name(i)}")
            print(f"  Memory: {torch.cuda.get_device_properties(i).total_memory / 1e9:.2f} GB")

    print(f"Number of CPU threads: {torch.get_num_threads()}")
    print("=" * 60)
