"""
Stabilizer code QEC with real syndrome measurement and decoding.

Implements the Steane [[7,1,3]] code as a concrete example of discrete QEC.
Provides syndrome measurement, lookup-table decoding, and error correction
on density matrices.

References:
    [1] Steane, "Error correcting codes in quantum theory"
        Phys. Rev. Lett. 77, 793 (1996)
    [2] Nielsen & Chuang, Ch. 10 "Quantum Error Correction"
    [3] Gottesman, "Stabilizer Codes and Quantum Error Correction" (1997)
"""

import numpy as np
from scipy import linalg
from typing import List, Dict, Tuple, Optional
from dataclasses import dataclass
from itertools import product

from .fidelity import uhlmann_fidelity
from .lindblad import normalize_density_matrix


# Pauli matrices
I2 = np.eye(2, dtype=complex)
X = np.array([[0, 1], [1, 0]], dtype=complex)
Y = np.array([[0, -1j], [1j, 0]], dtype=complex)
Z = np.array([[1, 0], [0, -1]], dtype=complex)


def tensor_product(*mats: np.ndarray) -> np.ndarray:
    """Compute tensor product of multiple matrices."""
    result = mats[0]
    for m in mats[1:]:
        result = np.kron(result, m)
    return result


def pauli_on_qubit(pauli: np.ndarray, qubit: int, n_qubits: int) -> np.ndarray:
    """Apply a single-qubit Pauli to a specific qubit in an n-qubit system."""
    ops = [I2] * n_qubits
    ops[qubit] = pauli
    return tensor_product(*ops)


# === Steane [[7,1,3]] Code ===

def steane_stabilizers() -> List[np.ndarray]:
    """
    Construct the 6 stabilizer generators of the Steane [[7,1,3]] code.

    X-type stabilizers (from classical [7,4,3] Hamming code):
        S1 = I I I X X X X
        S2 = I X X I I X X
        S3 = X I X I X I X

    Z-type stabilizers:
        S4 = I I I Z Z Z Z
        S5 = I Z Z I I Z Z
        S6 = Z I Z I Z I Z
    """
    n = 7

    # X-type stabilizer support sets (from Hamming code parity check)
    x_supports = [
        [3, 4, 5, 6],  # S1
        [1, 2, 5, 6],  # S2
        [0, 2, 4, 6],  # S3
    ]

    # Z-type stabilizer support sets (same pattern)
    z_supports = [
        [3, 4, 5, 6],  # S4
        [1, 2, 5, 6],  # S5
        [0, 2, 4, 6],  # S6
    ]

    stabilizers = []

    for support in x_supports:
        ops = [I2] * n
        for q in support:
            ops[q] = X
        stabilizers.append(tensor_product(*ops))

    for support in z_supports:
        ops = [I2] * n
        for q in support:
            ops[q] = Z
        stabilizers.append(tensor_product(*ops))

    return stabilizers


def steane_logical_operators() -> Tuple[np.ndarray, np.ndarray]:
    """
    Logical X and Z operators for the Steane code.

    X_L = X^{otimes 7}
    Z_L = Z^{otimes 7}
    """
    X_L = tensor_product(*[X] * 7)
    Z_L = tensor_product(*[Z] * 7)
    return X_L, Z_L


def steane_encode(logical_state: np.ndarray) -> np.ndarray:
    """
    Encode a single logical qubit into the Steane [[7,1,3]] code.

    |0_L> = (1/sqrt(8)) * sum_{v in C} |v>
    where C is the [7,4,3] Hamming code (even weight codewords for |0>).

    |1_L> = X_L |0_L>

    Args:
        logical_state: 2-component vector [alpha, beta] for alpha|0>+beta|1>

    Returns:
        128-component encoded state vector (2^7 = 128)
    """
    # [7,4,3] Hamming code codewords (even weight subset for |0_L>)
    # The code has 2^4 = 16 codewords. |0_L> uses the even weight subset.
    hamming_codewords_even = [
        [0, 0, 0, 0, 0, 0, 0],
        [1, 0, 1, 0, 1, 0, 1],
        [0, 1, 1, 0, 0, 1, 1],
        [1, 1, 0, 0, 1, 1, 0],
        [0, 0, 0, 1, 1, 1, 1],
        [1, 0, 1, 1, 0, 1, 0],
        [0, 1, 1, 1, 1, 0, 0],
        [1, 1, 0, 1, 0, 0, 1],
    ]

    # |0_L>
    dim = 2 ** 7
    zero_L = np.zeros(dim, dtype=complex)
    for cw in hamming_codewords_even:
        index = sum(b * 2 ** (6 - i) for i, b in enumerate(cw))
        zero_L[index] = 1.0
    zero_L /= np.linalg.norm(zero_L)

    # |1_L> = X_L |0_L>
    X_L, _ = steane_logical_operators()
    one_L = X_L @ zero_L

    # Encoded state
    alpha, beta = logical_state[0], logical_state[1]
    encoded = alpha * zero_L + beta * one_L
    encoded /= np.linalg.norm(encoded)

    return encoded


def measure_syndrome(
    state: np.ndarray,
    stabilizers: List[np.ndarray],
) -> np.ndarray:
    """
    Measure the syndrome of a quantum state.

    For pure states: syndrome bit k = 0 if <psi|S_k|psi> = +1, else 1.
    For density matrices: syndrome bit k based on Tr(S_k * rho).

    Args:
        state: State vector (d,) or density matrix (d x d)
        stabilizers: List of stabilizer operators

    Returns:
        Binary syndrome vector (n_stabilizers,)
    """
    syndrome = np.zeros(len(stabilizers), dtype=int)

    if state.ndim == 1:
        # Pure state
        for k, S in enumerate(stabilizers):
            expectation = np.real(state.conj() @ S @ state)
            syndrome[k] = 0 if expectation > 0 else 1
    else:
        # Density matrix
        for k, S in enumerate(stabilizers):
            expectation = np.real(np.trace(S @ state))
            syndrome[k] = 0 if expectation > 0 else 1

    return syndrome


def build_steane_decoder() -> Dict[tuple, Tuple[str, int]]:
    """
    Build lookup table decoder for the Steane [[7,1,3]] code.

    Maps each 6-bit syndrome to the most likely single-qubit error
    (error type and qubit index).

    The syndrome pattern directly identifies the error location via
    the Hamming code structure.

    Returns:
        Dict mapping syndrome tuple -> (error_type, qubit_index)
        error_type is 'X', 'Z', 'Y', or 'I' (no error)
    """
    decoder = {}

    # No error
    decoder[(0, 0, 0, 0, 0, 0)] = ('I', -1)

    # X errors: only Z-stabilizers (S4,S5,S6 = indices 3,4,5) are affected
    # Z-syndrome bits identify qubit via Hamming code
    z_support = {
        3: [3, 4, 5, 6],  # S4
        4: [1, 2, 5, 6],  # S5
        5: [0, 2, 4, 6],  # S6
    }

    for qubit in range(7):
        # X error on this qubit: which Z-stabilizers anticommute?
        z_syndrome = [0, 0, 0]
        for s_idx, (s_key, support) in enumerate(z_support.items()):
            if qubit in support:
                z_syndrome[s_idx] = 1

        x_syndrome = [0, 0, 0]  # X errors don't affect X-stabilizers
        full_syndrome = tuple(x_syndrome + z_syndrome)
        if full_syndrome != (0, 0, 0, 0, 0, 0):
            decoder[full_syndrome] = ('X', qubit)

    # Z errors: only X-stabilizers (S1,S2,S3 = indices 0,1,2) are affected
    x_support = {
        0: [3, 4, 5, 6],  # S1
        1: [1, 2, 5, 6],  # S2
        2: [0, 2, 4, 6],  # S3
    }

    for qubit in range(7):
        x_syndrome = [0, 0, 0]
        for s_idx, (s_key, support) in enumerate(x_support.items()):
            if qubit in support:
                x_syndrome[s_idx] = 1

        z_syndrome = [0, 0, 0]  # Z errors don't affect Z-stabilizers
        full_syndrome = tuple(x_syndrome + z_syndrome)
        if full_syndrome != (0, 0, 0, 0, 0, 0):
            if full_syndrome not in decoder:
                decoder[full_syndrome] = ('Z', qubit)

    # Y errors = X + Z: both sets of stabilizers affected
    for qubit in range(7):
        x_syndrome = [0, 0, 0]
        for s_idx, (s_key, support) in enumerate(x_support.items()):
            if qubit in support:
                x_syndrome[s_idx] = 1

        z_syndrome = [0, 0, 0]
        for s_idx, (s_key, support) in enumerate(z_support.items()):
            if qubit in support:
                z_syndrome[s_idx] = 1

        full_syndrome = tuple(x_syndrome + z_syndrome)
        if full_syndrome not in decoder:
            decoder[full_syndrome] = ('Y', qubit)

    return decoder


def decode_and_correct(
    state: np.ndarray,
    syndrome: np.ndarray,
    decoder: Dict[tuple, Tuple[str, int]],
) -> np.ndarray:
    """
    Decode syndrome and apply correction to the state.

    Args:
        state: State vector or density matrix
        syndrome: Binary syndrome vector
        decoder: Lookup table from build_steane_decoder()

    Returns:
        Corrected state
    """
    key = tuple(syndrome)
    if key not in decoder:
        return state  # Unknown syndrome, no correction

    error_type, qubit = decoder[key]

    if error_type == 'I':
        return state  # No error detected

    n_qubits = 7

    # Select correction operator
    if error_type == 'X':
        correction = pauli_on_qubit(X, qubit, n_qubits)
    elif error_type == 'Z':
        correction = pauli_on_qubit(Z, qubit, n_qubits)
    elif error_type == 'Y':
        correction = pauli_on_qubit(Y, qubit, n_qubits)
    else:
        return state

    if state.ndim == 1:
        return correction @ state
    else:
        return correction @ state @ correction.conj().T


def apply_single_qubit_error(
    state: np.ndarray,
    error_type: str,
    qubit: int,
    n_qubits: int = 7,
) -> np.ndarray:
    """
    Apply a single-qubit Pauli error to a state.

    Args:
        state: State vector or density matrix
        error_type: 'X', 'Y', or 'Z'
        qubit: Qubit index (0-indexed)
        n_qubits: Total number of qubits

    Returns:
        Errored state
    """
    pauli = {'X': X, 'Y': Y, 'Z': Z}[error_type]
    error_op = pauli_on_qubit(pauli, qubit, n_qubits)

    if state.ndim == 1:
        return error_op @ state
    else:
        return error_op @ state @ error_op.conj().T


def estimate_logical_error_rate(
    physical_error_rate: float,
    n_qubits: int = 7,
    n_trials: int = 1000,
    rng: Optional[np.random.Generator] = None,
) -> float:
    """
    Estimate logical error rate by Monte Carlo simulation.

    For each trial:
    1. Encode |0_L>
    2. Apply random single-qubit errors with probability p
    3. Measure syndrome and decode
    4. Check if logical state is preserved

    Args:
        physical_error_rate: Per-qubit error probability
        n_qubits: Number of physical qubits (7 for Steane)
        n_trials: Number of Monte Carlo trials
        rng: Random number generator

    Returns:
        Estimated logical error rate (fraction of trials with logical error)
    """
    if rng is None:
        rng = np.random.default_rng()

    stabilizers = steane_stabilizers()
    decoder = build_steane_decoder()
    X_L, Z_L = steane_logical_operators()

    # Encode |0_L>
    logical_zero = steane_encode(np.array([1.0, 0.0], dtype=complex))

    logical_errors = 0

    for _ in range(n_trials):
        state = logical_zero.copy()

        # Apply random errors
        for qubit in range(n_qubits):
            if rng.random() < physical_error_rate:
                # Random Pauli error
                error = rng.choice(['X', 'Y', 'Z'])
                state = apply_single_qubit_error(state, error, qubit, n_qubits)

        # Syndrome measurement and correction
        syndrome = measure_syndrome(state, stabilizers)
        state = decode_and_correct(state, syndrome, decoder)

        # Check if logical state is preserved
        # |0_L> should give +1 eigenvalue of Z_L
        z_expectation = np.real(state.conj() @ Z_L @ state)
        if z_expectation < 0:
            logical_errors += 1

    return logical_errors / n_trials
