"""
PyPhysDisc Core Module
======================
Core implementation of the co-evolutionary symbolic regression framework.
Designed for reusability across standard experiments.

Author: Ali Tozar
License: MIT
"""

import random
import operator
import numpy as np
from deap import base, creator, tools, gp

# --- SHARED WORKER CACHE ---
WORKER_CACHE = None

def init_worker(shared_data):
    """Initializer for multiprocessing to share data dictionary."""
    global WORKER_CACHE
    WORKER_CACHE = shared_data

def random_const():
    """Generates ephemeral random constants."""
    return random.uniform(-1, 1)

def eval_standard(individual, system_name, pset):
    """
    Standard fitness evaluation with linear scaling.
    Used for Experiment 1 (Noise) and Experiment 3 (Convergence).
    """
    try:
        w_idx = individual.window_idx
        if WORKER_CACHE is None: return (1e9,)
        
        data = WORKER_CACHE[system_name][w_idx]
        func = gp.compile(expr=individual, pset=pset)
        y_pred = func(*data['inputs'])
        
        # Validation checks
        if np.ndim(y_pred) == 0: y_pred = np.full_like(data['target'], y_pred)
        if np.isnan(y_pred).any() or np.isinf(y_pred).any(): return (1e9,)
        if np.var(y_pred) < 1e-10: return (1e9,)

        # Linear Scaling (Affine Transform)
        slope, intercept = np.polyfit(y_pred, data['target'], 1)
        y_opt = slope * y_pred + intercept
        
        # MSE & Parsimony
        mse = np.mean((y_opt - data['target'])**2)
        target_var = np.var(data['target'])
        norm_mse = mse / target_var if target_var > 1e-6 else mse
        
        penalty = len(individual) * 0.001 
        return (norm_mse + penalty,)
    except:
        return (1e9,)

class StandardOptimizer:
    """
    Standard Co-Evolutionary Optimizer.
    Configures DEAP toolbox for typical symbolic regression tasks.
    """
    def __init__(self, var_names, system_name, window_pool_size, use_square=True):
        self.var_names = var_names
        self.system_name = system_name
        self.w_pool_size = window_pool_size
        self.use_square = use_square
        self._setup()

    def _setup(self):
        # Clean up previous runs
        if hasattr(creator, "FitnessMin"): del creator.FitnessMin
        if hasattr(creator, "Individual"): del creator.Individual

        creator.create("FitnessMin", base.Fitness, weights=(-1.0,))
        creator.create("Individual", gp.PrimitiveTree, fitness=creator.FitnessMin, window_idx=None)

        self.pset = gp.PrimitiveSet("MAIN", len(self.var_names))
        self.pset.addPrimitive(operator.add, 2)
        self.pset.addPrimitive(operator.sub, 2)
        self.pset.addPrimitive(operator.mul, 2)
        self.pset.addPrimitive(operator.neg, 1)
        if self.use_square:
            self.pset.addPrimitive(np.square, 1)
        self.pset.addEphemeralConstant("rnd", random_const)
        for i, n in enumerate(self.var_names): self.pset.renameArguments(**{f"ARG{i}": n})

        self.toolbox = base.Toolbox()
        self.toolbox.register("expr", gp.genHalfAndHalf, pset=self.pset, min_=1, max_=4)
        
        def init_ind(pcls, expr_func):
            ind = tools.initIterate(pcls, expr_func)
            ind.window_idx = random.randint(0, self.w_pool_size - 1)
            return ind

        self.toolbox.register("individual", init_ind, creator.Individual, self.toolbox.expr)
        self.toolbox.register("population", tools.initRepeat, list, self.toolbox.individual)
        self.toolbox.register("compile", gp.compile, pset=self.pset)
        
        # Operators
        def mut_coupled(ind, expr, pset):
            if random.random() < 0.8: gp.mutUniform(ind, expr, pset)
            else: ind.window_idx = random.randint(0, self.w_pool_size - 1)
            return ind,

        def mate_coupled(ind1, ind2):
            gp.cxOnePoint(ind1, ind2)
            if random.random() < 0.5: ind1.window_idx, ind2.window_idx = ind2.window_idx, ind1.window_idx
            return ind1, ind2

        self.toolbox.register("select", tools.selDoubleTournament, fitness_size=7, parsimony_size=1.2, fitness_first=True)
        self.toolbox.register("mate", mate_coupled)
        self.toolbox.register("mutate", mut_coupled, expr=self.toolbox.expr, pset=self.pset)
        self.toolbox.register("evaluate", eval_standard, system_name=self.system_name, pset=self.pset)