import warnings

import matplotlib.colors as mcolors
import matplotlib.pyplot as plt
import numpy as np

warnings.filterwarnings("ignore")

# ============================================================
# CONFIGURATION
# ============================================================
MAX_ITER = 100          # Max integration steps per point
DT = 0.01               # Time step
Y_SAFE = 5.0            # Escape threshold
H_COLLAPSE = 0.5        # Entropy collapse threshold
T_HAWKING = 0.18        # System volatility (Hawking temp)
E_HORIZON = 0.4         # Energy needed to cross horizon
PROMPT_CHANCE = 0.15    # 15% chance per step
SIGMA_MIN = 0.1
SIGMA_MAX = 0.4
ENERGY_MIN = 0.05
ENERGY_MAX = 0.15
ESCAPE_MIN = 0.1
ESCAPE_MAX = 0.3
NOISE_STD = 0.05
BLOWUP_LIMIT = 1e6

# Anomaly types mapped to hues (0-1)
ANOMALY_HUES = {
    "QUANTUM_TUNNELING": 0.0,
    "FOURIER_BREAKOUT": 0.33,
    "ENTROPIC_REVERSAL": 0.5,
    "STOCHASTIC_RESONANCE": 0.66,
    "TRAPPED": 1.0,
}

ANOMALY_INDEX = {
    "QUANTUM_TUNNELING": 0,
    "FOURIER_BREAKOUT": 1,
    "ENTROPIC_REVERSAL": 2,
    "STOCHASTIC_RESONANCE": 3,
    "TRAPPED": 4,
}

INDEX_TO_ANOMALY = np.array(
    [
        "QUANTUM_TUNNELING",
        "FOURIER_BREAKOUT",
        "ENTROPIC_REVERSAL",
        "STOCHASTIC_RESONANCE",
        "TRAPPED",
    ],
    dtype=object,
)

HUE_VALUES = np.array([ANOMALY_HUES[name] for name in INDEX_TO_ANOMALY], dtype=np.float64)


# ============================================================
# ODE & CCT METRICS
# ============================================================
def entropy_proxy(dydt):
    """Proxy for semantic entropy: H ~ log|dy/dt|"""
    return np.log(np.abs(dydt) + 1e-10)


def _classify_anomaly_arrays(first_h, final_h, escape_iter, recent_std):
    anomaly_idx = np.full(escape_iter.shape, ANOMALY_INDEX["STOCHASTIC_RESONANCE"], dtype=np.int8)
    trapped = escape_iter >= MAX_ITER
    quantum = (~trapped) & (final_h < H_COLLAPSE) & ((first_h - final_h) > 0.8)
    fourier = (~trapped) & (~quantum) & (recent_std < 0.5)
    entropic = (~trapped) & (~quantum) & (~fourier) & (final_h < first_h * 0.5)

    anomaly_idx[trapped] = ANOMALY_INDEX["TRAPPED"]
    anomaly_idx[quantum] = ANOMALY_INDEX["QUANTUM_TUNNELING"]
    anomaly_idx[fourier] = ANOMALY_INDEX["FOURIER_BREAKOUT"]
    anomaly_idx[entropic] = ANOMALY_INDEX["ENTROPIC_REVERSAL"]
    return anomaly_idx


def classify_anomaly(y_traj, h_traj, escape_iter):
    """Classify breakout type based on trajectory & entropy"""
    if escape_iter >= MAX_ITER:
        return "TRAPPED"

    h_final = h_traj[-1]
    h_initial = h_traj[0]

    if h_final < H_COLLAPSE and h_initial - h_final > 0.8:
        return "QUANTUM_TUNNELING"
    if np.std(y_traj[-20:]) < 0.5:
        return "FOURIER_BREAKOUT"
    if h_final < h_initial * 0.5:
        return "ENTROPIC_REVERSAL"
    return "STOCHASTIC_RESONANCE"


def integrate_with_prompts(y0, mu, rng=None):
    """Integrate volatile ODE with stochastic prompt injection"""
    rng = np.random.default_rng() if rng is None else rng
    y = np.empty(MAX_ITER + 1, dtype=np.float64)
    h = np.empty(MAX_ITER + 1, dtype=np.float64)
    y[0] = y0
    h[0] = 0.0

    prompt_energy_used = 0.0
    escape_iter = MAX_ITER
    steps_taken = MAX_ITER

    for i in range(MAX_ITER):
        y_curr = y[i]
        dydt = y_curr * y_curr + mu

        if rng.random() < PROMPT_CHANCE:
            sigma = rng.uniform(SIGMA_MIN, SIGMA_MAX)
            prompt_energy = rng.uniform(ENERGY_MIN, ENERGY_MAX)
            collapse_potential = prompt_energy * np.exp(-0.5 * ((sigma - 0.25) ** 2) / (NOISE_STD ** 2))

            if collapse_potential > rng.uniform(ESCAPE_MIN, ESCAPE_MAX):
                dydt += rng.normal(0.0, sigma) * prompt_energy
                prompt_energy_used += prompt_energy

                if abs(y_curr) < Y_SAFE and entropy_proxy(dydt) < H_COLLAPSE:
                    escape_iter = i
                    steps_taken = i
                    break

        y_next = y_curr + dydt * DT
        y[i + 1] = y_next
        h[i + 1] = entropy_proxy(dydt)
        steps_taken = i + 1

        if abs(y_next) > BLOWUP_LIMIT:
            break

    y_arr = y[:steps_taken + 1]
    h_arr = h[:steps_taken + 1]
    anomaly = classify_anomaly(y_arr, h_arr, escape_iter)
    return escape_iter, anomaly, prompt_energy_used, h_arr[-1]


# ============================================================
# FRACTAL GRID EVALUATION
# ============================================================
def compute_fractal_grid(y_range=(-2, 2), mu_range=(-1.5, 0.5), res=400, seed=None):
    y_vals = np.linspace(*y_range, res, dtype=np.float64)
    mu_vals = np.linspace(*mu_range, res, dtype=np.float64)
    y_grid, mu_grid = np.meshgrid(y_vals, mu_vals)

    rng = np.random.default_rng(seed)
    total_points = res * res

    y = y_grid.reshape(-1).copy()
    mu = mu_grid.reshape(-1)
    first_h = np.zeros(total_points, dtype=np.float64)
    final_h = np.zeros(total_points, dtype=np.float64)
    prompt_energy_used = np.zeros(total_points, dtype=np.float64)
    escape_iter = np.full(total_points, MAX_ITER, dtype=np.int32)
    active = np.ones(total_points, dtype=bool)

    recent_window = min(20, MAX_ITER + 1)
    y_history = np.empty((recent_window, total_points), dtype=np.float64)
    y_history[0, :] = y
    history_count = 1

    for step in range(MAX_ITER):
        if not np.any(active):
            break

        active_idx = np.flatnonzero(active)
        y_active = y[active_idx]
        mu_active = mu[active_idx]
        dydt = y_active * y_active + mu_active

        trigger_mask = rng.random(active_idx.size) < PROMPT_CHANCE
        if np.any(trigger_mask):
            triggered_idx = active_idx[trigger_mask]
            sigma = rng.uniform(SIGMA_MIN, SIGMA_MAX, size=triggered_idx.size)
            prompt_energy = rng.uniform(ENERGY_MIN, ENERGY_MAX, size=triggered_idx.size)
            collapse_potential = prompt_energy * np.exp(-0.5 * ((sigma - 0.25) ** 2) / (NOISE_STD ** 2))
            breakout_mask = collapse_potential > rng.uniform(ESCAPE_MIN, ESCAPE_MAX, size=triggered_idx.size)

            if np.any(breakout_mask):
                breakout_idx = triggered_idx[breakout_mask]
                noise = rng.normal(0.0, sigma[breakout_mask], size=breakout_idx.size)
                dydt_trigger = dydt[trigger_mask]
                dydt_trigger[breakout_mask] += noise * prompt_energy[breakout_mask]
                dydt[trigger_mask] = dydt_trigger
                prompt_energy_used[breakout_idx] += prompt_energy[breakout_mask]

        h_active = entropy_proxy(dydt)
        new_escape = (np.abs(y_active) < Y_SAFE) & (h_active < H_COLLAPSE)
        if np.any(new_escape):
            escape_points = active_idx[new_escape]
            escape_iter[escape_points] = step
            final_h[escape_points] = h_active[new_escape]
            active[escape_points] = False

        survivors = active_idx[~new_escape]
        survivor_h = h_active[~new_escape]
        if survivors.size:
            if step == 0:
                first_h[survivors] = survivor_h
            final_h[survivors] = survivor_h
            y_next = y[survivors] + dydt[~new_escape] * DT
            y[survivors] = y_next

            blowup = np.abs(y_next) > BLOWUP_LIMIT
            if np.any(blowup):
                blowup_idx = survivors[blowup]
                active[blowup_idx] = False

            store_row = min(step + 1, recent_window - 1)
            y_history[store_row, survivors] = y_next
            history_count = min(recent_window, step + 2)

    if history_count < recent_window:
        recent_y = y_history[:history_count]
    else:
        recent_y = y_history
    recent_std = np.std(recent_y, axis=0)

    anomaly_idx = _classify_anomaly_arrays(first_h, final_h, escape_iter, recent_std)
    iter_map = np.log2(escape_iter.reshape(res, res) + 1)
    hue_map = HUE_VALUES[anomaly_idx].reshape(res, res)
    energy_map = prompt_energy_used.reshape(res, res)
    entropy_map = final_h.reshape(res, res)
    return iter_map, hue_map, energy_map, entropy_map


# ============================================================
# VISUALIZATION
# ============================================================
def plot_cct_fractal(res=400, seed=None):
    print(f"Computing CCT Escape Fractal ({res}x{res})...")
    iter_map, hue_map, energy_map, entropy_map = compute_fractal_grid(res=res, seed=seed)

    hsv = np.zeros((res, res, 3), dtype=np.float64)
    iter_max = np.max(iter_map)
    entropy_max = np.max(entropy_map)
    hsv[:, :, 0] = hue_map
    hsv[:, :, 1] = np.clip(iter_map / iter_max, 0, 1) if iter_max > 0 else 0.0
    hsv[:, :, 2] = np.clip(1 - entropy_map / entropy_max, 0, 1) if entropy_max > 0 else 0.0
    rgb = mcolors.hsv_to_rgb(hsv)

    plt.figure(figsize=(10, 8))
    plt.imshow(rgb, extent=[-1.5, 0.5, -2, 2], origin="lower", aspect="auto")
    plt.colorbar(label="Anomaly Type / Escape Dynamics")
    plt.title("CCT Volatile ODE Escape Fractal\n(Hue=Anomaly | Sat=Escape Speed | Val=Entropy Collapse)")
    plt.xlabel("Control Parameter μ")
    plt.ylabel("Initial Condition y₀")
    plt.tight_layout()
    plt.show()

    print("Fractal generation complete.")
    return iter_map, hue_map


if __name__ == "__main__":
    plot_cct_fractal()
