"""Deterministic quadratic benchmark suite for wormhole proof."""

from __future__ import annotations

from dataclasses import dataclass
from typing import Dict, Tuple

import numpy as np

from ..bridge import CCABridge


@dataclass(frozen=True)
class TwinQuadraticConfig:
    dim_donor: int = 8
    dim_receiver: int = 10
    shared_rank: int = 4
    history_samples: int = 256
    noise: float = 0.05
    cca_rank: int = 4
    cca_regularization: float = 1e-4
    trials: int = 64
    step_scale: float = 0.25


class TwinQuadraticBenchmark:
    """Analytic objective demonstrating bridge efficacy."""

    def __init__(self, cfg: TwinQuadraticConfig, seed: int = 7) -> None:
        self.cfg = cfg
        self.rng = np.random.default_rng(seed)
        self.true_map = self._sample_ground_truth()

    def _sample_ground_truth(self) -> np.ndarray:
        u = self.rng.normal(size=(self.cfg.dim_donor, self.cfg.shared_rank))
        v = self.rng.normal(size=(self.cfg.dim_receiver, self.cfg.shared_rank))
        return u @ v.T

    def generate_histories(self) -> Tuple[np.ndarray, np.ndarray]:
        donor = self.rng.normal(size=(self.cfg.history_samples, self.cfg.dim_donor))
        noise = self.cfg.noise * self.rng.normal(size=(self.cfg.history_samples, self.cfg.dim_receiver))
        receiver = donor @ self.true_map + noise
        return donor, receiver

    def evaluate(self) -> Dict[str, float]:
        donor_hist, receiver_hist = self.generate_histories()
        bridge = CCABridge(rank=self.cfg.cca_rank, regularization=self.cfg.cca_regularization)
        fit = bridge.fit(donor_hist, receiver_hist)

        errors = []
        baseline_losses = []
        wormhole_losses = []

        for _ in range(self.cfg.trials):
            delta = self.rng.normal(size=self.cfg.dim_donor)
            target = self.true_map.T @ delta
            pred = bridge.map_delta(delta)
            errors.append(np.linalg.norm(pred - target))

            state_a = self.rng.normal(size=self.cfg.dim_donor)
            state_b = self.rng.normal(size=self.cfg.dim_receiver)
            optimal_b = self.true_map.T @ state_a

            loss_before = 0.5 * (np.linalg.norm(state_a) ** 2 + np.linalg.norm(state_b - optimal_b) ** 2)
            mapped = bridge.map_delta(delta)
            state_b_updated = state_b + self.cfg.step_scale * mapped
            loss_after = 0.5 * (np.linalg.norm(state_a) ** 2 + np.linalg.norm(state_b_updated - optimal_b) ** 2)

            baseline_losses.append(loss_before)
            wormhole_losses.append(loss_after)

        loss_gain = np.array(baseline_losses) - np.array(wormhole_losses)

        return {
            "alignment_error": fit.alignment_error,
            "cca_cond_donor": fit.cond_donor,
            "cca_cond_receiver": fit.cond_receiver,
            "transport_mean_error": float(np.mean(errors)),
            "transport_std_error": float(np.std(errors)),
            "baseline_loss_mean": float(np.mean(baseline_losses)),
            "wormhole_loss_mean": float(np.mean(wormhole_losses)),
            "loss_gain_mean": float(np.mean(loss_gain)),
            "loss_gain_std": float(np.std(loss_gain)),
            "trials": float(self.cfg.trials),
        }
