from __future__ import annotations

import math
from typing import Dict, List, Optional

import numpy as np

from wormhole_proof.core.wormhole_smooth import SmoothGeometry

__all__ = ["geometry_to_payload"]


def _tail(sequence: List[List[float]] | List[float] | None, limit: int) -> List[List[float]] | List[float] | None:
    if sequence is None:
        return None
    if isinstance(sequence, list) and sequence and isinstance(sequence[0], (float, int)):
        return sequence
    return sequence[-limit:] if sequence else []


def _compute_pca_basis(data: np.ndarray, components: int = 3) -> tuple[np.ndarray, np.ndarray]:
    if data.size == 0:
        raise ValueError("cannot compute PCA on empty data")
    mean = data.mean(axis=0, keepdims=True)
    centered = data - mean
    _, _, vh = np.linalg.svd(centered, full_matrices=False)
    basis = vh[:components]
    return mean.squeeze(0), basis


def _project_sequence(sequence: List[List[float]], mean: np.ndarray, basis: np.ndarray) -> List[List[float]]:
    if not sequence:
        return []
    arr = np.asarray(sequence, dtype=float)
    proj = (arr - mean) @ basis.T
    return proj.tolist()


def geometry_to_payload(geometry: SmoothGeometry, limit: int = 256) -> Optional[Dict[str, object]]:
    if geometry is None:
        return None

    raw_sequences = {
        "state_a": _tail(geometry.state_a, limit),
        "state_b": _tail(geometry.state_b, limit),
        "packet": _tail(geometry.packet, limit),
        "donor_mode": _tail(geometry.donor_mode, limit),
        "receiver_mode": _tail(geometry.receiver_mode, limit),
        "hp_decode": _tail(geometry.hp_decode, limit),
        "hp_decode_bits": _tail(geometry.hp_decode_bits, limit),
        "hp_target": geometry.hp_target,
        "state_a_norm": geometry.state_a_norm[-limit:],
        "state_b_norm": geometry.state_b_norm[-limit:],
    }

    stack_candidates = []
    for key in ("state_a", "state_b", "packet", "donor_mode", "receiver_mode"):
        seq = raw_sequences.get(key) or []
        if seq:
            stack_candidates.append(np.asarray(seq, dtype=float))

    projected: Optional[Dict[str, List[List[float]]]] = None
    if stack_candidates:
        data = np.vstack(stack_candidates)
        try:
            mean, basis = _compute_pca_basis(data)
            projected = {
                key: _project_sequence(raw_sequences.get(key) or [], mean, basis)
                for key in ("state_a", "state_b", "packet", "donor_mode", "receiver_mode")
            }
        except ValueError:
            projected = None

    payload: Dict[str, object] = {**raw_sequences}
    if projected:
        payload["projected"] = projected
    return payload
